From ae1cb829f44283e51e49cef9bc5e42a2751f7571 Mon Sep 17 00:00:00 2001 From: Sakshi Patil Date: Sun, 23 Aug 2026 19:54:24 +0530 Subject: [PATCH] ix(postgres): randomize test_reader username to prevent CI collisions Resolves #5972. Appended a random UUID suffix to the Postgres test user creation logic to ensure parallel integration tests are completely isolated and do not collide. Signed-off-by: sakshipatil-hue --- .github/scripts/get_scm_version.py | 4 +- benchmarks/lsp_render_model_bench.py | 59 +- .../custom_materializations/custom_kind.py | 6 +- examples/sushi/config.py | 38 +- examples/sushi/macros/macros.py | 8 +- examples/sushi/models/disabled.py | 1 + examples/sushi/models/items.py | 4 +- examples/sushi/models/order_items.py | 6 +- examples/sushi/models/orders.py | 4 +- examples/sushi/models/raw_marketing.py | 8 +- examples/sushi/models/waiters.py | 4 +- examples/sushi/signals/__init__.py | 2 +- examples/sushi_dlt/sushi_pipeline.py | 8 +- examples/wursthall/models/db/order_f.py | 10 +- .../wursthall/models/src/customer_details.py | 3 +- .../models/src/order_item_details.py | 13 +- pdoc/cli.py | 7 +- sqlmesh/__init__.py | 35 +- sqlmesh/cli/__init__.py | 11 +- sqlmesh/cli/main.py | 56 +- sqlmesh/cli/project_init.py | 86 +- sqlmesh/core/_typing.py | 4 +- sqlmesh/core/analytics/__init__.py | 7 +- sqlmesh/core/analytics/collector.py | 59 +- sqlmesh/core/analytics/dispatcher.py | 24 +- sqlmesh/core/audit/__init__.py | 13 +- sqlmesh/core/audit/definition.py | 80 +- sqlmesh/core/config/__init__.py | 95 +- sqlmesh/core/config/base.py | 8 +- sqlmesh/core/config/common.py | 14 +- sqlmesh/core/config/connection.py | 226 +- sqlmesh/core/config/gateway.py | 11 +- sqlmesh/core/config/linter.py | 1 - sqlmesh/core/config/loader.py | 27 +- sqlmesh/core/config/model.py | 15 +- sqlmesh/core/config/root.py | 93 +- sqlmesh/core/config/run.py | 4 +- sqlmesh/core/config/scheduler.py | 34 +- sqlmesh/core/console.py | 743 ++-- sqlmesh/core/context.py | 455 ++- sqlmesh/core/context_diff.py | 62 +- sqlmesh/core/dialect.py | 201 +- sqlmesh/core/engine_adapter/__init__.py | 12 +- sqlmesh/core/engine_adapter/_typing.py | 10 +- sqlmesh/core/engine_adapter/athena.py | 125 +- sqlmesh/core/engine_adapter/base.py | 544 ++- sqlmesh/core/engine_adapter/base_postgres.py | 20 +- sqlmesh/core/engine_adapter/bigquery.py | 225 +- sqlmesh/core/engine_adapter/clickhouse.py | 169 +- sqlmesh/core/engine_adapter/databricks.py | 78 +- sqlmesh/core/engine_adapter/duckdb.py | 53 +- sqlmesh/core/engine_adapter/fabric.py | 108 +- sqlmesh/core/engine_adapter/mixins.py | 87 +- sqlmesh/core/engine_adapter/mssql.py | 141 +- sqlmesh/core/engine_adapter/mysql.py | 40 +- sqlmesh/core/engine_adapter/postgres.py | 20 +- sqlmesh/core/engine_adapter/redshift.py | 73 +- sqlmesh/core/engine_adapter/risingwave.py | 24 +- sqlmesh/core/engine_adapter/shared.py | 22 +- sqlmesh/core/engine_adapter/snowflake.py | 175 +- sqlmesh/core/engine_adapter/spark.py | 122 +- sqlmesh/core/engine_adapter/starrocks.py | 292 +- sqlmesh/core/engine_adapter/trino.py | 83 +- sqlmesh/core/environment.py | 26 +- sqlmesh/core/janitor.py | 43 +- sqlmesh/core/lineage.py | 8 +- sqlmesh/core/linter/definition.py | 13 +- sqlmesh/core/linter/helpers.py | 18 +- sqlmesh/core/linter/rule.py | 8 +- sqlmesh/core/linter/rules/builtin.py | 47 +- sqlmesh/core/loader.py | 110 +- sqlmesh/core/macros.py | 212 +- sqlmesh/core/metric/__init__.py | 10 +- sqlmesh/core/metric/definition.py | 20 +- sqlmesh/core/metric/rewriter.py | 16 +- sqlmesh/core/model/__init__.py | 84 +- sqlmesh/core/model/cache.py | 28 +- sqlmesh/core/model/common.py | 148 +- sqlmesh/core/model/decorator.py | 54 +- sqlmesh/core/model/definition.py | 398 ++- sqlmesh/core/model/kind.py | 174 +- sqlmesh/core/model/meta.py | 164 +- sqlmesh/core/model/schema.py | 12 +- sqlmesh/core/model/seed.py | 23 +- sqlmesh/core/node.py | 47 +- sqlmesh/core/notification_target.py | 50 +- sqlmesh/core/plan/__init__.py | 17 +- sqlmesh/core/plan/builder.py | 225 +- sqlmesh/core/plan/common.py | 38 +- sqlmesh/core/plan/definition.py | 59 +- sqlmesh/core/plan/evaluator.py | 87 +- sqlmesh/core/plan/explainer.py | 117 +- sqlmesh/core/plan/stages.py | 93 +- sqlmesh/core/reference.py | 8 +- sqlmesh/core/renderer.py | 114 +- sqlmesh/core/scheduler.py | 155 +- sqlmesh/core/schema_diff.py | 115 +- sqlmesh/core/schema_loader.py | 4 +- sqlmesh/core/selector.py | 60 +- sqlmesh/core/signal.py | 12 +- sqlmesh/core/snapshot/__init__.py | 88 +- sqlmesh/core/snapshot/cache.py | 14 +- sqlmesh/core/snapshot/categorizer.py | 4 +- sqlmesh/core/snapshot/definition.py | 344 +- sqlmesh/core/snapshot/evaluator.py | 429 ++- sqlmesh/core/snapshot/execution_tracker.py | 3 +- sqlmesh/core/state_sync/__init__.py | 11 +- sqlmesh/core/state_sync/base.py | 62 +- sqlmesh/core/state_sync/cache.py | 22 +- sqlmesh/core/state_sync/common.py | 59 +- sqlmesh/core/state_sync/db/environment.py | 68 +- sqlmesh/core/state_sync/db/facade.py | 147 +- sqlmesh/core/state_sync/db/interval.py | 85 +- sqlmesh/core/state_sync/db/migrator.py | 74 +- sqlmesh/core/state_sync/db/snapshot.py | 134 +- sqlmesh/core/state_sync/db/utils.py | 20 +- sqlmesh/core/state_sync/db/version.py | 14 +- sqlmesh/core/state_sync/export_import.py | 68 +- sqlmesh/core/table_diff.py | 137 +- sqlmesh/core/test/__init__.py | 10 +- sqlmesh/core/test/definition.py | 213 +- sqlmesh/core/test/discovery.py | 5 +- sqlmesh/core/test/runner.py | 29 +- sqlmesh/core/user.py | 10 +- sqlmesh/dbt/__init__.py | 8 +- sqlmesh/dbt/adapter.py | 141 +- sqlmesh/dbt/basemodel.py | 47 +- sqlmesh/dbt/builtin.py | 46 +- sqlmesh/dbt/column.py | 6 +- sqlmesh/dbt/common.py | 18 +- sqlmesh/dbt/context.py | 50 +- sqlmesh/dbt/loader.py | 102 +- sqlmesh/dbt/manifest.py | 148 +- sqlmesh/dbt/model.py | 124 +- sqlmesh/dbt/package.py | 24 +- sqlmesh/dbt/profile.py | 16 +- sqlmesh/dbt/project.py | 19 +- sqlmesh/dbt/relation.py | 1 - sqlmesh/dbt/seed.py | 4 +- sqlmesh/dbt/source.py | 4 +- sqlmesh/dbt/target.py | 65 +- sqlmesh/dbt/test.py | 36 +- sqlmesh/dbt/util.py | 6 +- sqlmesh/engines/spark/db_api/spark_session.py | 15 +- sqlmesh/integrations/dlt.py | 32 +- sqlmesh/integrations/github/cicd/command.py | 54 +- sqlmesh/integrations/github/cicd/config.py | 14 +- .../integrations/github/cicd/controller.py | 180 +- sqlmesh/integrations/slack.py | 10 +- sqlmesh/lsp/api.py | 8 +- sqlmesh/lsp/completions.py | 39 +- sqlmesh/lsp/context.py | 43 +- sqlmesh/lsp/custom.py | 3 +- sqlmesh/lsp/errors.py | 9 +- sqlmesh/lsp/helpers.py | 8 +- sqlmesh/lsp/hints.py | 2 +- sqlmesh/lsp/main.py | 305 +- sqlmesh/lsp/reference.py | 67 +- sqlmesh/lsp/rename.py | 29 +- sqlmesh/lsp/tests_ranges.py | 12 +- sqlmesh/lsp/uri.py | 3 +- sqlmesh/magics.py | 105 +- sqlmesh/migrations/v0000_baseline.py | 9 +- sqlmesh/migrations/v0063_change_signals.py | 2 +- .../v0064_join_when_matched_strings.py | 2 +- .../v0069_update_dev_table_suffix.py | 14 +- .../v0071_add_dev_version_to_intervals.py | 9 +- ...073_remove_symbolic_disable_restatement.py | 3 +- .../migrations/v0075_remove_validate_query.py | 3 +- ...v0078_warn_if_non_migratable_python_env.py | 8 +- .../migrations/v0081_update_partitioned_by.py | 3 +- .../migrations/v0085_deterministic_repr.py | 7 +- .../v0086_check_deterministic_bug.py | 5 +- .../v0087_normalize_blueprint_variables.py | 4 +- ...88_warn_about_variable_python_env_diffs.py | 7 +- .../v0090_add_forward_only_column.py | 2 +- ...add_dev_version_and_fingerprint_columns.py | 2 +- .../v0098_add_dbt_node_info_in_node.py | 4 +- .../v0102_normalize_python_env_payloads.py | 2 +- sqlmesh/utils/__init__.py | 20 +- sqlmesh/utils/cache.py | 12 +- sqlmesh/utils/concurrency.py | 8 +- sqlmesh/utils/config.py | 1 - sqlmesh/utils/connection_pool.py | 8 +- sqlmesh/utils/cron.py | 15 +- sqlmesh/utils/dag.py | 13 +- sqlmesh/utils/date.py | 31 +- sqlmesh/utils/errors.py | 15 +- sqlmesh/utils/git.py | 17 +- sqlmesh/utils/jinja.py | 87 +- sqlmesh/utils/lineage.py | 32 +- sqlmesh/utils/metaprogramming.py | 46 +- sqlmesh/utils/pandas.py | 6 +- sqlmesh/utils/process.py | 9 +- sqlmesh/utils/pydantic.py | 40 +- sqlmesh/utils/rich.py | 15 +- sqlmesh_dbt/cli.py | 30 +- sqlmesh_dbt/console.py | 4 +- sqlmesh_dbt/error.py | 8 +- sqlmesh_dbt/operations.py | 48 +- sqlmesh_dbt/options.py | 7 +- sqlmesh_dbt/selectors.py | 14 +- tests/cli/test_cli.py | 460 ++- tests/cli/test_integration_cli.py | 29 +- tests/cli/test_project_init.py | 21 +- tests/conftest.py | 105 +- tests/core/analytics/test_collector.py | 33 +- tests/core/analytics/test_dispatcher.py | 3 +- .../engine_adapter/integration/__init__.py | 109 +- .../engine_adapter/integration/conftest.py | 48 +- .../integration/test_freshness.py | 79 +- .../integration/test_integration.py | 596 +++- .../integration/test_integration_athena.py | 103 +- .../integration/test_integration_bigquery.py | 115 +- .../test_integration_clickhouse.py | 49 +- .../integration/test_integration_duckdb.py | 22 +- .../integration/test_integration_fabric.py | 32 +- .../integration/test_integration_postgres.py | 208 +- .../integration/test_integration_redshift.py | 73 +- .../test_integration_risingwave.py | 14 +- .../integration/test_integration_snowflake.py | 105 +- .../integration/test_integration_starrocks.py | 424 ++- .../integration/test_integration_trino.py | 40 +- tests/core/engine_adapter/test_athena.py | 118 +- tests/core/engine_adapter/test_base.py | 330 +- .../core/engine_adapter/test_base_postgres.py | 8 +- tests/core/engine_adapter/test_bigquery.py | 173 +- tests/core/engine_adapter/test_clickhouse.py | 299 +- tests/core/engine_adapter/test_databricks.py | 273 +- tests/core/engine_adapter/test_duckdb.py | 13 +- tests/core/engine_adapter/test_fabric.py | 39 +- tests/core/engine_adapter/test_mixins.py | 22 +- tests/core/engine_adapter/test_mssql.py | 96 +- tests/core/engine_adapter/test_mysql.py | 36 +- tests/core/engine_adapter/test_postgres.py | 49 +- tests/core/engine_adapter/test_redshift.py | 107 +- tests/core/engine_adapter/test_risingwave.py | 3 +- tests/core/engine_adapter/test_snowflake.py | 273 +- tests/core/engine_adapter/test_spark.py | 116 +- tests/core/engine_adapter/test_starrocks.py | 221 +- tests/core/engine_adapter/test_trino.py | 134 +- tests/core/integration/test_audits.py | 16 +- .../core/integration/test_auto_restatement.py | 37 +- tests/core/integration/test_aux_commands.py | 87 +- .../core/integration/test_change_scenarios.py | 368 +- tests/core/integration/test_config.py | 130 +- tests/core/integration/test_cron.py | 46 +- tests/core/integration/test_dbt.py | 43 +- tests/core/integration/test_dev_only_vde.py | 78 +- tests/core/integration/test_forward_only.py | 343 +- tests/core/integration/test_model_kinds.py | 141 +- tests/core/integration/test_multi_repo.py | 112 +- tests/core/integration/test_plan_options.py | 93 +- tests/core/integration/test_restatement.py | 315 +- tests/core/integration/test_run.py | 65 +- tests/core/integration/utils.py | 61 +- tests/core/linter/test_builtin.py | 15 +- tests/core/linter/test_helpers.py | 8 +- tests/core/metric/test_metric.py | 37 +- tests/core/state_sync/test_export_import.py | 133 +- tests/core/state_sync/test_state_sync.py | 522 ++- tests/core/test_audit.py | 216 +- tests/core/test_config.py | 329 +- tests/core/test_connection_config.py | 270 +- tests/core/test_context.py | 868 +++-- tests/core/test_dialect.py | 394 +-- tests/core/test_environment.py | 19 +- tests/core/test_execution_tracker.py | 17 +- tests/core/test_format.py | 50 +- tests/core/test_janitor.py | 50 +- tests/core/test_lineage.py | 13 +- tests/core/test_loader.py | 7 +- tests/core/test_macros.py | 179 +- tests/core/test_model.py | 3099 ++++++++--------- tests/core/test_notification_target.py | 18 +- tests/core/test_plan.py | 345 +- tests/core/test_plan_evaluator.py | 26 +- tests/core/test_plan_stages.py | 184 +- tests/core/test_reference.py | 10 +- tests/core/test_rule.py | 3 +- tests/core/test_scheduler.py | 268 +- tests/core/test_schema_diff.py | 175 +- tests/core/test_schema_loader.py | 39 +- tests/core/test_selector_dbt.py | 27 +- tests/core/test_selector_native.py | 134 +- tests/core/test_snapshot.py | 914 +++-- tests/core/test_snapshot_evaluator.py | 921 ++--- tests/core/test_table_diff.py | 188 +- tests/core/test_test.py | 761 ++-- tests/dbt/cli/conftest.py | 5 +- tests/dbt/cli/test_global_flags.py | 30 +- tests/dbt/cli/test_list.py | 26 +- tests/dbt/cli/test_operations.py | 40 +- tests/dbt/cli/test_options.py | 8 +- tests/dbt/cli/test_run.py | 25 +- tests/dbt/cli/test_selectors.py | 33 +- tests/dbt/conftest.py | 16 +- tests/dbt/test_adapter.py | 128 +- tests/dbt/test_config.py | 156 +- tests/dbt/test_custom_materializations.py | 39 +- tests/dbt/test_docs.py | 2 +- tests/dbt/test_integration.py | 135 +- tests/dbt/test_manifest.py | 29 +- tests/dbt/test_model.py | 136 +- tests/dbt/test_test.py | 12 +- tests/dbt/test_transformation.py | 438 ++- tests/dbt/test_util.py | 1 + tests/engines/spark/test_db_api.py | 6 +- tests/fixtures/dbt/sushi_test/config.py | 4 +- tests/integrations/github/cicd/conftest.py | 28 +- tests/integrations/github/cicd/test_config.py | 28 +- .../github/cicd/test_github_commands.py | 309 +- .../github/cicd/test_github_controller.py | 122 +- .../github/cicd/test_github_event.py | 35 +- .../github/cicd/test_integration.py | 707 ++-- tests/integrations/jupyter/test_magics.py | 112 +- tests/lsp/test_code_actions.py | 28 +- tests/lsp/test_completions.py | 48 +- tests/lsp/test_diagnostics.py | 18 +- tests/lsp/test_document_highlight.py | 14 +- tests/lsp/test_hints.py | 13 +- tests/lsp/test_reference.py | 31 +- tests/lsp/test_reference_cte.py | 22 +- tests/lsp/test_reference_cte_find_all.py | 57 +- tests/lsp/test_reference_external_model.py | 6 +- tests/lsp/test_reference_macro.py | 3 +- tests/lsp/test_reference_macro_find_all.py | 106 +- tests/lsp/test_reference_macro_multi.py | 11 +- .../lsp/test_reference_model_column_prefix.py | 63 +- tests/lsp/test_reference_model_find_all.py | 181 +- tests/lsp/test_rename_cte.py | 42 +- tests/setup.py | 3 +- tests/test_forking.py | 16 +- tests/utils/pandas.py | 4 +- tests/utils/test_aws.py | 5 +- tests/utils/test_cache.py | 20 +- tests/utils/test_concurrency.py | 18 +- tests/utils/test_connection_pool.py | 8 +- tests/utils/test_conversions.py | 7 +- tests/utils/test_date.py | 60 +- tests/utils/test_filesystem.py | 4 +- tests/utils/test_git_client.py | 30 +- tests/utils/test_helpers.py | 33 +- tests/utils/test_jinja.py | 56 +- tests/utils/test_metaprogramming.py | 76 +- tests/utils/test_pydantic.py | 10 +- tests/utils/test_windows.py | 15 +- tests/utils/test_yaml.py | 6 +- tests/web/conftest.py | 9 +- tests/web/test_lineage.py | 337 +- tests/web/test_main.py | 52 +- .../tests/tcloud/mock_tcloud/__init__.py | 2 +- .../extension/tests/tcloud/mock_tcloud/cli.py | 7 +- web/client/src/workers/sqlglot/sqlglot.py | 13 +- web/server/api/endpoints/__init__.py | 16 +- web/server/api/endpoints/commands.py | 30 +- web/server/api/endpoints/environments.py | 4 +- web/server/api/endpoints/files.py | 24 +- web/server/api/endpoints/lineage.py | 13 +- web/server/api/endpoints/meta.py | 4 +- web/server/api/endpoints/models.py | 33 +- web/server/api/endpoints/plan.py | 19 +- web/server/api/endpoints/table_diff.py | 37 +- web/server/console.py | 79 +- web/server/models.py | 36 +- web/server/openapi.py | 4 +- web/server/settings.py | 4 +- web/server/utils.py | 6 +- web/server/watcher.py | 22 +- 369 files changed, 21144 insertions(+), 12597 deletions(-) diff --git a/.github/scripts/get_scm_version.py b/.github/scripts/get_scm_version.py index 79dfee9e5d..b31a572daa 100644 --- a/.github/scripts/get_scm_version.py +++ b/.github/scripts/get_scm_version.py @@ -1,4 +1,4 @@ from setuptools_scm import get_version -version = get_version(root='../../', relative_to=__file__) -print(version.split('+')[0]) +version = get_version(root="../../", relative_to=__file__) +print(version.split("+")[0]) diff --git a/benchmarks/lsp_render_model_bench.py b/benchmarks/lsp_render_model_bench.py index f41f5f2d22..b56a7bbaf8 100644 --- a/benchmarks/lsp_render_model_bench.py +++ b/benchmarks/lsp_render_model_bench.py @@ -1,15 +1,16 @@ #!/usr/bin/env python import asyncio -import pyperf -import os import logging +import os from pathlib import Path + +import pyperf from lsprotocol import types +from pygls.client import JsonRPCClient -from sqlmesh.lsp.custom import RenderModelRequest, RENDER_MODEL_FEATURE +from sqlmesh.lsp.custom import RENDER_MODEL_FEATURE, RenderModelRequest from sqlmesh.lsp.uri import URI -from pygls.client import JsonRPCClient # Suppress debug logging during benchmark logging.getLogger().setLevel(logging.WARNING) @@ -17,28 +18,28 @@ class LSPClient(JsonRPCClient): """A custom LSP client for benchmarking.""" - + def __init__(self): super().__init__() self.render_model_result = None self.initialized = asyncio.Event() - + # Register handlers for notifications we expect from the server @self.feature(types.WINDOW_SHOW_MESSAGE) def handle_show_message(_): # Silently ignore show message notifications during benchmark pass - + @self.feature(types.WINDOW_LOG_MESSAGE) def handle_log_message(_): # Silently ignore log message notifications during benchmark pass - + async def initialize_server(self): """Send initialization request to server.""" # Get the sushi example directory sushi_dir = Path(__file__).parent.parent / "examples" / "sushi" - + response = await self.protocol.send_request_async( types.INITIALIZE, types.InitializeParams( @@ -47,13 +48,12 @@ async def initialize_server(self): capabilities=types.ClientCapabilities(), workspace_folders=[ types.WorkspaceFolder( - uri=URI.from_path(sushi_dir).value, - name="sushi" + uri=URI.from_path(sushi_dir).value, name="sushi" ) - ] - ) + ], + ), ) - + # Send initialized notification self.protocol.notify(types.INITIALIZED, types.InitializedParams()) self.initialized.set() @@ -63,56 +63,53 @@ async def initialize_server(self): async def benchmark_render_model_async(client: LSPClient, model_path: Path): """Benchmark the render_model request.""" uri = URI.from_path(model_path).value - + # Send render_model request result = await client.protocol.send_request_async( - RENDER_MODEL_FEATURE, - RenderModelRequest(textDocumentUri=uri) + RENDER_MODEL_FEATURE, RenderModelRequest(textDocumentUri=uri) ) - + return result def benchmark_render_model(loops): """Synchronous wrapper for the benchmark.""" + async def run(): # Create client client = LSPClient() - + # Start the SQLMesh LSP server as a subprocess await client.start_io("python", "-m", "sqlmesh.lsp.main") - + # Initialize the server await client.initialize_server() - + # Get a model file to test with sushi_dir = Path(__file__).parent.parent / "examples" / "sushi" model_path = sushi_dir / "models" / "customers.sql" - + # Warm up await benchmark_render_model_async(client, model_path) - + # Run benchmark t0 = pyperf.perf_counter() for _ in range(loops): await benchmark_render_model_async(client, model_path) dt = pyperf.perf_counter() - t0 - + # Clean up await client.stop() - + return dt - + return asyncio.run(run()) def main(): runner = pyperf.Runner() - runner.bench_time_func( - "lsp_render_model", - benchmark_render_model - ) + runner.bench_time_func("lsp_render_model", benchmark_render_model) if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/examples/custom_materializations/custom_materializations/custom_kind.py b/examples/custom_materializations/custom_materializations/custom_kind.py index 8a0eabcfa7..b823a2b151 100644 --- a/examples/custom_materializations/custom_materializations/custom_kind.py +++ b/examples/custom_materializations/custom_materializations/custom_kind.py @@ -2,7 +2,7 @@ import typing as t -from sqlmesh import CustomMaterialization, CustomKind, Model +from sqlmesh import CustomKind, CustomMaterialization, Model from sqlmesh.utils.pydantic import validate_string if t.TYPE_CHECKING: @@ -15,7 +15,9 @@ def custom_property(self) -> str: return validate_string(self.materialization_properties.get("custom_property")) -class CustomFullWithCustomKindMaterialization(CustomMaterialization[ExtendedCustomKind]): +class CustomFullWithCustomKindMaterialization( + CustomMaterialization[ExtendedCustomKind] +): NAME = "custom_full_with_custom_kind" def insert( diff --git a/examples/sushi/config.py b/examples/sushi/config.py index b985e24ec5..5f4c74ea4a 100644 --- a/examples/sushi/config.py +++ b/examples/sushi/config.py @@ -1,23 +1,16 @@ import os -from sqlmesh.core.config.common import VirtualEnvironmentMode, TableNamingConvention -from sqlmesh.core.config import ( - AutoCategorizationMode, - BigQueryConnectionConfig, - CategorizerConfig, - Config, - DuckDBConnectionConfig, - EnvironmentSuffixTarget, - GatewayConfig, - ModelDefaultsConfig, - PlanConfig, -) +from sqlmesh.core.config import (AutoCategorizationMode, + BigQueryConnectionConfig, CategorizerConfig, + Config, DuckDBConnectionConfig, + EnvironmentSuffixTarget, GatewayConfig, + ModelDefaultsConfig, PlanConfig) +from sqlmesh.core.config.common import (TableNamingConvention, + VirtualEnvironmentMode) from sqlmesh.core.config.linter import LinterConfig -from sqlmesh.core.notification_target import ( - BasicSMTPNotificationTarget, - SlackApiNotificationTarget, - SlackWebhookNotificationTarget, -) +from sqlmesh.core.notification_target import (BasicSMTPNotificationTarget, + SlackApiNotificationTarget, + SlackWebhookNotificationTarget) from sqlmesh.core.user import User, UserRole CURRENT_FILE_PATH = os.path.abspath(__file__) @@ -64,7 +57,9 @@ gateways={ "bq": GatewayConfig( connection=BigQueryConnectionConfig(), - state_connection=DuckDBConnectionConfig(database=f"{DATA_DIR}/bigquery.duckdb"), + state_connection=DuckDBConnectionConfig( + database=f"{DATA_DIR}/bigquery.duckdb" + ), ) }, default_gateway="bq", @@ -123,7 +118,12 @@ roles=[UserRole.REQUIRED_APPROVER], notification_targets=[ SlackApiNotificationTarget( - notify_on=["apply_start", "apply_failure", "apply_end", "audit_failure"], + notify_on=[ + "apply_start", + "apply_failure", + "apply_end", + "audit_failure", + ], token=os.getenv("ADMIN_SLACK_API_TOKEN"), channel="UXXXXXXXXX", # User's Slack member ID ), diff --git a/examples/sushi/macros/macros.py b/examples/sushi/macros/macros.py index 763f6d62b2..4c9f7f4838 100644 --- a/examples/sushi/macros/macros.py +++ b/examples/sushi/macros/macros.py @@ -6,7 +6,9 @@ @macro() def incremental_by_ds(evaluator, column: exp.Column): - return between(evaluator, column, evaluator.locals["start_date"], evaluator.locals["end_date"]) + return between( + evaluator, column, evaluator.locals["start_date"], evaluator.locals["end_date"] + ) @macro() @@ -14,7 +16,9 @@ def assert_has_columns(evaluator, model, columns_to_types): if evaluator.runtime_stage == "creating": expected_schema = { column_type.name: exp.maybe_parse( - column_type.text("expression"), into=exp.DataType, dialect=evaluator.dialect + column_type.text("expression"), + into=exp.DataType, + dialect=evaluator.dialect, ) for column_type in columns_to_types.expressions } diff --git a/examples/sushi/models/disabled.py b/examples/sushi/models/disabled.py index 1ed1058d1c..c424cb8bd5 100644 --- a/examples/sushi/models/disabled.py +++ b/examples/sushi/models/disabled.py @@ -1,4 +1,5 @@ import typing as t + from sqlmesh import ExecutionContext, model diff --git a/examples/sushi/models/items.py b/examples/sushi/models/items.py index 54c9442dc5..f17767a7fd 100644 --- a/examples/sushi/models/items.py +++ b/examples/sushi/models/items.py @@ -50,7 +50,9 @@ @model( "sushi.items", kind=dict( - name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, time_column="event_date", batch_size=30 + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + time_column="event_date", + batch_size=30, ), start="1 week ago", cron="@daily", diff --git a/examples/sushi/models/order_items.py b/examples/sushi/models/order_items.py index 9d4dc551e3..43211bebb5 100644 --- a/examples/sushi/models/order_items.py +++ b/examples/sushi/models/order_items.py @@ -37,7 +37,11 @@ def get_items_table(context: ExecutionContext) -> str: audits=[ ( "NOT_NULL", - {"columns": [to_column(c) for c in ("id", "order_id", "item_id", "quantity")]}, + { + "columns": [ + to_column(c) for c in ("id", "order_id", "item_id", "quantity") + ] + }, ), ("assert_order_items_quantity_exceeds_threshold", {"quantity": 0}), ], diff --git a/examples/sushi/models/orders.py b/examples/sushi/models/orders.py index 8d8718a3e3..4ddb0e06b1 100644 --- a/examples/sushi/models/orders.py +++ b/examples/sushi/models/orders.py @@ -17,7 +17,9 @@ "sushi.orders", description="Table of sushi orders.", kind=dict( - name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, time_column="event_date", batch_size=30 + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + time_column="event_date", + batch_size=30, ), start="1 week ago", cron="@daily", diff --git a/examples/sushi/models/raw_marketing.py b/examples/sushi/models/raw_marketing.py index b17c471895..c3d0e25508 100644 --- a/examples/sushi/models/raw_marketing.py +++ b/examples/sushi/models/raw_marketing.py @@ -55,14 +55,18 @@ def execute( df_new = pd.DataFrame( { "customer_id": random.sample(range(0, 100), k=num_customers), - "status": np.random.choice(["active", "inactive"], size=num_customers, p=[0.8, 0.2]), + "status": np.random.choice( + ["active", "inactive"], size=num_customers, p=[0.8, 0.2] + ), "updated_at": [exec_time] * num_customers, } ) # clickhouse returns a dataframe with no columns if the query is empty, so we can't merge if not df_existing.empty: - df = df_new.merge(df_existing, on="customer_id", how="left", suffixes=(None, "_old")) + df = df_new.merge( + df_existing, on="customer_id", how="left", suffixes=(None, "_old") + ) else: df = df_new df["status_old"] = pd.NA diff --git a/examples/sushi/models/waiters.py b/examples/sushi/models/waiters.py index f9e26eda82..3b9a20a5a3 100644 --- a/examples/sushi/models/waiters.py +++ b/examples/sushi/models/waiters.py @@ -62,7 +62,9 @@ def entrypoint(evaluator: MacroEvaluator) -> exp.Select: name = ".".join([f'"{default_catalog}"', name]) - assert parent_snapshot_name == name, f"Snapshot Name: {parent_snapshot_name}, Name: {name}" + assert ( + parent_snapshot_name == name + ), f"Snapshot Name: {parent_snapshot_name}, Name: {name}" excluded = {"id", "customer_id", "start_ts", "end_ts"} projections = [] diff --git a/examples/sushi/signals/__init__.py b/examples/sushi/signals/__init__.py index bd7c839fce..5e2f74379e 100644 --- a/examples/sushi/signals/__init__.py +++ b/examples/sushi/signals/__init__.py @@ -1,6 +1,6 @@ import typing as t -from sqlmesh import signal, DatetimeRanges +from sqlmesh import DatetimeRanges, signal @signal() diff --git a/examples/sushi_dlt/sushi_pipeline.py b/examples/sushi_dlt/sushi_pipeline.py index 3a44a4897e..abc85090e6 100644 --- a/examples/sushi_dlt/sushi_pipeline.py +++ b/examples/sushi_dlt/sushi_pipeline.py @@ -1,4 +1,5 @@ import typing as t + import dlt @@ -76,7 +77,12 @@ def sushi_menu() -> t.Iterator[t.Dict[str, t.Any]]: { "id": 3, "name": "Temaki", - "fillings": ["Tuna Temaki", "Salmon Temaki", "Vegetable Temaki", "Ebi Temaki"], + "fillings": [ + "Tuna Temaki", + "Salmon Temaki", + "Vegetable Temaki", + "Ebi Temaki", + ], "details": { "preparation": "Hand Roll", "ingredients": ["Seaweed", "Rice", "Fish", "Vegetables"], diff --git a/examples/wursthall/models/db/order_f.py b/examples/wursthall/models/db/order_f.py index d682d55f02..9e00706bd6 100644 --- a/examples/wursthall/models/db/order_f.py +++ b/examples/wursthall/models/db/order_f.py @@ -53,8 +53,7 @@ def execute( ) df_order_item_f = context.fetchdf( - parse_one( - f""" + parse_one(f""" SELECT order_id, customer_id, @@ -64,14 +63,15 @@ def execute( FROM {order_item_f_table_name} WHERE order_ds BETWEEN '{to_ds(start)}' AND '{to_ds(end)}' - """ - ), + """), quote_identifiers=True, ) df_order_item_f = df_order_item_f.merge(df_item_d, how="inner", on="item_id") df_order_item_f["item_price"] = 1.00 - df_order_item_f["item_total"] = df_order_item_f["item_price"] * df_order_item_f["quantity"] + df_order_item_f["item_total"] = ( + df_order_item_f["item_price"] * df_order_item_f["quantity"] + ) df_order_item_f = ( df_order_item_f.groupby(["order_id", "customer_id", "order_ds"], dropna=False) .agg(order_total=("item_total", "sum")) diff --git a/examples/wursthall/models/src/customer_details.py b/examples/wursthall/models/src/customer_details.py index 44c56a60e5..ad29493c9c 100644 --- a/examples/wursthall/models/src/customer_details.py +++ b/examples/wursthall/models/src/customer_details.py @@ -5,7 +5,8 @@ import pandas as pd # noqa: TID253 from faker import Faker -from models.src.shared import DATA_START_DATE_STR, iter_dates, set_seed # type: ignore +from models.src.shared import (DATA_START_DATE_STR, iter_dates, # type: ignore + set_seed) from sqlmesh import model from sqlmesh.core.model import IncrementalByTimeRangeKind, TimeColumn diff --git a/examples/wursthall/models/src/order_item_details.py b/examples/wursthall/models/src/order_item_details.py index 852250d2e2..c4a53c6010 100644 --- a/examples/wursthall/models/src/order_item_details.py +++ b/examples/wursthall/models/src/order_item_details.py @@ -6,7 +6,8 @@ import numpy as np # noqa: TID253 import pandas as pd # noqa: TID253 from faker import Faker -from models.src.shared import DATA_START_DATE_STR, iter_dates, set_seed # type: ignore +from models.src.shared import (DATA_START_DATE_STR, iter_dates, # type: ignore + set_seed) from sqlglot import parse_one from sqlmesh import ExecutionContext, model @@ -54,16 +55,14 @@ def execute( # project and we want to ensure that the resulting query is properly quoted in # the target dialect before executing it df_customers = context.fetchdf( - parse_one( - f""" + parse_one(f""" SELECT id AS customer_id, register_ds FROM {customer_details_table_name} WHERE register_ds <= '{to_ds(end)}' - """ - ), + """), quote_identifiers=True, ) @@ -99,7 +98,9 @@ def execute( ) ): item_id = str( - df_menu_items.iloc[[random.choice(range(num_menu_items))]]["item_id"].values[0] + df_menu_items.iloc[[random.choice(range(num_menu_items))]][ + "item_id" + ].values[0] ) quantity = np.random.choice(range(1, 4), p=[0.8, 0.1, 0.1]) results.append( diff --git a/pdoc/cli.py b/pdoc/cli.py index 9301ae0444..5baf4dcb53 100755 --- a/pdoc/cli.py +++ b/pdoc/cli.py @@ -4,10 +4,9 @@ from pathlib import Path from unittest import mock -from pdoc.__main__ import cli, parser - # Need this import or else import_module doesn't work import sqlmesh +from pdoc.__main__ import cli, parser def mocked_import(*args, **kwargs): @@ -29,7 +28,9 @@ def mocked_import(*args, **kwargs): opts.logo_link = "https://tobikodata.com" opts.footer_text = "Copyright Tobiko Data Inc. 2022" opts.template_directory = Path(__file__).parent.joinpath("templates").absolute() - opts.edit_url = ["sqlmesh=https://github.com/SQLMesh/sqlmesh/tree/main/sqlmesh/"] + opts.edit_url = [ + "sqlmesh=https://github.com/SQLMesh/sqlmesh/tree/main/sqlmesh/" + ] with mock.patch("pdoc.__main__.parser", **{"parse_args.return_value": opts}): cli() diff --git a/sqlmesh/__init__.py b/sqlmesh/__init__.py index 577a3aaf02..6f65a7df73 100644 --- a/sqlmesh/__init__.py +++ b/sqlmesh/__init__.py @@ -20,25 +20,26 @@ from sqlmesh.core import constants as c from sqlmesh.core.config import Config as Config -from sqlmesh.core.context import Context as Context, ExecutionContext as ExecutionContext +from sqlmesh.core.context import Context as Context +from sqlmesh.core.context import ExecutionContext as ExecutionContext from sqlmesh.core.engine_adapter import EngineAdapter as EngineAdapter -from sqlmesh.core.macros import SQL as SQL, macro as macro -from sqlmesh.core.model import Model as Model, model as model +from sqlmesh.core.macros import SQL as SQL +from sqlmesh.core.macros import macro as macro +from sqlmesh.core.model import Model as Model +from sqlmesh.core.model import model as model +from sqlmesh.core.model.kind import CustomKind as CustomKind from sqlmesh.core.signal import signal as signal from sqlmesh.core.snapshot import Snapshot as Snapshot -from sqlmesh.core.snapshot.evaluator import ( - CustomMaterialization as CustomMaterialization, -) -from sqlmesh.core.model.kind import CustomKind as CustomKind -from sqlmesh.utils import ( - debug_mode_enabled as debug_mode_enabled, - enable_debug_mode as enable_debug_mode, - str_to_bool, -) +from sqlmesh.core.snapshot.evaluator import \ + CustomMaterialization as CustomMaterialization +from sqlmesh.utils import debug_mode_enabled as debug_mode_enabled +from sqlmesh.utils import enable_debug_mode as enable_debug_mode +from sqlmesh.utils import str_to_bool from sqlmesh.utils.date import DatetimeRanges as DatetimeRanges try: - from sqlmesh._version import __version__ as __version__, __version_tuple__ as __version_tuple__ + from sqlmesh._version import __version__ as __version__ + from sqlmesh._version import __version_tuple__ as __version_tuple__ except ImportError: pass @@ -178,7 +179,9 @@ def remove_excess_logs( log_file_dir = log_file_dir or c.DEFAULT_LOG_FILE_DIR log_path_prefix = Path(log_file_dir) / LOG_FILENAME_PREFIX - for path in list(sorted(glob.glob(f"{log_path_prefix}*.log"), reverse=True))[log_limit:]: + for path in list(sorted(glob.glob(f"{log_path_prefix}*.log"), reverse=True))[ + log_limit: + ]: os.remove(path) @@ -222,7 +225,9 @@ def configure_logging( if write_to_file: os.makedirs(str(log_file_dir), exist_ok=True) - filename = f"{log_path_prefix}{datetime.now().strftime('%Y_%m_%d_%H_%M_%S')}.log" + filename = ( + f"{log_path_prefix}{datetime.now().strftime('%Y_%m_%d_%H_%M_%S')}.log" + ) file_handler = logging.FileHandler(filename, mode="w", encoding="utf-8") # the log files should always log at least info so that users will always have diff --git a/sqlmesh/cli/__init__.py b/sqlmesh/cli/__init__.py index 3b417eb478..034064acb8 100644 --- a/sqlmesh/cli/__init__.py +++ b/sqlmesh/cli/__init__.py @@ -4,6 +4,7 @@ import click from sqlglot.errors import SqlglotError + from sqlmesh.core.context import Context from sqlmesh.utils import debug_mode_enabled from sqlmesh.utils.errors import SQLMeshError @@ -21,11 +22,17 @@ def error_handler( def wrapper(*args: t.List[t.Any], **kwargs: t.Any) -> DECORATOR_RETURN_TYPE: context_or_obj = args[0] sqlmesh_context = ( - context_or_obj.obj if isinstance(context_or_obj, click.Context) else context_or_obj + context_or_obj.obj + if isinstance(context_or_obj, click.Context) + else context_or_obj ) if not isinstance(sqlmesh_context, Context): sqlmesh_context = None - handler = _debug_exception_handler if debug_mode_enabled() else _default_exception_handler + handler = ( + _debug_exception_handler + if debug_mode_enabled() + else _default_exception_handler + ) return handler(sqlmesh_context, lambda: func(*args, **kwargs)) return wrapper diff --git a/sqlmesh/cli/main.py b/sqlmesh/cli/main.py index 5574a892cb..42f365a689 100644 --- a/sqlmesh/cli/main.py +++ b/sqlmesh/cli/main.py @@ -11,12 +11,8 @@ from sqlmesh import configure_logging, remove_excess_logs from sqlmesh.cli import error_handler from sqlmesh.cli import options as opt -from sqlmesh.cli.project_init import ( - InitCliMode, - ProjectTemplate, - init_example_project, - interactive_init, -) +from sqlmesh.cli.project_init import (InitCliMode, ProjectTemplate, + init_example_project, interactive_init) from sqlmesh.core.analytics import cli_analytics from sqlmesh.core.config import load_configs from sqlmesh.core.console import configure_console, get_console @@ -201,7 +197,9 @@ def init( try: project_template = ProjectTemplate(template.lower()) except ValueError: - template_strings = "', '".join([template.value for template in ProjectTemplate]) + template_strings = "', '".join( + [template.value for template in ProjectTemplate] + ) raise click.ClickException( f"Invalid project template '{template}'. Please specify one of '{template_strings}'." ) @@ -220,7 +218,9 @@ def init( console = srich.console - project_template, engine_type, cli_mode = interactive_init(ctx.obj, console, project_template) + project_template, engine_type, cli_mode = interactive_init( + ctx.obj, console, project_template + ) config_path = init_example_project( path=ctx.obj, @@ -234,7 +234,9 @@ def init( engine_install_text = "" if engine_type and engine_type not in ("duckdb", "motherduck"): install_text = ( - "pyspark" if engine_type == "spark" else f"sqlmesh\\[{engine_type.replace('_', '')}]" + "pyspark" + if engine_type == "spark" + else f"sqlmesh\\[{engine_type.replace('_', '')}]" ) engine_install_text = f'• Run command in CLI to install your SQL engine\'s Python dependencies: pip install "{install_text}"\n' # interactive init does not support DLT template @@ -279,7 +281,9 @@ def init( type=str, help="The SQL dialect to render the query as.", ) -@click.option("--no-format", is_flag=True, help="Disable fancy formatting of the query.") +@click.option( + "--no-format", is_flag=True, help="Disable fancy formatting of the query." +) @opt.format_options @click.pass_context @error_handler @@ -392,7 +396,9 @@ def format( ctx: click.Context, paths: t.Optional[t.Tuple[str, ...]] = None, **kwargs: t.Any ) -> None: """Format all SQL models and audits.""" - if not ctx.obj.format(**{k: v for k, v in kwargs.items() if v is not None}, paths=paths): + if not ctx.obj.format( + **{k: v for k, v in kwargs.items() if v is not None}, paths=paths + ): ctx.exit(1) @@ -618,7 +624,9 @@ def plan( @click.pass_context @error_handler @cli_analytics -def run(ctx: click.Context, environment: t.Optional[str] = None, **kwargs: t.Any) -> None: +def run( + ctx: click.Context, environment: t.Optional[str] = None, **kwargs: t.Any +) -> None: """Evaluate missing intervals for the target environment.""" context = ctx.obj select_models = kwargs.pop("select_model") or None @@ -677,7 +685,9 @@ def janitor( The janitor cleans up old environments and expired snapshots. """ - ctx.obj.run_janitor(ignore_ttl, force_delete=force_delete, environment=environment, **kwargs) + ctx.obj.run_janitor( + ignore_ttl, force_delete=force_delete, environment=environment, **kwargs + ) @cli.command("destroy") @@ -838,7 +848,9 @@ def audit( execution_time: t.Optional[TimeLike] = None, ) -> None: """Run audits for the target model(s).""" - if not obj.audit(models=models, start=start, end=end, execution_time=execution_time): + if not obj.audit( + models=models, start=start, end=end, execution_time=execution_time + ): exit(1) @@ -1112,7 +1124,9 @@ def rewrite(obj: Context, sql: str, read: str = "", write: str = "") -> None: https://sqlmesh.readthedocs.io/en/latest/concepts/metrics/overview/ """ obj.console.show_sql( - obj.rewrite(sql, dialect=read).sql(pretty=True, dialect=write or obj.config.dialect), + obj.rewrite(sql, dialect=read).sql( + pretty=True, dialect=write or obj.config.dialect + ), ) @@ -1185,10 +1199,14 @@ def dlt_refresh( """Attaches to a DLT pipeline with the option to update specific or all missing tables in the SQLMesh project.""" from sqlmesh.integrations.dlt import generate_dlt_models - sqlmesh_models = generate_dlt_models(ctx.obj, pipeline, list(table or []), force, dlt_path) + sqlmesh_models = generate_dlt_models( + ctx.obj, pipeline, list(table or []), force, dlt_path + ) if sqlmesh_models: model_names = "\n".join([f"- {model_name}" for model_name in sqlmesh_models]) - ctx.obj.console.log_success(f"Updated SQLMesh project with models:\n{model_names}") + ctx.obj.console.log_success( + f"Updated SQLMesh project with models:\n{model_names}" + ) else: ctx.obj.console.log_success("All SQLMesh models are up to date.") @@ -1301,7 +1319,9 @@ def state_export( @click.pass_obj @error_handler @cli_analytics -def state_import(obj: Context, input_file: Path, replace: bool, no_confirm: bool) -> None: +def state_import( + obj: Context, input_file: Path, replace: bool, no_confirm: bool +) -> None: """Import a state export file back into the state database""" confirm = not no_confirm obj.import_state(input_file=input_file, clear=replace, confirm=confirm) diff --git a/sqlmesh/cli/project_init.py b/sqlmesh/cli/project_init.py index e3132a6de3..fe5d5057db 100644 --- a/sqlmesh/cli/project_init.py +++ b/sqlmesh/cli/project_init.py @@ -1,21 +1,19 @@ import typing as t +from dataclasses import dataclass from enum import Enum from pathlib import Path -from dataclasses import dataclass -from rich.prompt import Prompt + from rich.console import Console +from rich.prompt import Prompt + +from sqlmesh.core.config.common import (DBT_PROJECT_FILENAME, + VirtualEnvironmentMode) +from sqlmesh.core.config.connection import (CONNECTION_CONFIG_TO_TYPE, + DIALECT_TO_TYPE, + INIT_DISPLAY_INFO_TO_TYPE) from sqlmesh.integrations.dlt import generate_dlt_models_and_settings from sqlmesh.utils.date import yesterday_ds from sqlmesh.utils.errors import SQLMeshError -from sqlmesh.core.config.common import VirtualEnvironmentMode - -from sqlmesh.core.config.common import DBT_PROJECT_FILENAME -from sqlmesh.core.config.connection import ( - CONNECTION_CONFIG_TO_TYPE, - DIALECT_TO_TYPE, - INIT_DISPLAY_INFO_TO_TYPE, -) - PRIMITIVES = (str, int, bool, float) @@ -42,21 +40,22 @@ def _gen_config( ) -> str: project_dialect = dialect or DIALECT_TO_TYPE.get(engine_type) - connection_settings = ( - settings - or """ type: duckdb + connection_settings = settings or """ type: duckdb database: db.db""" - ) if not settings and template != ProjectTemplate.DBT: - doc_link = "https://sqlmesh.readthedocs.io/en/stable/integrations/engines{engine_link}" + doc_link = ( + "https://sqlmesh.readthedocs.io/en/stable/integrations/engines{engine_link}" + ) engine_link = "" if engine_type in CONNECTION_CONFIG_TO_TYPE: required_fields = [] non_required_fields = [] - for name, field in CONNECTION_CONFIG_TO_TYPE[engine_type].model_fields.items(): + for name, field in CONNECTION_CONFIG_TO_TYPE[ + engine_type + ].model_fields.items(): field_name = field.alias or name default_value = field.get_default() @@ -169,7 +168,9 @@ def _gen_config( # sqlmesh plan prod # Specify `prod` to apply changes to production """ - return default_configs[template] + (flow_cli_mode if cli_mode == InitCliMode.FLOW else "") + return default_configs[template] + ( + flow_cli_mode if cli_mode == InitCliMode.FLOW else "" + ) @dataclass @@ -327,7 +328,10 @@ def init_example_project( f"Found an existing config file '{config_path}'.\n\nPlease change to another directory or remove the existing file." ) - if template == ProjectTemplate.DBT and not Path(root_path, DBT_PROJECT_FILENAME).exists(): + if ( + template == ProjectTemplate.DBT + and not Path(root_path, DBT_PROJECT_FILENAME).exists() + ): raise SQLMeshError( "Required dbt project file 'dbt_project.yml' not found in the current directory.\n\nPlease add it or change directories before running `sqlmesh init` to set up your project." ) @@ -356,7 +360,9 @@ def init_example_project( "Please provide a DLT pipeline with the `--dlt-pipeline` flag to generate a SQLMesh project from DLT." ) - _create_config(config_path, engine_type, dialect, settings, start, template, cli_mode) + _create_config( + config_path, engine_type, dialect, settings, start, template, cli_mode + ) if template == ProjectTemplate.DBT: return config_path @@ -364,7 +370,9 @@ def init_example_project( if template == ProjectTemplate.DLT: _create_object_files( - models_path, {model[0].split(".")[-1]: model[1] for model in dlt_models}, "sql" + models_path, + {model[0].split(".")[-1]: model[1] for model in dlt_models}, + "sql", ) return config_path @@ -397,7 +405,9 @@ def _create_config( template: ProjectTemplate, cli_mode: InitCliMode, ) -> None: - project_config = _gen_config(engine_type, settings, start, template, cli_mode, dialect) + project_config = _gen_config( + engine_type, settings, start, template, cli_mode, dialect + ) _write_file( config_path, @@ -405,7 +415,9 @@ def _create_config( ) -def _create_object_files(path: Path, object_dict: t.Dict[str, str], file_extension: str) -> None: +def _create_object_files( + path: Path, object_dict: t.Dict[str, str], file_extension: str +) -> None: for object_name, object_def in object_dict.items(): # file name is table component of catalog.schema.table _write_file(path / f"{object_name.split('.')[-1]}.{file_extension}", object_def) @@ -424,7 +436,9 @@ def interactive_init( console.print("──────────────────────────────") console.print("Welcome to SQLMesh!") - project_template = _init_template_prompt(console) if not project_template else project_template + project_template = ( + _init_template_prompt(console) if not project_template else project_template + ) if project_template == ProjectTemplate.DBT: return (project_template, None, None) @@ -436,7 +450,10 @@ def interactive_init( def _init_integer_prompt( - console: Console, err_msg_entity: str, num_options: int, retry_func: t.Callable[[t.Any], t.Any] + console: Console, + err_msg_entity: str, + num_options: int, + retry_func: t.Callable[[t.Any], t.Any], ) -> int: err_msg = "\nERROR: '{option_str}' is not a valid {err_msg_entity} number - please enter a number between 1 and {num_options} or exit with control+c\n" while True: @@ -451,7 +468,9 @@ def _init_integer_prompt( if value_error or option_num < 1 or option_num > num_options: console.print( err_msg.format( - option_str=option_str, err_msg_entity=err_msg_entity, num_options=num_options + option_str=option_str, + err_msg_entity=err_msg_entity, + num_options=num_options, ), style="red", ) @@ -460,10 +479,14 @@ def _init_integer_prompt( return option_num -def _init_display_choices(values_dict: t.Dict[str, str], console: Console) -> t.Dict[int, str]: +def _init_display_choices( + values_dict: t.Dict[str, str], console: Console +) -> t.Dict[int, str]: display_num_to_value = {} for i, value_str in enumerate(values_dict.keys()): - console.print(f" \\[{i + 1}] {' ' if i < 9 else ''}{value_str} {values_dict[value_str]}") + console.print( + f" \\[{i + 1}] {' ' if i < 9 else ''}{value_str} {values_dict[value_str]}" + ) display_num_to_value[i + 1] = value_str console.print("") return display_num_to_value @@ -496,9 +519,12 @@ def _init_engine_prompt(console: Console) -> str: # INIT_DISPLAY_INFO_TO_TYPE is a dict of {engine_type: (display_order, display_name)} DISPLAY_NAME_TO_TYPE = {v[1]: k for k, v in INIT_DISPLAY_INFO_TO_TYPE.items()} ordered_engine_display_names = { - info[1]: "" for info in sorted(INIT_DISPLAY_INFO_TO_TYPE.values(), key=lambda x: x[0]) + info[1]: "" + for info in sorted(INIT_DISPLAY_INFO_TO_TYPE.values(), key=lambda x: x[0]) } - display_num_to_display_name = _init_display_choices(ordered_engine_display_names, console) + display_num_to_display_name = _init_display_choices( + ordered_engine_display_names, console + ) engine_num = _init_integer_prompt( console, "engine", len(ordered_engine_display_names), _init_engine_prompt diff --git a/sqlmesh/core/_typing.py b/sqlmesh/core/_typing.py index 2bc69e901b..dd59c9a376 100644 --- a/sqlmesh/core/_typing.py +++ b/sqlmesh/core/_typing.py @@ -9,7 +9,9 @@ TableName = t.Union[str, exp.Table] SchemaName = t.Union[str, exp.Table] SessionProperties = t.Dict[str, t.Union[exp.Expr, str, int, float, bool]] - CustomMaterializationProperties = t.Dict[str, t.Union[exp.Expr, str, int, float, bool]] + CustomMaterializationProperties = t.Dict[ + str, t.Union[exp.Expr, str, int, float, bool] + ] if sys.version_info >= (3, 11): diff --git a/sqlmesh/core/analytics/__init__.py b/sqlmesh/core/analytics/__init__.py index fcf6d52064..8c08afbd2c 100644 --- a/sqlmesh/core/analytics/__init__.py +++ b/sqlmesh/core/analytics/__init__.py @@ -6,7 +6,8 @@ from functools import wraps from sqlmesh.core.analytics.collector import AnalyticsCollector -from sqlmesh.core.analytics.dispatcher import AsyncEventDispatcher, NoopEventDispatcher +from sqlmesh.core.analytics.dispatcher import (AsyncEventDispatcher, + NoopEventDispatcher) from sqlmesh.utils import str_to_bool if t.TYPE_CHECKING: @@ -105,7 +106,9 @@ def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _T: break if should_log: - collector.on_python_api_command(command_name=func.__name__, command_args=kwargs) + collector.on_python_api_command( + command_name=func.__name__, command_args=kwargs + ) return func(*args, **kwargs) diff --git a/sqlmesh/core/analytics/collector.py b/sqlmesh/core/analytics/collector.py index cfdb60aadd..ded217fc95 100644 --- a/sqlmesh/core/analytics/collector.py +++ b/sqlmesh/core/analytics/collector.py @@ -7,7 +7,8 @@ from pathlib import Path from sqlmesh.core import constants as c -from sqlmesh.core.analytics.dispatcher import AsyncEventDispatcher, EventDispatcher +from sqlmesh.core.analytics.dispatcher import (AsyncEventDispatcher, + EventDispatcher) from sqlmesh.utils import random_id from sqlmesh.utils.date import now_timestamp from sqlmesh.utils.hashing import md5 @@ -49,7 +50,9 @@ def on_cicd_command( cicd_config: The CICD bot configuration. """ additional_args = {} - if cicd_bot_config is not None and getattr(cicd_bot_config, "FIELDS_FOR_ANALYTICS", None): + if cicd_bot_config is not None and getattr( + cicd_bot_config, "FIELDS_FOR_ANALYTICS", None + ): additional_args["cicd_bot_config"] = cicd_bot_config.dict( include=cicd_bot_config.FIELDS_FOR_ANALYTICS, mode="json" ) @@ -76,10 +79,15 @@ def on_cli_command( parent_command_names: The names of the parent commands. """ self._on_command( - "CLI_COMMAND", command_name, command_args, parent_command_names=parent_command_names + "CLI_COMMAND", + command_name, + command_args, + parent_command_names=parent_command_names, ) - def on_magic_command(self, *, command_name: str, command_args: t.Collection[str]) -> None: + def on_magic_command( + self, *, command_name: str, command_args: t.Collection[str] + ) -> None: """Called when a Notebook magic command is executed. Args: @@ -88,7 +96,9 @@ def on_magic_command(self, *, command_name: str, command_args: t.Collection[str] """ self._on_command("MAGIC_COMMAND", command_name, command_args) - def on_python_api_command(self, *, command_name: str, command_args: t.Collection[str]) -> None: + def on_python_api_command( + self, *, command_name: str, command_args: t.Collection[str] + ) -> None: """Called when a Python method is called directly. Args: @@ -164,7 +174,9 @@ def on_plan_apply_start( { "plan_id": plan.plan_id, "engine_type": engine_type.lower() if engine_type is not None else None, - "state_sync_type": state_sync_type.lower() if state_sync_type is not None else None, + "state_sync_type": ( + state_sync_type.lower() if state_sync_type is not None else None + ), "scheduler_type": scheduler_type.lower(), "is_dev": plan.is_dev, "skip_backfill": plan.skip_backfill, @@ -184,7 +196,9 @@ def on_plan_apply_start( }, ) - def on_plan_apply_end(self, *, plan_id: str, error: t.Optional[t.Any] = None) -> None: + def on_plan_apply_end( + self, *, plan_id: str, error: t.Optional[t.Any] = None + ) -> None: """Called after the plan application ends. Args: @@ -200,7 +214,9 @@ def on_plan_apply_end(self, *, plan_id: str, error: t.Optional[t.Any] = None) -> }, ) - def on_snapshots_created(self, *, new_snapshots: t.Collection[Snapshot], plan_id: str) -> None: + def on_snapshots_created( + self, *, new_snapshots: t.Collection[Snapshot], plan_id: str + ) -> None: """Called after new snapshots were created and stored in the SQLMesh state. Args: @@ -217,19 +233,27 @@ def on_snapshots_created(self, *, new_snapshots: t.Collection[Snapshot], plan_id "identifier": snapshot.identifier, "version": snapshot.version, "node_type": snapshot.node_type.lower(), - "model_kind": snapshot.model.kind.name.value.lower() - if snapshot.is_model - else None, + "model_kind": ( + snapshot.model.kind.name.value.lower() + if snapshot.is_model + else None + ), "is_sql": snapshot.model.is_sql if snapshot.is_model else None, "change_category": ( - snapshot.change_category.name.lower() if snapshot.change_category else None + snapshot.change_category.name.lower() + if snapshot.change_category + else None ), "dialect": getattr(snapshot.node, "dialect", None), - "audits_count": len(snapshot.model.audits) if snapshot.is_model else None, + "audits_count": ( + len(snapshot.model.audits) if snapshot.is_model else None + ), "effective_from_set": snapshot.effective_from is not None, } ) - self._add_event("SNAPSHOTS_CREATED", {"plan_id": plan_id, "snapshots": snapshots}) + self._add_event( + "SNAPSHOTS_CREATED", {"plan_id": plan_id, "snapshots": snapshots} + ) def on_run_start(self, *, engine_type: str, state_sync_type: str) -> str: """Called after a run starts. @@ -253,7 +277,12 @@ def on_run_start(self, *, engine_type: str, state_sync_type: str) -> str: return run_id def on_run_end( - self, *, run_id: str, succeeded: bool, interrupted: bool, error: t.Optional[t.Any] = None + self, + *, + run_id: str, + succeeded: bool, + interrupted: bool, + error: t.Optional[t.Any] = None, ) -> None: """Called after a run ends. diff --git a/sqlmesh/core/analytics/dispatcher.py b/sqlmesh/core/analytics/dispatcher.py index 4ebb98b391..93404dd70c 100644 --- a/sqlmesh/core/analytics/dispatcher.py +++ b/sqlmesh/core/analytics/dispatcher.py @@ -45,10 +45,14 @@ def __init__( ) def emit(self, events: t.List[t.Dict[str, t.Any]]) -> None: - data = json.dumps({"events": events, "versions": self._versions}).encode("utf-8") + data = json.dumps({"events": events, "versions": self._versions}).encode( + "utf-8" + ) data = gzip.compress(data) response = self._session.post( - self.sqlmesh_url, data=data, timeout=(self.connect_timeout, self.read_timeout) + self.sqlmesh_url, + data=data, + timeout=(self.connect_timeout, self.read_timeout), ) raise_for_status(response) @@ -113,7 +117,9 @@ def __init__( self._events_lock = Lock() self._shutdown_event = Event() - self._emitter_thread = Thread(target=self._run_flush, name="event-emitter", daemon=True) + self._emitter_thread = Thread( + target=self._run_flush, name="event-emitter", daemon=True + ) self._emitter_thread.start() @cached_property @@ -137,12 +143,16 @@ def flush(self) -> None: except Exception as e: logger.info("Failed to emit events: %s", e) if isinstance(e, ApiClientError): - if e.code == 429 and self._emit_interval_sec < self._max_emit_interval_sec: + if ( + e.code == 429 + and self._emit_interval_sec < self._max_emit_interval_sec + ): self._emit_interval_sec = min( self._emit_interval_sec * 2, self._max_emit_interval_sec ) logger.debug( - "Increasing the emit interval to %s seconds", self._emit_interval_sec + "Increasing the emit interval to %s seconds", + self._emit_interval_sec, ) elif e.code in (400, 403, 404, 405, 426): logger.info( @@ -165,7 +175,9 @@ def shutdown(self, flush: bool = True) -> None: self.flush() self.emitter.close() - def _add_events(self, events: t.List[t.Dict[str, t.Any]], prepend: bool = False) -> None: + def _add_events( + self, events: t.List[t.Dict[str, t.Any]], prepend: bool = False + ) -> None: with self._events_lock: if not prepend: self._events.extend(events) diff --git a/sqlmesh/core/audit/__init__.py b/sqlmesh/core/audit/__init__.py index 65f77a8eca..3e4c414b86 100644 --- a/sqlmesh/core/audit/__init__.py +++ b/sqlmesh/core/audit/__init__.py @@ -1,7 +1,6 @@ -from sqlmesh.core.audit.definition import ( - Audit as Audit, - ModelAudit as ModelAudit, - StandaloneAudit as StandaloneAudit, - load_audit as load_audit, - load_multiple_audits as load_multiple_audits, -) +from sqlmesh.core.audit.definition import Audit as Audit +from sqlmesh.core.audit.definition import ModelAudit as ModelAudit +from sqlmesh.core.audit.definition import StandaloneAudit as StandaloneAudit +from sqlmesh.core.audit.definition import load_audit as load_audit +from sqlmesh.core.audit.definition import \ + load_multiple_audits as load_multiple_audits diff --git a/sqlmesh/core/audit/definition.py b/sqlmesh/core/audit/definition.py index 4c90151ee4..dbd3cefd17 100644 --- a/sqlmesh/core/audit/definition.py +++ b/sqlmesh/core/audit/definition.py @@ -11,25 +11,22 @@ from sqlmesh.core import dialect as d from sqlmesh.core.macros import MacroRegistry, macro -from sqlmesh.core.model.common import ( - bool_validator, - default_catalog_validator, - depends_on_validator, - sort_python_env, - sorted_python_env_payloads, -) -from sqlmesh.core.model.common import make_python_env, single_value_or_tuple, ParsableSql -from sqlmesh.core.node import _Node, DbtInfoMixin, DbtNodeInfo +from sqlmesh.core.model.common import (ParsableSql, bool_validator, + default_catalog_validator, + depends_on_validator, make_python_env, + single_value_or_tuple, sort_python_env, + sorted_python_env_payloads) +from sqlmesh.core.node import DbtInfoMixin, DbtNodeInfo, _Node from sqlmesh.core.renderer import QueryRenderer from sqlmesh.utils.date import TimeLike -from sqlmesh.utils.errors import AuditConfigError, SQLMeshError, raise_config_error +from sqlmesh.utils.errors import (AuditConfigError, SQLMeshError, + raise_config_error) from sqlmesh.utils.hashing import hash_data -from sqlmesh.utils.jinja import ( - JinjaMacroRegistry, - extract_macro_references_and_variables, -) +from sqlmesh.utils.jinja import (JinjaMacroRegistry, + extract_macro_references_and_variables) from sqlmesh.utils.metaprogramming import Executable -from sqlmesh.utils.pydantic import PydanticModel, field_validator, model_validator +from sqlmesh.utils.pydantic import (PydanticModel, field_validator, + model_validator) if t.TYPE_CHECKING: from sqlmesh.core._typing import Self @@ -111,10 +108,16 @@ def audit_map_validator(cls: t.Type, v: t.Any, values: t.Any) -> t.Dict[str, t.A if isinstance(v, dict): dialect = get_dialect(values) return { - key: value if isinstance(value, exp.Expr) else d.parse_one(str(value), dialect=dialect) + key: ( + value + if isinstance(value, exp.Expr) + else d.parse_one(str(value), dialect=dialect) + ) for key, value in v.items() } - raise_config_error("Defaults must be a tuple of exp.EQ or a dict", error_type=AuditConfigError) + raise_config_error( + "Defaults must be a tuple of exp.EQ or a dict", error_type=AuditConfigError + ) return {} @@ -132,7 +135,9 @@ class ModelAudit(PydanticModel, AuditMixin, DbtInfoMixin, frozen=True): standalone: t.Literal[False] = False query_: ParsableSql = Field(alias="query") defaults: t.Dict[str, exp.Expr] = {} - expressions_: t.Optional[t.List[ParsableSql]] = Field(default=None, alias="expressions") + expressions_: t.Optional[t.List[ParsableSql]] = Field( + default=None, alias="expressions" + ) jinja_macros: JinjaMacroRegistry = JinjaMacroRegistry() formatting: t.Optional[bool] = Field(default=None, exclude=True) dbt_node_info_: t.Optional[DbtNodeInfo] = Field(alias="dbt_node_info", default=None) @@ -168,7 +173,9 @@ class StandaloneAudit(_Node, AuditMixin): standalone: t.Literal[True] = True query_: ParsableSql = Field(alias="query") defaults: t.Dict[str, exp.Expr] = {} - expressions_: t.Optional[t.List[ParsableSql]] = Field(default=None, alias="expressions") + expressions_: t.Optional[t.List[ParsableSql]] = Field( + default=None, alias="expressions" + ) jinja_macros: JinjaMacroRegistry = JinjaMacroRegistry() default_catalog: t.Optional[str] = None depends_on_: t.Optional[t.Set[str]] = Field(default=None, alias="depends_on") @@ -188,7 +195,9 @@ class StandaloneAudit(_Node, AuditMixin): @model_validator(mode="after") def _node_root_validator(self) -> Self: if self.blocking: - raise AuditConfigError(f"Standalone audits cannot be blocking: '{self.name}'.") + raise AuditConfigError( + f"Standalone audits cannot be blocking: '{self.name}'." + ) return self def render_audit_query( @@ -342,7 +351,9 @@ def render_definition( else: expression = exp.Property( this=field_info.alias or field_name, - value=META_FIELD_CONVERTER.get(field_name, exp.to_identifier)(field_value), + value=META_FIELD_CONVERTER.get(field_name, exp.to_identifier)( + field_value + ), ) if field_name == "name": expressions.insert(0, expression) @@ -355,7 +366,9 @@ def render_definition( jinja_expressions = [] python_expressions = [] if include_python: - python_env = d.PythonCode(expressions=sorted_python_env_payloads(self.python_env)) + python_env = d.PythonCode( + expressions=sorted_python_env_payloads(self.python_env) + ) if python_env.expressions: python_expressions.append(python_env) @@ -443,7 +456,9 @@ def load_audit( extra_fields = audit_class.extra_fields(set(meta_fields)) if extra_fields: - _raise_config_error(f"Invalid extra fields {extra_fields} in the audit definition", path) + _raise_config_error( + f"Invalid extra fields {extra_fields} in the audit definition", path + ) if not isinstance(query, exp.Query) and not isinstance(query, d.JinjaQuery): _raise_config_error("Missing SELECT query in the audit definition", path) @@ -451,11 +466,15 @@ def load_audit( extra_kwargs: t.Dict[str, t.Any] = {} if is_standalone: - jinja_macro_refrences, referenced_variables = extract_macro_references_and_variables( - *(gen(s) for s in statements), - gen(query), + jinja_macro_refrences, referenced_variables = ( + extract_macro_references_and_variables( + *(gen(s) for s in statements), + gen(query), + ) + ) + jinja_macros = (jinja_macros or JinjaMacroRegistry()).trim( + jinja_macro_refrences ) - jinja_macros = (jinja_macros or JinjaMacroRegistry()).trim(jinja_macro_refrences) for jinja_macro in jinja_macros.root_macros.values(): referenced_variables.update( extract_macro_references_and_variables(jinja_macro.definition)[1] @@ -476,9 +495,12 @@ def load_audit( dialect = meta_fields.pop("dialect", dialect) or "" - parsable_query = ParsableSql.from_parsed_expression(query, dialect, use_meta_sql=True) + parsable_query = ParsableSql.from_parsed_expression( + query, dialect, use_meta_sql=True + ) parsable_statements = [ - ParsableSql.from_parsed_expression(s, dialect, use_meta_sql=True) for s in statements + ParsableSql.from_parsed_expression(s, dialect, use_meta_sql=True) + for s in statements ] try: diff --git a/sqlmesh/core/config/__init__.py b/sqlmesh/core/config/__init__.py index 50d2d9a5a2..16a0e5deb9 100644 --- a/sqlmesh/core/config/__init__.py +++ b/sqlmesh/core/config/__init__.py @@ -1,42 +1,61 @@ -from sqlmesh.core.config.categorizer import ( - AutoCategorizationMode as AutoCategorizationMode, - CategorizerConfig as CategorizerConfig, -) -from sqlmesh.core.config.common import ( - EnvironmentSuffixTarget as EnvironmentSuffixTarget, - TableNamingConvention as TableNamingConvention, -) -from sqlmesh.core.config.connection import ( - AthenaConnectionConfig as AthenaConnectionConfig, - BaseDuckDBConnectionConfig as BaseDuckDBConnectionConfig, - BigQueryConnectionConfig as BigQueryConnectionConfig, - ConnectionConfig as ConnectionConfig, - DatabricksConnectionConfig as DatabricksConnectionConfig, - DuckDBConnectionConfig as DuckDBConnectionConfig, - FabricConnectionConfig as FabricConnectionConfig, - GCPPostgresConnectionConfig as GCPPostgresConnectionConfig, - MotherDuckConnectionConfig as MotherDuckConnectionConfig, - MSSQLConnectionConfig as MSSQLConnectionConfig, - MySQLConnectionConfig as MySQLConnectionConfig, - PostgresConnectionConfig as PostgresConnectionConfig, - RedshiftConnectionConfig as RedshiftConnectionConfig, - SnowflakeConnectionConfig as SnowflakeConnectionConfig, - SparkConnectionConfig as SparkConnectionConfig, - StarRocksConnectionConfig as StarRocksConnectionConfig, - TrinoConnectionConfig as TrinoConnectionConfig, - parse_connection_config as parse_connection_config, -) +from sqlmesh.core.config.categorizer import \ + AutoCategorizationMode as AutoCategorizationMode +from sqlmesh.core.config.categorizer import \ + CategorizerConfig as CategorizerConfig +from sqlmesh.core.config.common import \ + EnvironmentSuffixTarget as EnvironmentSuffixTarget +from sqlmesh.core.config.common import \ + TableNamingConvention as TableNamingConvention +from sqlmesh.core.config.connection import \ + AthenaConnectionConfig as AthenaConnectionConfig +from sqlmesh.core.config.connection import \ + BaseDuckDBConnectionConfig as BaseDuckDBConnectionConfig +from sqlmesh.core.config.connection import \ + BigQueryConnectionConfig as BigQueryConnectionConfig +from sqlmesh.core.config.connection import ConnectionConfig as ConnectionConfig +from sqlmesh.core.config.connection import \ + DatabricksConnectionConfig as DatabricksConnectionConfig +from sqlmesh.core.config.connection import \ + DuckDBConnectionConfig as DuckDBConnectionConfig +from sqlmesh.core.config.connection import \ + FabricConnectionConfig as FabricConnectionConfig +from sqlmesh.core.config.connection import \ + GCPPostgresConnectionConfig as GCPPostgresConnectionConfig +from sqlmesh.core.config.connection import \ + MotherDuckConnectionConfig as MotherDuckConnectionConfig +from sqlmesh.core.config.connection import \ + MSSQLConnectionConfig as MSSQLConnectionConfig +from sqlmesh.core.config.connection import \ + MySQLConnectionConfig as MySQLConnectionConfig +from sqlmesh.core.config.connection import \ + PostgresConnectionConfig as PostgresConnectionConfig +from sqlmesh.core.config.connection import \ + RedshiftConnectionConfig as RedshiftConnectionConfig +from sqlmesh.core.config.connection import \ + SnowflakeConnectionConfig as SnowflakeConnectionConfig +from sqlmesh.core.config.connection import \ + SparkConnectionConfig as SparkConnectionConfig +from sqlmesh.core.config.connection import \ + StarRocksConnectionConfig as StarRocksConnectionConfig +from sqlmesh.core.config.connection import \ + TrinoConnectionConfig as TrinoConnectionConfig +from sqlmesh.core.config.connection import \ + parse_connection_config as parse_connection_config from sqlmesh.core.config.gateway import GatewayConfig as GatewayConfig -from sqlmesh.core.config.loader import ( - load_config_from_paths as load_config_from_paths, - load_config_from_yaml as load_config_from_yaml, - load_configs as load_configs, -) -from sqlmesh.core.config.migration import MigrationConfig as MigrationConfig -from sqlmesh.core.config.model import ModelDefaultsConfig as ModelDefaultsConfig -from sqlmesh.core.config.naming import NameInferenceConfig as NameInferenceConfig from sqlmesh.core.config.linter import LinterConfig as LinterConfig +from sqlmesh.core.config.loader import \ + load_config_from_paths as load_config_from_paths +from sqlmesh.core.config.loader import \ + load_config_from_yaml as load_config_from_yaml +from sqlmesh.core.config.loader import load_configs as load_configs +from sqlmesh.core.config.migration import MigrationConfig as MigrationConfig +from sqlmesh.core.config.model import \ + ModelDefaultsConfig as ModelDefaultsConfig +from sqlmesh.core.config.naming import \ + NameInferenceConfig as NameInferenceConfig from sqlmesh.core.config.plan import PlanConfig as PlanConfig -from sqlmesh.core.config.root import Config as Config, DbtConfig as DbtConfig +from sqlmesh.core.config.root import Config as Config +from sqlmesh.core.config.root import DbtConfig as DbtConfig from sqlmesh.core.config.run import RunConfig as RunConfig -from sqlmesh.core.config.scheduler import BuiltInSchedulerConfig as BuiltInSchedulerConfig +from sqlmesh.core.config.scheduler import \ + BuiltInSchedulerConfig as BuiltInSchedulerConfig diff --git a/sqlmesh/core/config/base.py b/sqlmesh/core/config/base.py index 0da36e4754..7256ba2b27 100644 --- a/sqlmesh/core/config/base.py +++ b/sqlmesh/core/config/base.py @@ -74,12 +74,16 @@ def _update_pydantic_config(old: BaseConfig, new: BaseConfig) -> PydanticModel: combined = old.copy() for key, value in new.items(): if not isinstance(value, list): - raise ConfigError("KEY_EXTEND behavior requires list values in dictionary.") + raise ConfigError( + "KEY_EXTEND behavior requires list values in dictionary." + ) old_value = combined.get(key) if old_value: if not isinstance(old_value, list): - raise ConfigError("KEY_EXTEND behavior requires list values in dictionary.") + raise ConfigError( + "KEY_EXTEND behavior requires list values in dictionary." + ) combined[key] = old_value + value else: diff --git a/sqlmesh/core/config/common.py b/sqlmesh/core/config/common.py index dca472d7a9..4b15921f96 100644 --- a/sqlmesh/core/config/common.py +++ b/sqlmesh/core/config/common.py @@ -1,15 +1,21 @@ from __future__ import annotations +import re import typing as t from enum import Enum -import re from sqlmesh.utils import classproperty from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.pydantic import field_validator # Config files that can be present in the project dir -ALL_CONFIG_FILENAMES = ("config.py", "config.yml", "config.yaml", "sqlmesh.yml", "sqlmesh.yaml") +ALL_CONFIG_FILENAMES = ( + "config.py", + "config.yml", + "config.yaml", + "sqlmesh.yml", + "sqlmesh.yaml", +) # For personal paths (~/.sqlmesh/) where python config is not supported YAML_CONFIG_FILENAMES = tuple(n for n in ALL_CONFIG_FILENAMES if not n.endswith(".py")) @@ -173,7 +179,9 @@ def _validate_type(v: t.Any) -> None: )(_variables_validator) -def compile_regex_mapping(value: t.Dict[str | re.Pattern, t.Any]) -> t.Dict[re.Pattern, t.Any]: +def compile_regex_mapping( + value: t.Dict[str | re.Pattern, t.Any], +) -> t.Dict[re.Pattern, t.Any]: """ Utility function to compile a dict of { "string regex pattern" : "string value" } into { "": "string value" } """ diff --git a/sqlmesh/core/config/connection.py b/sqlmesh/core/config/connection.py index 73fe1b9300..c7a5960f45 100644 --- a/sqlmesh/core/config/connection.py +++ b/sqlmesh/core/config/connection.py @@ -13,8 +13,8 @@ from sys import version_info import pydantic -from pydantic import Field, computed_field from packaging import version +from pydantic import Field, computed_field from pydantic_core import from_json from sqlglot import exp from sqlglot.errors import ParseError @@ -22,24 +22,18 @@ from sqlmesh.core import engine_adapter from sqlmesh.core.config.base import BaseConfig -from sqlmesh.core.config.common import ( - compile_regex_mapping, - concurrent_tasks_validator, - http_headers_validator, -) +from sqlmesh.core.config.common import (compile_regex_mapping, + concurrent_tasks_validator, + http_headers_validator) from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.engine_adapter.shared import CatalogSupport from sqlmesh.utils import debug_mode_enabled, str_to_bool from sqlmesh.utils.aws import validate_s3_uri from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.pydantic import ( - ValidationInfo, - field_validator, - get_concrete_types_from_typehint, - model_validator, - validation_data, - validation_error_message, -) +from sqlmesh.utils.pydantic import (ValidationInfo, field_validator, + get_concrete_types_from_typehint, + model_validator, validation_data, + validation_error_message) if t.TYPE_CHECKING: from sqlmesh.core._typing import Self @@ -67,13 +61,18 @@ def _get_engine_import_validator( - import_name: str, engine_type: str, extra_name: t.Optional[str] = None, decorate: bool = True + import_name: str, + engine_type: str, + extra_name: t.Optional[str] = None, + decorate: bool = True, ) -> t.Callable: extra_name = extra_name or engine_type def validate(cls: t.Any, data: t.Any) -> t.Any: check_import = ( - str_to_bool(str(data.pop("check_import", True))) if isinstance(data, dict) else True + str_to_bool(str(data.pop("check_import", True))) + if isinstance(data, dict) + else True ) if not check_import: return data @@ -166,7 +165,11 @@ def _connection_factory_with_kwargs(self) -> t.Callable[[], t.Any]: self._connection_factory, **{ **self._static_connection_kwargs, - **{k: v for k, v in self.dict().items() if k in self._connection_kwargs_keys}, + **{ + k: v + for k, v in self.dict().items() + if k in self._connection_kwargs_keys + }, }, ) @@ -175,7 +178,9 @@ def connection_validator(self) -> t.Callable[[], None]: return self.create_engine_adapter().ping def create_engine_adapter( - self, register_comments_override: bool = False, concurrent_tasks: t.Optional[int] = None + self, + register_comments_override: bool = False, + concurrent_tasks: t.Optional[int] = None, ) -> EngineAdapter: """Returns a new instance of the Engine Adapter.""" @@ -219,7 +224,9 @@ def _expand_json_strings_to_concrete_types(cls, data: t.Any) -> t.Any: """ if data and isinstance(data, dict): for maybe_json_field_name in cls._get_list_and_dict_field_names(): - if (value := data.get(maybe_json_field_name)) and isinstance(value, str): + if (value := data.get(maybe_json_field_name)) and isinstance( + value, str + ): # crude JSON check as we dont want to try and parse every string we get value = value.strip() if value.startswith("{") or value.startswith("["): @@ -274,7 +281,9 @@ def to_sql(self, alias: str) -> str: if self.encrypted: options.append("ENCRYPTED") if self.data_inlining_row_limit is not None: - options.append(f"DATA_INLINING_ROW_LIMIT {self.data_inlining_row_limit}") + options.append( + f"DATA_INLINING_ROW_LIMIT {self.data_inlining_row_limit}" + ) if self.metadata_schema is not None: options.append(f"METADATA_SCHEMA '{self.metadata_schema}'") @@ -284,7 +293,9 @@ def to_sql(self, alias: str) -> str: # MotherDuck does not support aliasing alias_sql = ( - f" AS {alias}" if not (self.type == "motherduck" or self.path.startswith("md:")) else "" + f" AS {alias}" + if not (self.type == "motherduck" or self.path.startswith("md:")) + else "" ) return f"ATTACH IF NOT EXISTS '{path}'{alias_sql}{options_sql}" @@ -364,12 +375,16 @@ def _cursor_init(self) -> t.Optional[t.Callable[[t.Any], None]]: def init(cursor: duckdb.DuckDBPyConnection) -> None: for extension in self.extensions: - extension = extension if isinstance(extension, dict) else {"name": extension} + extension = ( + extension if isinstance(extension, dict) else {"name": extension} + ) install_command = f"INSTALL {extension['name']}" if extension.get("repository"): - install_command = f"{install_command} FROM {extension['repository']}" + install_command = ( + f"{install_command} FROM {extension['repository']}" + ) if extension.get("force_install"): install_command = f"FORCE {install_command}" @@ -378,7 +393,9 @@ def init(cursor: duckdb.DuckDBPyConnection) -> None: cursor.execute(install_command) cursor.execute(f"LOAD {extension['name']}") except Exception as e: - raise ConfigError(f"Failed to load extension {extension['name']}: {e}") + raise ConfigError( + f"Failed to load extension {extension['name']}: {e}" + ) if self.connector_config: option_names = list(self.connector_config) @@ -389,7 +406,9 @@ def init(cursor: duckdb.DuckDBPyConnection) -> None: option_names, ) - existing_values = {field: setting for field, setting in cursor.fetchall()} + existing_values = { + field: setting for field, setting in cursor.fetchall() + } # only set connector_config items if the values differ from what is already set # trying to set options like 'temp_directory' even to the same value can throw errors like: @@ -415,7 +434,9 @@ def init(cursor: duckdb.DuckDBPyConnection) -> None: ) else: if isinstance(self.secrets, list): - secrets_items = [(secret_dict, "") for secret_dict in self.secrets] + secrets_items = [ + (secret_dict, "") for secret_dict in self.secrets + ] else: secrets_items = [ (secret_dict, secret_name) @@ -466,13 +487,11 @@ def init(cursor: duckdb.DuckDBPyConnection) -> None: # set it as the default catalog. # If a user tried to attach a MotherDuck database/share which has already by attached via # `ATTACH 'md:'`, then we don't want to raise since this is expected. - if ( - not ( - 'database with name "memory" already exists' in str(e) - and path_options == ":memory:" - ) - and f"""database with name "{path_options.path.replace("md:", "")}" already exists""" - not in str(e) + if not ( + 'database with name "memory" already exists' in str(e) + and path_options == ":memory:" + ) and f"""database with name "{path_options.path.replace("md:", "")}" already exists""" not in str( + e ): raise e if i == 0 and not getattr(self, "database", None): @@ -481,7 +500,9 @@ def init(cursor: duckdb.DuckDBPyConnection) -> None: return init def create_engine_adapter( - self, register_comments_override: bool = False, concurrent_tasks: t.Optional[int] = None + self, + register_comments_override: bool = False, + concurrent_tasks: t.Optional[int] = None, ) -> EngineAdapter: """Checks if another engine adapter has already been created that shares a catalog that points to the same data file. If so, it uses that same adapter instead of creating a new one. As a result, any additional configuration @@ -562,7 +583,9 @@ def _static_connection_kwargs(self) -> t.Dict[str, t.Any]: # Attach single MD database instead of all databases on the account connection_str += f"{self.database}?attach_mode=single" if self.token: - connection_str += f"{'&' if self.database else '?'}motherduck_token={self.token}" + connection_str += ( + f"{'&' if self.database else '?'}motherduck_token={self.token}" + ) return {"database": connection_str, "config": custom_user_agent_config} @property @@ -638,7 +661,8 @@ def _validate_authenticator(cls, data: t.Any) -> t.Any: if not isinstance(data, dict): return data - from snowflake.connector.network import DEFAULT_AUTHENTICATOR, OAUTH_AUTHENTICATOR + from snowflake.connector.network import (DEFAULT_AUTHENTICATOR, + OAUTH_AUTHENTICATOR) auth = data.get("authenticator") auth = auth.upper() if auth else DEFAULT_AUTHENTICATOR @@ -651,7 +675,9 @@ def _validate_authenticator(cls, data: t.Any) -> t.Any: and not data.get("private_key") and (not user or not password) ): - raise ConfigError("User and password must be provided if using default authentication") + raise ConfigError( + "User and password must be provided if using default authentication" + ) if auth == OAUTH_AUTHENTICATOR and not data.get("token"): raise ConfigError("Token must be provided if using oauth authentication") @@ -663,7 +689,9 @@ def _validate_authenticator(cls, data: t.Any) -> t.Any: ) @classmethod - def _get_private_key(cls, values: t.Dict[str, t.Optional[str]], auth: str) -> t.Optional[bytes]: + def _get_private_key( + cls, values: t.Dict[str, t.Optional[str]], auth: str + ) -> t.Optional[bytes]: """ source: https://github.com/dbt-labs/dbt-snowflake/blob/0374b4ec948982f2ac8ec0c95d53d672ad19e09c/dbt/adapters/snowflake/connections.py#L247C5-L285C1 @@ -672,22 +700,24 @@ def _get_private_key(cls, values: t.Dict[str, t.Optional[str]], auth: str) -> t. # Start custom code from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization - from snowflake.connector.network import ( - DEFAULT_AUTHENTICATOR, - KEY_PAIR_AUTHENTICATOR, - ) + from snowflake.connector.network import (DEFAULT_AUTHENTICATOR, + KEY_PAIR_AUTHENTICATOR) private_key = values.get("private_key") private_key_path = values.get("private_key_path") private_key_passphrase = values.get("private_key_passphrase") user = values.get("user") password = values.get("password") - auth = auth if auth and auth != DEFAULT_AUTHENTICATOR else KEY_PAIR_AUTHENTICATOR + auth = ( + auth if auth and auth != DEFAULT_AUTHENTICATOR else KEY_PAIR_AUTHENTICATOR + ) if not private_key and not private_key_path: return None if private_key and private_key_path: - raise ConfigError("Cannot specify both `private_key` and `private_key_path`") + raise ConfigError( + "Cannot specify both `private_key` and `private_key_path`" + ) if auth != KEY_PAIR_AUTHENTICATOR: raise ConfigError( f"Private key or private key path can only be provided when using {KEY_PAIR_AUTHENTICATOR} authentication" @@ -850,14 +880,17 @@ def _databricks_connect_validator(cls, data: t.Any) -> t.Any: if not isinstance(data, dict): return data - from sqlmesh.core.engine_adapter.databricks import DatabricksEngineAdapter + from sqlmesh.core.engine_adapter.databricks import \ + DatabricksEngineAdapter if DatabricksEngineAdapter.can_access_spark_session( bool(data.get("disable_spark_session")) ): return data - databricks_connect_use_serverless = data.get("databricks_connect_use_serverless") + databricks_connect_use_serverless = data.get( + "databricks_connect_use_serverless" + ) server_hostname, http_path, access_token, auth_type = ( data.get("server_hostname"), data.get("http_path"), @@ -885,7 +918,9 @@ def _databricks_connect_validator(cls, data: t.Any) -> t.Any: if not data.get("databricks_connect_access_token"): data["databricks_connect_access_token"] = access_token if not data.get("databricks_connect_server_hostname"): - data["databricks_connect_server_hostname"] = f"https://{server_hostname}" + data["databricks_connect_server_hostname"] = ( + f"https://{server_hostname}" + ) if not databricks_connect_use_serverless and not data.get( "databricks_connect_cluster_id" ): @@ -945,7 +980,8 @@ def _extra_engine_config(self) -> t.Dict[str, t.Any]: @property def use_spark_session_only(self) -> bool: - from sqlmesh.core.engine_adapter.databricks import DatabricksEngineAdapter + from sqlmesh.core.engine_adapter.databricks import \ + DatabricksEngineAdapter return ( DatabricksEngineAdapter.can_access_spark_session(self.disable_spark_session) @@ -965,7 +1001,8 @@ def _connection_factory(self) -> t.Callable: @property def _static_connection_kwargs(self) -> t.Dict[str, t.Any]: - from sqlmesh.core.engine_adapter.databricks import DatabricksEngineAdapter + from sqlmesh.core.engine_adapter.databricks import \ + DatabricksEngineAdapter if not self.use_spark_session_only: conn_kwargs: t.Dict[str, t.Any] = { @@ -978,14 +1015,17 @@ def _static_connection_kwargs(self) -> t.Dict[str, t.Any]: # if a client_secret exists, then a client_id also exists and we are using M2M # ref: https://docs.databricks.com/en/dev-tools/python-sql-connector.html#oauth-machine-to-machine-m2m-authentication # ref: https://github.com/databricks/databricks-sql-python/blob/main/examples/m2m_oauth.py - from databricks.sdk.core import Config, oauth_service_principal + from databricks.sdk.core import (Config, + oauth_service_principal) config = Config( host=f"https://{self.server_hostname}", client_id=self.oauth_client_id, client_secret=self.oauth_client_secret, ) - conn_kwargs["credentials_provider"] = lambda: oauth_service_principal(config) + conn_kwargs["credentials_provider"] = ( + lambda: oauth_service_principal(config) + ) else: # if auth_type is set to an 'oauth' type but no client_id/secret are set, then we are using U2M # ref: https://docs.databricks.com/en/dev-tools/python-sql-connector.html#oauth-user-to-machine-u2m-authentication @@ -1096,7 +1136,9 @@ class BigQueryConnectionConfig(ConnectionConfig): DISPLAY_NAME: t.ClassVar[t.Literal["BigQuery"]] = "BigQuery" DISPLAY_ORDER: t.ClassVar[t.Literal[4]] = 4 - _engine_import_validator = _get_engine_import_validator("google.cloud.bigquery", "bigquery") + _engine_import_validator = _get_engine_import_validator( + "google.cloud.bigquery", "bigquery" + ) @field_validator("execution_project") def validate_execution_project( @@ -1229,7 +1271,9 @@ class GCPPostgresConnectionConfig(ConnectionConfig): password: t.Optional[str] = None enable_iam_auth: t.Optional[bool] = None db: str - ip_type: t.Union[t.Literal["public"], t.Literal["private"], t.Literal["psc"]] = "public" + ip_type: t.Union[t.Literal["public"], t.Literal["private"], t.Literal["psc"]] = ( + "public" + ) # Keyfile Auth keyfile: t.Optional[str] = None keyfile_json: t.Optional[t.Dict[str, t.Any]] = None @@ -1376,7 +1420,9 @@ class RedshiftConnectionConfig(ConnectionConfig): DISPLAY_NAME: t.ClassVar[t.Literal["Redshift"]] = "Redshift" DISPLAY_ORDER: t.ClassVar[t.Literal[7]] = 7 - _engine_import_validator = _get_engine_import_validator("redshift_connector", "redshift") + _engine_import_validator = _get_engine_import_validator( + "redshift_connector", "redshift" + ) @property def _connection_kwargs_keys(self) -> t.Set[str]: @@ -1653,7 +1699,9 @@ def _connection_factory(self) -> t.Callable: # with the `pyodbc` equivalent for documented parameters. if not SUPPORTS_MSSQL_PYTHON_DRIVER: - raise ConfigError("The `mssql-python` driver requires Python 3.10 or higher.") + raise ConfigError( + "The `mssql-python` driver requires Python 3.10 or higher." + ) import mssql_python @@ -1690,7 +1738,9 @@ def connect_mssql_python(**kwargs: t.Any) -> t.Callable: # - https://github.com/microsoft/mssql-python/wiki/Connection-to-SQL-Database # - https://github.com/microsoft/mssql-python/wiki/Connection#timeout conn_str_parts.append(f"ConnectRetryCount={login_attempts}") - conn_str_parts.append(f"ConnectRetryInterval={min(int(login_timeout), 60)}") + conn_str_parts.append( + f"ConnectRetryInterval={min(int(login_timeout), 60)}" + ) # Standard SQL Server authentication if user: @@ -1824,7 +1874,9 @@ def connect_pyodbc(**kwargs: t.Any) -> t.Callable: import pyodbc - conn = pyodbc.connect(conn_str, autocommit=kwargs.get("autocommit", False)) + conn = pyodbc.connect( + conn_str, autocommit=kwargs.get("autocommit", False) + ) # Set up output converters for MSSQL-specific data types # Handle SQL type -155 (DATETIMEOFFSET) which is not yet supported by pyodbc @@ -2032,7 +2084,9 @@ class TrinoConnectionConfig(ConnectionConfig): timezone: t.Optional[str] = None # Basic/LDAP password: t.Optional[str] = None - verify: t.Optional[bool] = None # disable SSL verification (ignored if `cert` is provided) + verify: t.Optional[bool] = ( + None # disable SSL verification (ignored if `cert` is provided) + ) # LDAP impersonation_user: t.Optional[str] = None # Kerberos @@ -2111,13 +2165,21 @@ def _validate_timestamp_mapping( @model_validator(mode="after") def _root_validator(self) -> Self: port = self.port - if self.http_scheme == "http" and not self.method.is_no_auth and not self.method.is_basic: - raise ConfigError("HTTP scheme can only be used with no-auth or basic method") + if ( + self.http_scheme == "http" + and not self.method.is_no_auth + and not self.method.is_basic + ): + raise ConfigError( + "HTTP scheme can only be used with no-auth or basic method" + ) if port is None: self.port = 80 if self.http_scheme == "http" else 443 - if (self.method.is_ldap or self.method.is_basic) and (not self.password or not self.user): + if (self.method.is_ldap or self.method.is_basic) and ( + not self.password or not self.user + ): raise ConfigError( f"Username and Password must be provided if using {self.method.value} authentication" ) @@ -2168,13 +2230,9 @@ def _connection_factory(self) -> t.Callable: @property def _static_connection_kwargs(self) -> t.Dict[str, t.Any]: - from trino.auth import ( - BasicAuthentication, - CertificateAuthentication, - JWTAuthentication, - KerberosAuthentication, - OAuth2Authentication, - ) + from trino.auth import (BasicAuthentication, CertificateAuthentication, + JWTAuthentication, KerberosAuthentication, + OAuth2Authentication) auth: t.Optional[ t.Union[ @@ -2186,7 +2244,9 @@ def _static_connection_kwargs(self) -> t.Dict[str, t.Any]: ] ] = None if self.method.is_basic or self.method.is_ldap: - assert self.password is not None # for mypy since validator already checks this + assert ( + self.password is not None + ) # for mypy since validator already checks this auth = BasicAuthentication(self.user, self.password) elif self.method.is_kerberos: if self.keytab: @@ -2210,7 +2270,9 @@ def _static_connection_kwargs(self) -> t.Dict[str, t.Any]: elif self.method.is_certificate: assert self.client_certificate is not None assert self.client_private_key is not None - auth = CertificateAuthentication(self.client_certificate, self.client_private_key) + auth = CertificateAuthentication( + self.client_certificate, self.client_private_key + ) return { "auth": auth, @@ -2273,7 +2335,9 @@ class ClickhouseConnectionConfig(ConnectionConfig): DISPLAY_NAME: t.ClassVar[t.Literal["ClickHouse"]] = "ClickHouse" DISPLAY_ORDER: t.ClassVar[t.Literal[6]] = 6 - _engine_import_validator = _get_engine_import_validator("clickhouse_connect", "clickhouse") + _engine_import_validator = _get_engine_import_validator( + "clickhouse_connect", "clickhouse" + ) @field_validator("virtual_catalog") def validate_virtual_catalog(cls, v: t.Optional[str]) -> t.Optional[str]: @@ -2428,10 +2492,14 @@ def _root_validator(self) -> Self: s3_warehouse_location = self.s3_warehouse_location if not work_group and not s3_staging_dir: - raise ConfigError("At least one of work_group or s3_staging_dir must be set") + raise ConfigError( + "At least one of work_group or s3_staging_dir must be set" + ) if s3_staging_dir: - self.s3_staging_dir = validate_s3_uri(s3_staging_dir, base=True, error_type=ConfigError) + self.s3_staging_dir = validate_s3_uri( + s3_staging_dir, base=True, error_type=ConfigError + ) if s3_warehouse_location: self.s3_warehouse_location = validate_s3_uri( @@ -2605,12 +2673,16 @@ def _connection_factory(self) -> t.Callable: CONNECTION_CONFIG_TO_TYPE = { # Map all subclasses of ConnectionConfig to the value of their `type_` field. tpe.all_field_infos()["type_"].default: tpe - for tpe in subclasses(__name__, ConnectionConfig, exclude=_CONNECTION_CONFIG_EXCLUDE) + for tpe in subclasses( + __name__, ConnectionConfig, exclude=_CONNECTION_CONFIG_EXCLUDE + ) } DIALECT_TO_TYPE = { tpe.all_field_infos()["type_"].default: tpe.DIALECT - for tpe in subclasses(__name__, ConnectionConfig, exclude=_CONNECTION_CONFIG_EXCLUDE) + for tpe in subclasses( + __name__, ConnectionConfig, exclude=_CONNECTION_CONFIG_EXCLUDE + ) } INIT_DISPLAY_INFO_TO_TYPE = { @@ -2618,7 +2690,9 @@ def _connection_factory(self) -> t.Callable: tpe.DISPLAY_ORDER, tpe.DISPLAY_NAME, ) - for tpe in subclasses(__name__, ConnectionConfig, exclude=_CONNECTION_CONFIG_EXCLUDE) + for tpe in subclasses( + __name__, ConnectionConfig, exclude=_CONNECTION_CONFIG_EXCLUDE + ) } diff --git a/sqlmesh/core/config/gateway.py b/sqlmesh/core/config/gateway.py index a51557c4d7..08f5d26fe4 100644 --- a/sqlmesh/core/config/gateway.py +++ b/sqlmesh/core/config/gateway.py @@ -4,13 +4,12 @@ from sqlmesh.core import constants as c from sqlmesh.core.config.base import BaseConfig -from sqlmesh.core.config.model import ModelDefaultsConfig from sqlmesh.core.config.common import variables_validator -from sqlmesh.core.config.connection import ( - SerializableConnectionConfig, - connection_config_validator, -) -from sqlmesh.core.config.scheduler import SchedulerConfig, scheduler_config_validator +from sqlmesh.core.config.connection import (SerializableConnectionConfig, + connection_config_validator) +from sqlmesh.core.config.model import ModelDefaultsConfig +from sqlmesh.core.config.scheduler import (SchedulerConfig, + scheduler_config_validator) class GatewayConfig(BaseConfig): diff --git a/sqlmesh/core/config/linter.py b/sqlmesh/core/config/linter.py index 11d700c540..7b0a10a94e 100644 --- a/sqlmesh/core/config/linter.py +++ b/sqlmesh/core/config/linter.py @@ -6,7 +6,6 @@ from sqlglot.helper import ensure_collection from sqlmesh.core.config.base import BaseConfig - from sqlmesh.utils.pydantic import field_validator diff --git a/sqlmesh/core/config/loader.py b/sqlmesh/core/config/loader.py index a3b6b213ab..50a45b156a 100644 --- a/sqlmesh/core/config/loader.py +++ b/sqlmesh/core/config/loader.py @@ -5,16 +5,14 @@ import typing as t from pathlib import Path -from pydantic import ValidationError from dotenv import load_dotenv +from pydantic import ValidationError from sqlglot.helper import ensure_list from sqlmesh.core import constants as c -from sqlmesh.core.config.common import ( - ALL_CONFIG_FILENAMES, - YAML_CONFIG_FILENAMES, - DBT_PROJECT_FILENAME, -) +from sqlmesh.core.config.common import (ALL_CONFIG_FILENAMES, + DBT_PROJECT_FILENAME, + YAML_CONFIG_FILENAMES) from sqlmesh.core.config.model import ModelDefaultsConfig from sqlmesh.core.config.root import Config from sqlmesh.utils import env_vars, merge_dicts, sys_path @@ -107,7 +105,9 @@ def load_config_from_paths( parent_path = path.parent if parent_path in visited_folders: - raise ConfigError(f"Multiple configuration files found in folder '{parent_path}'.") + raise ConfigError( + f"Multiple configuration files found in folder '{parent_path}'." + ) visited_folders.add(parent_path) extension = path.name.split(".")[-1].lower() @@ -125,7 +125,9 @@ def load_config_from_paths( ) except ValidationError as e: raise ConfigError( - validation_error_message(e, f"Invalid project config '{config_name}':") + validation_error_message( + e, f"Invalid project config '{config_name}':" + ) + "\n\nVerify your config.py." ) else: @@ -168,7 +170,9 @@ def load_config_from_paths( # any config within yaml files will get overlayed on top of it. if not python_config: potential_project_files = [f / DBT_PROJECT_FILENAME for f in visited_folders] - dbt_project_file = next((f for f in potential_project_files if f.exists()), None) + dbt_project_file = next( + (f for f in potential_project_files if f.exists()), None + ) if dbt_project_file: from sqlmesh.dbt.loader import sqlmesh_config @@ -253,7 +257,10 @@ def load_config_from_env() -> t.Dict[str, t.Any]: for key, value in os.environ.items(): key = key.lower() - if key.startswith(f"{c.SQLMESH}__") and key != (c.DISABLE_SQLMESH_STATE_MIGRATION).lower(): + if ( + key.startswith(f"{c.SQLMESH}__") + and key != (c.DISABLE_SQLMESH_STATE_MIGRATION).lower() + ): segments = key.split("__")[1:] if not segments or not segments[-1]: raise ConfigError(f"Invalid SQLMesh configuration variable '{key}'.") diff --git a/sqlmesh/core/config/model.py b/sqlmesh/core/config/model.py index e7af572cc7..14cb494806 100644 --- a/sqlmesh/core/config/model.py +++ b/sqlmesh/core/config/model.py @@ -3,16 +3,13 @@ import typing as t from sqlglot import exp -from sqlmesh.core.dialect import parse_one, extract_func_call + from sqlmesh.core.config.base import BaseConfig -from sqlmesh.core.model.kind import ( - ModelKind, - OnDestructiveChange, - model_kind_validator, - on_destructive_change_validator, - on_additive_change_validator, - OnAdditiveChange, -) +from sqlmesh.core.dialect import extract_func_call, parse_one +from sqlmesh.core.model.kind import (ModelKind, OnAdditiveChange, + OnDestructiveChange, model_kind_validator, + on_additive_change_validator, + on_destructive_change_validator) from sqlmesh.core.model.meta import FunctionCall from sqlmesh.core.node import IntervalUnit, cron_tz_validator from sqlmesh.utils.date import TimeLike diff --git a/sqlmesh/core/config/root.py b/sqlmesh/core/config/root.py index b36b7dadc1..68150a17dc 100644 --- a/sqlmesh/core/config/root.py +++ b/sqlmesh/core/config/root.py @@ -13,40 +13,36 @@ from sqlmesh.cicd.config import CICDBotConfig from sqlmesh.core import constants as c -from sqlmesh.core.console import get_console -from sqlmesh.core.config.common import ( - EnvironmentSuffixTarget, - TableNamingConvention, - VirtualEnvironmentMode, -) from sqlmesh.core.config.base import BaseConfig, UpdateStrategy -from sqlmesh.core.config.common import variables_validator, compile_regex_mapping -from sqlmesh.core.config.connection import ( - ConnectionConfig, - DuckDBConnectionConfig, - SerializableConnectionConfig, - connection_config_validator, -) +from sqlmesh.core.config.common import (EnvironmentSuffixTarget, + TableNamingConvention, + VirtualEnvironmentMode, + compile_regex_mapping, + variables_validator) +from sqlmesh.core.config.connection import (ConnectionConfig, + DuckDBConnectionConfig, + SerializableConnectionConfig, + connection_config_validator) +from sqlmesh.core.config.dbt import DbtConfig from sqlmesh.core.config.format import FormatConfig from sqlmesh.core.config.gateway import GatewayConfig from sqlmesh.core.config.janitor import JanitorConfig +from sqlmesh.core.config.linter import LinterConfig as LinterConfig from sqlmesh.core.config.migration import MigrationConfig from sqlmesh.core.config.model import ModelDefaultsConfig -from sqlmesh.core.config.naming import NameInferenceConfig as NameInferenceConfig -from sqlmesh.core.config.linter import LinterConfig as LinterConfig +from sqlmesh.core.config.naming import \ + NameInferenceConfig as NameInferenceConfig from sqlmesh.core.config.plan import PlanConfig from sqlmesh.core.config.run import RunConfig -from sqlmesh.core.config.dbt import DbtConfig -from sqlmesh.core.config.scheduler import ( - BuiltInSchedulerConfig, - SchedulerConfig, - scheduler_config_validator, -) +from sqlmesh.core.config.scheduler import (BuiltInSchedulerConfig, + SchedulerConfig, + scheduler_config_validator) from sqlmesh.core.config.ui import UIConfig +from sqlmesh.core.console import get_console from sqlmesh.core.loader import Loader, SqlMeshLoader from sqlmesh.core.notification_target import NotificationTarget from sqlmesh.core.user import User -from sqlmesh.utils.date import to_timestamp, now +from sqlmesh.utils.date import now, to_timestamp from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.pydantic import model_validator @@ -72,7 +68,9 @@ def gateways_ensure_dict(value: t.Dict[str, t.Any]) -> t.Dict[str, t.Any]: return value -def validate_regex_key_dict(value: t.Dict[str | re.Pattern, t.Any]) -> t.Dict[re.Pattern, t.Any]: +def validate_regex_key_dict( + value: t.Dict[str | re.Pattern, t.Any], +) -> t.Dict[re.Pattern, t.Any]: return compile_regex_mapping(value) @@ -102,8 +100,12 @@ def _canonicalize(obj: object) -> object: RegexKeyDict = t.Dict[re.Pattern, str] else: NoPastTTLString = t.Annotated[str, BeforeValidator(validate_no_past_ttl)] - GatewayDict = t.Annotated[t.Dict[str, GatewayConfig], BeforeValidator(gateways_ensure_dict)] - RegexKeyDict = t.Annotated[t.Dict[re.Pattern, str], BeforeValidator(validate_regex_key_dict)] + GatewayDict = t.Annotated[ + t.Dict[str, GatewayConfig], BeforeValidator(gateways_ensure_dict) + ] + RegexKeyDict = t.Annotated[ + t.Dict[re.Pattern, str], BeforeValidator(validate_regex_key_dict) + ] class Config(BaseConfig): @@ -171,7 +173,9 @@ class Config(BaseConfig): username: str = "" physical_schema_mapping: RegexKeyDict = {} environment_suffix_target: EnvironmentSuffixTarget = EnvironmentSuffixTarget.default - physical_table_naming_convention: TableNamingConvention = TableNamingConvention.default + physical_table_naming_convention: TableNamingConvention = ( + TableNamingConvention.default + ) virtual_environment_mode: VirtualEnvironmentMode = VirtualEnvironmentMode.default gateway_managed_virtual_layer: bool = False infer_python_dependencies: bool = True @@ -242,7 +246,9 @@ def _normalize_and_validate_fields(cls, data: t.Any) -> t.Any: "Only one of `physical_schema_override` and `physical_schema_mapping` can be specified." ) - physical_schema_override: t.Dict[str, str] = data.pop("physical_schema_override") + physical_schema_override: t.Dict[str, str] = data.pop( + "physical_schema_override" + ) # translate physical_schema_override to physical_schema_mapping data["physical_schema_mapping"] = { f"^{k}$": v for k, v in physical_schema_override.items() @@ -290,7 +296,9 @@ def _inherit_project_config_in_cicd_bot(self) -> Self: if self.cicd_bot: # inherit the project-level settings into the CICD bot if they have not been explicitly overridden if self.cicd_bot.auto_categorize_changes_ is None: - self.cicd_bot.auto_categorize_changes_ = self.plan.auto_categorize_changes + self.cicd_bot.auto_categorize_changes_ = ( + self.plan.auto_categorize_changes + ) if self.cicd_bot.pr_include_unmodified_ is None: self.cicd_bot.pr_include_unmodified_ = self.plan.include_unmodified @@ -308,9 +316,9 @@ def get_default_test_connection( if default_catalog is None else { # transpile catalog name from main connection dialect to DuckDB - exp.parse_identifier(default_catalog, dialect=default_catalog_dialect).sql( - dialect="duckdb" - ): ":memory:" + exp.parse_identifier( + default_catalog, dialect=default_catalog_dialect + ).sql(dialect="duckdb"): ":memory:" } ) ) @@ -322,7 +330,9 @@ def get_gateway(self, name: t.Optional[str] = None) -> GatewayConfig: # Normalize default_gateway name to lowercase for lookup default_key = self.default_gateway.lower() if default_key not in self.gateways: - raise ConfigError(f"Missing gateway with name '{self.default_gateway}'") + raise ConfigError( + f"Missing gateway with name '{self.default_gateway}'" + ) return self.gateways[default_key] if "" in self.gateways: @@ -337,11 +347,15 @@ def get_gateway(self, name: t.Optional[str] = None) -> GatewayConfig: return self.gateways[lookup_key] if name is not None: - raise ConfigError("Gateway name is not supported when only one gateway is configured.") + raise ConfigError( + "Gateway name is not supported when only one gateway is configured." + ) return self.gateways def get_connection(self, gateway_name: t.Optional[str] = None) -> ConnectionConfig: - connection = self.get_gateway(gateway_name).connection or self.default_connection + connection = ( + self.get_gateway(gateway_name).connection or self.default_connection + ) if connection is None: msg = f" for gateway '{gateway_name}'" if gateway_name else "" raise ConfigError(f"No connection configured{msg}.") @@ -358,8 +372,11 @@ def get_test_connection( default_catalog: t.Optional[str] = None, default_catalog_dialect: t.Optional[str] = None, ) -> ConnectionConfig: - return self.get_gateway(gateway_name).test_connection or self.get_default_test_connection( - default_catalog=default_catalog, default_catalog_dialect=default_catalog_dialect + return self.get_gateway( + gateway_name + ).test_connection or self.get_default_test_connection( + default_catalog=default_catalog, + default_catalog_dialect=default_catalog_dialect, ) def get_scheduler(self, gateway_name: t.Optional[str] = None) -> SchedulerConfig: @@ -384,6 +401,8 @@ def dialect(self) -> t.Optional[str]: def fingerprint(self) -> str: return str( zlib.crc32( - pickle.dumps(_canonicalize(self.dict(exclude={"loader", "notification_targets"}))) + pickle.dumps( + _canonicalize(self.dict(exclude={"loader", "notification_targets"})) + ) ) ) diff --git a/sqlmesh/core/config/run.py b/sqlmesh/core/config/run.py index dc610c67f6..c2cbbbc80b 100644 --- a/sqlmesh/core/config/run.py +++ b/sqlmesh/core/config/run.py @@ -16,7 +16,9 @@ class RunConfig(BaseConfig): environment_check_interval: int = 30 environment_check_max_wait: int = 6 * 60 * 60 # 6 hours by default - @field_validator("environment_check_interval", "environment_check_max_wait", mode="after") + @field_validator( + "environment_check_interval", "environment_check_max_wait", mode="after" + ) @classmethod def _validate_positive_int(cls, v: int) -> int: if v <= 0: diff --git a/sqlmesh/core/config/scheduler.py b/sqlmesh/core/config/scheduler.py index 4cce9b0f76..bd63d0fcc9 100644 --- a/sqlmesh/core/config/scheduler.py +++ b/sqlmesh/core/config/scheduler.py @@ -4,15 +4,12 @@ import typing as t from pydantic import Field, ValidationError - from sqlglot.helper import subclasses + +from sqlmesh.core.config import DuckDBConnectionConfig from sqlmesh.core.config.base import BaseConfig from sqlmesh.core.console import get_console -from sqlmesh.core.plan import ( - BuiltInPlanEvaluator, - PlanEvaluator, -) -from sqlmesh.core.config import DuckDBConnectionConfig +from sqlmesh.core.plan import BuiltInPlanEvaluator, PlanEvaluator from sqlmesh.core.state_sync import EngineAdapterStateSync, StateSync from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.hashing import md5 @@ -21,7 +18,7 @@ if t.TYPE_CHECKING: from sqlmesh.core.context import GenericContext -from sqlmesh.utils.config import sensitive_fields, excluded_fields +from sqlmesh.utils.config import excluded_fields, sensitive_fields class SchedulerConfig(abc.ABC): @@ -47,7 +44,9 @@ def create_state_sync(self, context: GenericContext) -> StateSync: """ @abc.abstractmethod - def get_default_catalog_per_gateway(self, context: GenericContext) -> t.Dict[str, str]: + def get_default_catalog_per_gateway( + self, context: GenericContext + ) -> t.Dict[str, str]: """Returns the default catalog for each gateway. Args: @@ -66,7 +65,8 @@ def state_sync_fingerprint(self, context: GenericContext) -> str: class _EngineAdapterStateSyncSchedulerConfig(SchedulerConfig): def create_state_sync(self, context: GenericContext) -> StateSync: state_connection = ( - context.config.get_state_connection(context.gateway) or context.connection_config + context.config.get_state_connection(context.gateway) + or context.connection_config ) warehouse_connection = context.config.get_connection(context.gateway) @@ -83,7 +83,9 @@ def create_state_sync(self, context: GenericContext) -> StateSync: + " This can cause SQLMesh to hang. Overriding the duckdb state connection config to use multi threaded mode." ) # this triggers multithreaded mode and has to happen before the engine adapter is created below - state_connection.concurrent_tasks = warehouse_connection.concurrent_tasks + state_connection.concurrent_tasks = ( + warehouse_connection.concurrent_tasks + ) engine_adapter = state_connection.create_engine_adapter() if state_connection.is_forbidden_for_state_sync: @@ -105,12 +107,16 @@ def create_state_sync(self, context: GenericContext) -> StateSync: schema = context.config.get_state_schema(context.gateway) return EngineAdapterStateSync( - engine_adapter, schema=schema, cache_dir=context.cache_dir, console=context.console + engine_adapter, + schema=schema, + cache_dir=context.cache_dir, + console=context.console, ) def state_sync_fingerprint(self, context: GenericContext) -> str: state_connection = ( - context.config.get_state_connection(context.gateway) or context.connection_config + context.config.get_state_connection(context.gateway) + or context.connection_config ) return md5( [ @@ -136,7 +142,9 @@ def create_plan_evaluator(self, context: GenericContext) -> PlanEvaluator: console=context.console, ) - def get_default_catalog_per_gateway(self, context: GenericContext) -> t.Dict[str, str]: + def get_default_catalog_per_gateway( + self, context: GenericContext + ) -> t.Dict[str, str]: default_catalogs_per_gateway: t.Dict[str, str] = {} unsupported_gateways = [] diff --git a/sqlmesh/core/console.py b/sqlmesh/core/console.py index f9a758b405..8ad5df886b 100644 --- a/sqlmesh/core/console.py +++ b/sqlmesh/core/console.py @@ -2,25 +2,20 @@ import abc import datetime +import logging +import textwrap import typing as t import unittest import uuid -import logging -import textwrap -from humanize import metric, naturalsize from itertools import zip_longest from pathlib import Path + +from humanize import metric, naturalsize from hyperscript import h from rich.console import Console as RichConsole from rich.live import Live -from rich.progress import ( - BarColumn, - Progress, - SpinnerColumn, - TaskID, - TextColumn, - TimeElapsedColumn, -) +from rich.progress import (BarColumn, Progress, SpinnerColumn, TaskID, + TextColumn, TimeElapsedColumn) from rich.prompt import Confirm, Prompt from rich.status import Status from rich.syntax import Syntax @@ -28,42 +23,38 @@ from rich.tree import Tree from sqlglot import exp -from sqlmesh.core.schema_diff import TableAlterOperation -from sqlmesh.core.test.result import ModelTextTestResult from sqlmesh.core.environment import EnvironmentNamingInfo, EnvironmentSummary from sqlmesh.core.linter.rule import RuleViolation from sqlmesh.core.model import Model -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotChangeCategory, - SnapshotId, - SnapshotInfoLike, -) -from sqlmesh.core.snapshot.definition import Interval, Intervals, SnapshotTableInfo +from sqlmesh.core.schema_diff import TableAlterOperation +from sqlmesh.core.snapshot import (Snapshot, SnapshotChangeCategory, + SnapshotId, SnapshotInfoLike) +from sqlmesh.core.snapshot.definition import (Interval, Intervals, + SnapshotTableInfo) from sqlmesh.core.snapshot.execution_tracker import QueryExecutionStats from sqlmesh.core.test import ModelTest -from sqlmesh.utils import rich as srich +from sqlmesh.core.test.result import ModelTextTestResult from sqlmesh.utils import Verbosity +from sqlmesh.utils import rich as srich from sqlmesh.utils.concurrency import NodeExecutionFailedError -from sqlmesh.utils.date import time_like_to_str, to_date, yesterday_ds, to_ds, make_inclusive -from sqlmesh.utils.errors import ( - PythonModelEvalError, - NodeAuditsErrors, - format_destructive_change_msg, - format_additive_change_msg, -) +from sqlmesh.utils.date import (make_inclusive, time_like_to_str, to_date, + to_ds, yesterday_ds) +from sqlmesh.utils.errors import (NodeAuditsErrors, PythonModelEvalError, + format_additive_change_msg, + format_destructive_change_msg) from sqlmesh.utils.rich import strip_ansi_codes if t.TYPE_CHECKING: import ipywidgets as widgets - from sqlglot import exp from sqlglot.dialects.dialect import DialectType - from sqlmesh.core.context_diff import ContextDiff - from sqlmesh.core.plan import Plan, EvaluatablePlan, PlanBuilder, SnapshotIntervals - from sqlmesh.core.table_diff import TableDiff, RowDiff, SchemaDiff + from sqlmesh.core.config.connection import ConnectionConfig + from sqlmesh.core.context_diff import ContextDiff + from sqlmesh.core.plan import (EvaluatablePlan, Plan, PlanBuilder, + SnapshotIntervals) from sqlmesh.core.state_sync import Versions + from sqlmesh.core.table_diff import RowDiff, SchemaDiff, TableDiff LayoutWidget = t.TypeVar("LayoutWidget", bound=t.Union[widgets.VBox, widgets.HBox]) @@ -220,11 +211,15 @@ class EnvironmentsConsole(abc.ABC): """Console for displaying environments""" @abc.abstractmethod - def print_environments(self, environments_summary: t.List[EnvironmentSummary]) -> None: + def print_environments( + self, environments_summary: t.List[EnvironmentSummary] + ) -> None: """Prints all environment names along with expiry datetime.""" @abc.abstractmethod - def show_intervals(self, snapshot_intervals: t.Dict[Snapshot, SnapshotIntervals]) -> None: + def show_intervals( + self, snapshot_intervals: t.Dict[Snapshot, SnapshotIntervals] + ) -> None: """Show ready intervals""" @@ -296,7 +291,10 @@ def show_schema_diff(self, schema_diff: SchemaDiff) -> None: @abc.abstractmethod def show_row_diff( - self, row_diff: RowDiff, show_sample: bool = True, skip_grain_check: bool = False + self, + row_diff: RowDiff, + show_sample: bool = True, + skip_grain_check: bool = False, ) -> None: """Show table summary diff.""" @@ -350,7 +348,9 @@ def log_additive_change( class UnitTestConsole(abc.ABC): @abc.abstractmethod - def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> None: + def log_test_results( + self, result: ModelTextTestResult, target_dialect: str + ) -> None: """Display the test result and output. Args: @@ -477,7 +477,9 @@ def start_promotion_progress( """Indicates that a new snapshot promotion progress has begun.""" @abc.abstractmethod - def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) -> None: + def update_promotion_progress( + self, snapshot: SnapshotInfoLike, promoted: bool + ) -> None: """Update the snapshot promotion progress.""" @abc.abstractmethod @@ -668,7 +670,9 @@ def start_promotion_progress( ) -> None: pass - def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) -> None: + def update_promotion_progress( + self, snapshot: SnapshotInfoLike, promoted: bool + ) -> None: pass def stop_promotion_progress(self, success: bool = True) -> None: @@ -772,7 +776,9 @@ def plan( if auto_apply: plan_builder.apply() - def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> None: + def log_test_results( + self, result: ModelTextTestResult, target_dialect: str + ) -> None: pass def show_sql(self, sql: str) -> None: @@ -816,7 +822,9 @@ def log_additive_change( def log_error(self, message: str) -> None: pass - def log_warning(self, short_message: str, long_message: t.Optional[str] = None) -> None: + def log_warning( + self, short_message: str, long_message: t.Optional[str] = None + ) -> None: logger.warning(long_message or short_message) def log_success(self, message: str) -> None: @@ -839,7 +847,9 @@ def show_table_diff( self.show_table_diff_summary(table_diff) self.show_schema_diff(table_diff.schema_diff()) self.show_row_diff( - table_diff.row_diff(temp_schema=temp_schema, skip_grain_check=skip_grain_check), + table_diff.row_diff( + temp_schema=temp_schema, skip_grain_check=skip_grain_check + ), show_sample=show_sample, skip_grain_check=skip_grain_check, ) @@ -869,14 +879,21 @@ def show_schema_diff(self, schema_diff: SchemaDiff) -> None: pass def show_row_diff( - self, row_diff: RowDiff, show_sample: bool = True, skip_grain_check: bool = False + self, + row_diff: RowDiff, + show_sample: bool = True, + skip_grain_check: bool = False, ) -> None: pass - def print_environments(self, environments_summary: t.List[EnvironmentSummary]) -> None: + def print_environments( + self, environments_summary: t.List[EnvironmentSummary] + ) -> None: pass - def show_intervals(self, snapshot_intervals: t.Dict[Snapshot, SnapshotIntervals]) -> None: + def show_intervals( + self, snapshot_intervals: t.Dict[Snapshot, SnapshotIntervals] + ) -> None: pass def show_linter_violations( @@ -989,7 +1006,9 @@ def __init__( self.dialect = dialect self.ignore_warnings = ignore_warnings - def _limit_model_names(self, tree: Tree, verbosity: Verbosity = Verbosity.DEFAULT) -> Tree: + def _limit_model_names( + self, tree: Tree, verbosity: Verbosity = Verbosity.DEFAULT + ) -> Tree: """Trim long indirectly modified model lists below threshold.""" modified_length = len(tree.children) if ( @@ -1032,7 +1051,8 @@ def start_evaluation_progress( if not self.evaluation_progress_live: self.evaluation_total_progress = make_progress_bar( - "Executing model batches" if not audit_only else "Auditing models", self.console + "Executing model batches" if not audit_only else "Auditing models", + self.console, ) self.evaluation_model_progress = Progress( @@ -1051,7 +1071,8 @@ def start_evaluation_progress( self.evaluation_progress_live.start() batch_sizes = { - snapshot: len(intervals) for snapshot, intervals in batched_intervals.items() + snapshot: len(intervals) + for snapshot, intervals in batched_intervals.items() } message = "Executing" if not audit_only else "Auditing" self.evaluation_total_task = self.evaluation_total_progress.add_task( @@ -1061,7 +1082,9 @@ def start_evaluation_progress( # determine column widths self.evaluation_column_widths["annotation"] = ( _calculate_annotation_str_len( - batched_intervals, self.AUDIT_PADDING, len(" (123.4m rows, 123.4 KiB)") + batched_intervals, + self.AUDIT_PADDING, + len(" (123.4m rows, 123.4 KiB)"), ) + 3 # brackets and opening escape backslash ) @@ -1069,14 +1092,20 @@ def start_evaluation_progress( len( snapshot.display_name( environment_naming_info, - default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) ) for snapshot in batched_intervals ) largest_batch_size = max(batch_sizes.values()) - self.evaluation_column_widths["batch"] = len(str(largest_batch_size)) * 2 + 3 # [X/X] + self.evaluation_column_widths["batch"] = ( + len(str(largest_batch_size)) * 2 + 3 + ) # [X/X] self.evaluation_column_widths["duration"] = 8 self.evaluation_model_batch_sizes = batch_sizes @@ -1086,16 +1115,25 @@ def start_evaluation_progress( def start_snapshot_evaluation_progress( self, snapshot: Snapshot, audit_only: bool = False ) -> None: - if self.evaluation_model_progress and snapshot.name not in self.evaluation_model_tasks: + if ( + self.evaluation_model_progress + and snapshot.name not in self.evaluation_model_tasks + ): display_name = snapshot.display_name( self.environment_naming_info, - self.default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + self.default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) - self.evaluation_model_tasks[snapshot.name] = self.evaluation_model_progress.add_task( - f"{'Evaluating' if not audit_only else 'Auditing'} {display_name}...", - view_name=display_name, - total=self.evaluation_model_batch_sizes[snapshot], + self.evaluation_model_tasks[snapshot.name] = ( + self.evaluation_model_progress.add_task( + f"{'Evaluating' if not audit_only else 'Auditing'} {display_name}...", + view_name=display_name, + total=self.evaluation_model_batch_sizes[snapshot], + ) ) def update_snapshot_evaluation_progress( @@ -1118,17 +1156,25 @@ def update_snapshot_evaluation_progress( ): total_batches = self.evaluation_model_batch_sizes[snapshot] batch_num = str(batch_idx + 1).rjust(len(str(total_batches))) - batch = f"[{batch_num}/{total_batches}]".ljust(self.evaluation_column_widths["batch"]) + batch = f"[{batch_num}/{total_batches}]".ljust( + self.evaluation_column_widths["batch"] + ) if duration_ms: display_name = snapshot.display_name( self.environment_naming_info, - self.default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + self.default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ).ljust(self.evaluation_column_widths["name"]) annotation = _create_evaluation_model_annotation( - snapshot, _format_evaluation_model_interval(snapshot, interval), execution_stats + snapshot, + _format_evaluation_model_interval(snapshot, interval), + execution_stats, ) audits_str = "" if num_audits_passed: @@ -1159,9 +1205,12 @@ def update_snapshot_evaluation_progress( ) model_task_id = self.evaluation_model_tasks[snapshot.name] - self.evaluation_model_progress.update(model_task_id, refresh=True, advance=1) + self.evaluation_model_progress.update( + model_task_id, refresh=True, advance=1 + ) if ( - self.evaluation_model_progress._tasks[model_task_id].completed >= total_batches + self.evaluation_model_progress._tasks[model_task_id].completed + >= total_batches or audit_only ): self.evaluation_model_progress.remove_task(model_task_id) @@ -1210,8 +1259,12 @@ def update_signal_progress( """Updates the signal checking progress.""" tree = Tree(f"[{signal_idx + 1}/{total_signals}] {signal_name} {duration:.2f}s") - formatted_check_intervals = [_format_signal_interval(snapshot, i) for i in check_intervals] - formatted_ready_intervals = [_format_signal_interval(snapshot, i) for i in ready_intervals] + formatted_check_intervals = [ + _format_signal_interval(snapshot, i) for i in check_intervals + ] + formatted_ready_intervals = [ + _format_signal_interval(snapshot, i) for i in ready_intervals + ] if not formatted_check_intervals: formatted_check_intervals = ["no intervals"] @@ -1233,12 +1286,16 @@ def update_signal_progress( num_check_intervals = len(formatted_check_intervals) if num_check_intervals > 3: formatted_check_intervals = formatted_check_intervals[:3] - formatted_check_intervals.append(f"... and {num_check_intervals - 3} more") + formatted_check_intervals.append( + f"... and {num_check_intervals - 3} more" + ) num_ready_intervals = len(formatted_ready_intervals) if num_ready_intervals > 3: formatted_ready_intervals = formatted_ready_intervals[:3] - formatted_ready_intervals.append(f"... and {num_ready_intervals - 3} more") + formatted_ready_intervals.append( + f"... and {num_ready_intervals - 3} more" + ) check = ", ".join(formatted_check_intervals) tree.add(f"Check: {check}") @@ -1274,7 +1331,9 @@ def start_creation_progress( ) -> None: """Indicates that a new creation progress has begun.""" if self.creation_progress is None: - self.creation_progress = make_progress_bar("Updating physical layer", self.console) + self.creation_progress = make_progress_bar( + "Updating physical layer", self.console + ) self._print("") self.creation_progress.start() @@ -1289,7 +1348,11 @@ def start_creation_progress( len( snapshot.display_name( environment_naming_info, - default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) ) @@ -1305,10 +1368,16 @@ def update_creation_progress(self, snapshot: SnapshotInfoLike) -> None: if self.verbosity >= Verbosity.VERBOSE: msg = snapshot.display_name( self.environment_naming_info, - self.default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + self.default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ).ljust(self.creation_column_widths["name"]) - self.creation_progress.live.console.print(msg + " [green]created[/green]") + self.creation_progress.live.console.print( + msg + " [green]created[/green]" + ) self.creation_progress.update(self.creation_task, refresh=True, advance=1) def stop_creation_progress(self, success: bool = True) -> None: @@ -1382,7 +1451,9 @@ def start_destroy( "potentially external resources created by other tools in these schemas.\n" ) - if not self._confirm("Are you ABSOLUTELY SURE you want to proceed with deletion?"): + if not self._confirm( + "Are you ABSOLUTELY SURE you want to proceed with deletion?" + ): self.log_error("Destroy operation cancelled.") return False return True @@ -1420,7 +1491,11 @@ def start_promotion_progress( len( snapshot.display_name( environment_naming_info, - default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) ) @@ -1430,7 +1505,9 @@ def start_promotion_progress( self.environment_naming_info = environment_naming_info self.default_catalog = default_catalog - def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) -> None: + def update_promotion_progress( + self, snapshot: SnapshotInfoLike, promoted: bool + ) -> None: """Update the snapshot promotion progress.""" if ( self.promotion_progress is not None @@ -1441,7 +1518,11 @@ def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) if self.verbosity >= Verbosity.VERBOSE: display_name = snapshot.display_name( self.environment_naming_info, - self.default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + self.default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ).ljust(self.promotion_column_widths["name"]) action_str = "" @@ -1452,7 +1533,9 @@ def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) else "[green]created[/green]" ) action_str = action_str or "[red]dropped[/red]" - self.promotion_progress.live.console.print(f"{display_name} {action_str}") + self.promotion_progress.live.console.print( + f"{display_name} {action_str}" + ) self.promotion_progress.update(self.promotion_task, refresh=True, advance=1) def stop_promotion_progress(self, success: bool = True) -> None: @@ -1471,7 +1554,9 @@ def stop_promotion_progress(self, success: bool = True) -> None: def start_snapshot_migration_progress(self, total_tasks: int) -> None: """Indicates that a new snapshot migration progress has begun.""" if self.migration_progress is None: - self.migration_progress = make_progress_bar("Migrating snapshots", self.console) + self.migration_progress = make_progress_bar( + "Migrating snapshots", self.console + ) self.migration_progress.start() self.migration_task = self.migration_progress.add_task( @@ -1482,7 +1567,9 @@ def start_snapshot_migration_progress(self, total_tasks: int) -> None: def update_snapshot_migration_progress(self, num_tasks: int) -> None: """Update the migration progress.""" if self.migration_progress is not None and self.migration_task is not None: - self.migration_progress.update(self.migration_task, refresh=True, advance=num_tasks) + self.migration_progress.update( + self.migration_task, refresh=True, advance=num_tasks + ) def log_migration_status(self, success: bool = True) -> None: """Log the migration status.""" @@ -1502,7 +1589,9 @@ def stop_snapshot_migration_progress(self, success: bool = True) -> None: def start_env_migration_progress(self, total_tasks: int) -> None: """Indicates that a new environment migration has begun.""" if self.env_migration_progress is None: - self.env_migration_progress = make_progress_bar("Migrating environments", self.console) + self.env_migration_progress = make_progress_bar( + "Migrating environments", self.console + ) self.env_migration_progress.start() self.env_migration_task = self.env_migration_progress.add_task( "Migrating environments...", @@ -1511,7 +1600,10 @@ def start_env_migration_progress(self, total_tasks: int) -> None: def update_env_migration_progress(self, num_tasks: int) -> None: """Update the environment migration progress.""" - if self.env_migration_progress is not None and self.env_migration_task is not None: + if ( + self.env_migration_progress is not None + and self.env_migration_task is not None + ): self.env_migration_progress.update( self.env_migration_task, refresh=True, advance=num_tasks ) @@ -1537,7 +1629,9 @@ def start_state_export( self.state_export_progress = None if local_only: - self.log_status_update(f"Exporting [b]local[/b] state to '{output_file.as_posix()}'\n") + self.log_status_update( + f"Exporting [b]local[/b] state to '{output_file.as_posix()}'\n" + ) self.log_warning( "Local state exports just contain the model versions in your local context. Therefore, the resulting file cannot be imported." ) @@ -1548,9 +1642,13 @@ def start_state_export( if gateway: self.log_status_update(f"[b]Gateway[/b]: [green]{gateway}[/green]") if state_connection_config: - self.print_connection_config(state_connection_config, title="State Connection") + self.print_connection_config( + state_connection_config, title="State Connection" + ) if environment_names: - heading = "Environments" if len(environment_names) > 1 else "Environment" + heading = ( + "Environments" if len(environment_names) > 1 else "Environment" + ) self.log_status_update( f"[b]{heading}[/b]: [yellow]{', '.join(environment_names)}[/yellow]" ) @@ -1561,7 +1659,9 @@ def start_state_export( self.log_status_update("") if should_continue: - self.state_export_progress = make_progress_bar("{task.description}", self.console) + self.state_export_progress = make_progress_bar( + "{task.description}", self.console + ) assert isinstance(self.state_export_progress, Progress) self.state_export_version_task = self.state_export_progress.add_task( @@ -1590,7 +1690,9 @@ def update_state_export_progress( if self.state_export_progress: if self.state_export_version_task is not None: if version_count is not None: - self.state_export_progress.start_task(self.state_export_version_task) + self.state_export_progress.start_task( + self.state_export_version_task + ) self.state_export_progress.update( self.state_export_version_task, total=version_count, @@ -1602,7 +1704,9 @@ def update_state_export_progress( if self.state_export_snapshot_task is not None: if snapshot_count is not None: - self.state_export_progress.start_task(self.state_export_snapshot_task) + self.state_export_progress.start_task( + self.state_export_snapshot_task + ) self.state_export_progress.update( self.state_export_snapshot_task, total=snapshot_count, @@ -1610,11 +1714,15 @@ def update_state_export_progress( refresh=True, ) if snapshots_complete: - self.state_export_progress.stop_task(self.state_export_snapshot_task) + self.state_export_progress.stop_task( + self.state_export_snapshot_task + ) if self.state_export_environment_task is not None: if environment_count is not None: - self.state_export_progress.start_task(self.state_export_environment_task) + self.state_export_progress.start_task( + self.state_export_environment_task + ) self.state_export_progress.update( self.state_export_environment_task, total=environment_count, @@ -1622,7 +1730,9 @@ def update_state_export_progress( refresh=True, ) if environments_complete: - self.state_export_progress.stop_task(self.state_export_environment_task) + self.state_export_progress.stop_task( + self.state_export_environment_task + ) def stop_state_export(self, success: bool, output_file: Path) -> None: if self.state_export_progress: @@ -1632,7 +1742,9 @@ def stop_state_export(self, success: bool, output_file: Path) -> None: self.log_status_update("") if success: - self.log_success(f"State exported successfully to '{output_file.as_posix()}'") + self.log_success( + f"State exported successfully to '{output_file.as_posix()}'" + ) else: self.log_error("State export failed!") @@ -1669,7 +1781,9 @@ def start_state_import( self.log_status_update("") if should_continue: - self.state_import_progress = make_progress_bar("{task.description}", self.console) + self.state_import_progress = make_progress_bar( + "{task.description}", self.console + ) self.state_import_info = Tree("[bold]State File Information:") @@ -1704,28 +1818,38 @@ def update_state_import_progress( if state_file_version: self.state_import_info.add(f"File Version: {state_file_version}") if versions: - self.state_import_info.add(f"SQLMesh version: {versions.sqlmesh_version}") + self.state_import_info.add( + f"SQLMesh version: {versions.sqlmesh_version}" + ) self.state_import_info.add( f"SQLMesh migration version: {versions.schema_version}" ) - self.state_import_info.add(f"SQLGlot version: {versions.sqlglot_version}\n") + self.state_import_info.add( + f"SQLGlot version: {versions.sqlglot_version}\n" + ) self._print(self.state_import_info) version_count = len(versions.model_dump()) if self.state_import_version_task is not None: - self.state_import_progress.start_task(self.state_import_version_task) + self.state_import_progress.start_task( + self.state_import_version_task + ) self.state_import_progress.update( self.state_import_version_task, total=version_count, completed=version_count, ) - self.state_import_progress.stop_task(self.state_import_version_task) + self.state_import_progress.stop_task( + self.state_import_version_task + ) if self.state_import_snapshot_task is not None: if snapshot_count is not None: - self.state_import_progress.start_task(self.state_import_snapshot_task) + self.state_import_progress.start_task( + self.state_import_snapshot_task + ) self.state_import_progress.update( self.state_import_snapshot_task, completed=snapshot_count, @@ -1734,11 +1858,15 @@ def update_state_import_progress( ) if snapshots_complete: - self.state_import_progress.stop_task(self.state_import_snapshot_task) + self.state_import_progress.stop_task( + self.state_import_snapshot_task + ) if self.state_import_environment_task is not None: if environment_count is not None: - self.state_import_progress.start_task(self.state_import_environment_task) + self.state_import_progress.start_task( + self.state_import_environment_task + ) self.state_import_progress.update( self.state_import_environment_task, completed=environment_count, @@ -1747,7 +1875,9 @@ def update_state_import_progress( ) if environments_complete: - self.state_import_progress.stop_task(self.state_import_environment_task) + self.state_import_progress.stop_task( + self.state_import_environment_task + ) def stop_state_import(self, success: bool, input_file: Path) -> None: if self.state_import_progress: @@ -1757,7 +1887,9 @@ def stop_state_import(self, success: bool, input_file: Path) -> None: self.log_status_update("") if success: - self.log_success(f"State imported successfully from '{input_file.as_posix()}'") + self.log_success( + f"State imported successfully from '{input_file.as_posix()}'" + ) else: self.log_error("State import failed!") @@ -1875,7 +2007,10 @@ def plan( ) self._show_options_after_categorization( - plan_builder, auto_apply, default_catalog=default_catalog, no_prompts=no_prompts + plan_builder, + auto_apply, + default_catalog=default_catalog, + no_prompts=no_prompts, ) if auto_apply: @@ -1891,7 +2026,9 @@ def _show_summary_tree_for( no_diff: bool = True, ) -> None: added_snapshot_ids = { - s_id for s_id in context_diff.added if snapshot_selector(context_diff.snapshots[s_id]) + s_id + for s_id in context_diff.added + if snapshot_selector(context_diff.snapshots[s_id]) } removed_snapshot_ids = { s_id @@ -1964,7 +2101,9 @@ def _add_modified_models( direct.add( f"[direct]{display_name}" if no_diff - else Syntax(f"{display_name}\n{context_diff.text_diff(name)}", "sql") + else Syntax( + f"{display_name}\n{context_diff.text_diff(name)}", "sql" + ) ) elif context_diff.indirectly_modified(name): indirect.add(f"[indirect]{display_name}") @@ -1972,7 +2111,9 @@ def _add_modified_models( metadata.add( f"[metadata]{display_name}" if no_diff - else Syntax(f"{display_name}\n{context_diff.text_diff(name)}", "sql") + else Syntax( + f"{display_name}\n{context_diff.text_diff(name)}", "sql" + ) ) if direct.children: tree.add(direct) @@ -1999,7 +2140,9 @@ def _show_options_after_categorization( if not no_prompts: self._prompt_backfill(plan_builder, auto_apply, default_catalog) - backfill_or_preview = "preview" if plan.is_dev and plan.forward_only else "backfill" + backfill_or_preview = ( + "preview" if plan.is_dev and plan.forward_only else "backfill" + ) if not auto_apply and self._confirm( f"Apply - {backfill_or_preview.capitalize()} Tables" ): @@ -2039,7 +2182,11 @@ def _prompt_categorize( snapshot = model_fqn_to_snapshot[model_fqn] display_name = snapshot.display_name( plan.environment_naming_info, - default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) tree.add( @@ -2076,7 +2223,9 @@ def _prompt_categorize( ) indirect_tree = None - for child_sid in sorted(plan.indirectly_modified.get(snapshot.snapshot_id, set())): + for child_sid in sorted( + plan.indirectly_modified.get(snapshot.snapshot_id, set()) + ): child_snapshot = plan.context_diff.snapshots[child_sid] if not indirect_tree: indirect_tree = Tree("[indirect]Indirectly Modified Children:") @@ -2093,7 +2242,9 @@ def _prompt_categorize( snapshot, plan_builder, auto_apply, default_catalog ) - def _show_categorized_snapshots(self, plan: Plan, default_catalog: t.Optional[str]) -> None: + def _show_categorized_snapshots( + self, plan: Plan, default_catalog: t.Optional[str] + ) -> None: context_diff = plan.context_diff for snapshot in plan.categorized: @@ -2103,7 +2254,9 @@ def _show_categorized_snapshots(self, plan: Plan, default_catalog: t.Optional[st f"\n[bold][direct]Directly Modified: {snapshot.display_name(plan.environment_naming_info, default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, dialect=self.dialect)} ({category_str})" ) indirect_tree = None - for child_sid in sorted(plan.indirectly_modified.get(snapshot.snapshot_id, set())): + for child_sid in sorted( + plan.indirectly_modified.get(snapshot.snapshot_id, set()) + ): child_snapshot = context_diff.snapshots[child_sid] if not indirect_tree: indirect_tree = Tree("[indirect]Indirectly Modified Children:") @@ -2115,7 +2268,9 @@ def _show_categorized_snapshots(self, plan: Plan, default_catalog: t.Optional[st f"[indirect]{child_snapshot.display_name(plan.environment_naming_info, default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, dialect=self.dialect)} ({child_category_str})" ) if indirect_tree: - indirect_tree = self._limit_model_names(indirect_tree, self.verbosity) + indirect_tree = self._limit_model_names( + indirect_tree, self.verbosity + ) elif context_diff.metadata_updated(snapshot.name): tree = Tree( f"\n[bold][metadata]Metadata Updated: {snapshot.display_name(plan.environment_naming_info, default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, dialect=self.dialect)}" @@ -2158,7 +2313,10 @@ def _show_missing_dates(self, plan: Plan, default_catalog: t.Optional[str]) -> N self._print(backfill) def _prompt_effective_from( - self, plan_builder: PlanBuilder, auto_apply: bool, default_catalog: t.Optional[str] + self, + plan_builder: PlanBuilder, + auto_apply: bool, + default_catalog: t.Optional[str], ) -> None: if not plan_builder.build().effective_from: effective_from = self._prompt( @@ -2168,7 +2326,10 @@ def _prompt_effective_from( plan_builder.set_effective_from(effective_from) def _prompt_backfill( - self, plan_builder: PlanBuilder, auto_apply: bool, default_catalog: t.Optional[str] + self, + plan_builder: PlanBuilder, + auto_apply: bool, + default_catalog: t.Optional[str], ) -> None: plan = plan_builder.build() is_forward_only_dev = plan.is_dev and plan.forward_only @@ -2185,7 +2346,9 @@ def _prompt_backfill( default_start = yesterday_ds() else: if plan.provided_start: - blank_meaning = f"starting from '{time_like_to_str(plan.provided_start)}'" + blank_meaning = ( + f"starting from '{time_like_to_str(plan.provided_start)}'" + ) else: blank_meaning = "from the beginning of history" default_start = None @@ -2220,7 +2383,9 @@ def _prompt_promote(self, plan_builder: PlanBuilder) -> None: ): plan_builder.apply() - def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> None: + def log_test_results( + self, result: ModelTextTestResult, target_dialect: str + ) -> None: # We don't log the test results if no tests were ran if not result.testsRun: return @@ -2229,9 +2394,7 @@ def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> self._log_test_details(result) - message = ( - f"Ran {result.testsRun} tests against {target_dialect} in {result.duration} seconds." - ) + message = f"Ran {result.testsRun} tests against {target_dialect} in {result.duration} seconds." if result.wasSuccessful(): self._print("=" * divider_length) self._print( @@ -2290,12 +2453,18 @@ def log_models_updated_during_restatement( for restated_snapshot, updated_snapshot in snapshots: display_name = restated_snapshot.display_name( environment_naming_info, - default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) current_branch = tree.add(display_name) current_branch.add(f"restated version: '{restated_snapshot.version}'") - current_branch.add(f"currently active version: '{updated_snapshot.version}'") + current_branch.add( + f"currently active version: '{updated_snapshot.version}'" + ) self._print(tree) self._print("") # newline spacer @@ -2308,10 +2477,14 @@ def log_destructive_change( error: bool = True, ) -> None: if error: - self._print(format_destructive_change_msg(snapshot_name, alter_operations, dialect)) + self._print( + format_destructive_change_msg(snapshot_name, alter_operations, dialect) + ) else: self.log_warning( - format_destructive_change_msg(snapshot_name, alter_operations, dialect, error) + format_destructive_change_msg( + snapshot_name, alter_operations, dialect, error + ) ) def log_additive_change( @@ -2322,16 +2495,22 @@ def log_additive_change( error: bool = True, ) -> None: if error: - self._print(format_additive_change_msg(snapshot_name, alter_operations, dialect)) + self._print( + format_additive_change_msg(snapshot_name, alter_operations, dialect) + ) else: self.log_warning( - format_additive_change_msg(snapshot_name, alter_operations, dialect, error) + format_additive_change_msg( + snapshot_name, alter_operations, dialect, error + ) ) def log_error(self, message: str) -> None: self._print(f"[red]{message}[/red]") - def log_warning(self, short_message: str, long_message: t.Optional[str] = None) -> None: + def log_warning( + self, short_message: str, long_message: t.Optional[str] = None + ) -> None: logger.warning(long_message or short_message) if not self.ignore_warnings: if long_message: @@ -2340,11 +2519,15 @@ def log_warning(self, short_message: str, long_message: t.Optional[str] = None) if isinstance(handler, logging.FileHandler): file_path = handler.baseFilename break - file_path_msg = f" Learn more in logs: {file_path}\n" if file_path else "" + file_path_msg = ( + f" Learn more in logs: {file_path}\n" if file_path else "" + ) short_message = f"{short_message}{file_path_msg}" message_lstrip = short_message.lstrip() leading_ws = short_message[: -len(message_lstrip)] - message_formatted = f"{leading_ws}[yellow]\\[WARNING] {message_lstrip}[/yellow]" + message_formatted = ( + f"{leading_ws}[yellow]\\[WARNING] {message_lstrip}[/yellow]" + ) self._print(message_formatted) def log_success(self, message: str) -> None: @@ -2352,7 +2535,9 @@ def log_success(self, message: str) -> None: def loading_start(self, message: t.Optional[str] = None) -> uuid.UUID: id = uuid.uuid4() - self.loading_status[id] = Status(message or "", console=self.console, spinner="line") + self.loading_status[id] = Status( + message or "", console=self.console, spinner="line" + ) self.loading_status[id].start() return id @@ -2369,7 +2554,9 @@ def show_table_diff_details( if models_to_diff: m_tree = Tree("\n[b]Models to compare:") for m in models_to_diff: - m_tree.add(f"[{self.TABLE_DIFF_SOURCE_BLUE}]{m}[/{self.TABLE_DIFF_SOURCE_BLUE}]") + m_tree.add( + f"[{self.TABLE_DIFF_SOURCE_BLUE}]{m}[/{self.TABLE_DIFF_SOURCE_BLUE}]" + ) self._print(m_tree) self._print("") @@ -2397,15 +2584,19 @@ def start_table_diff_progress(self, models_to_diff: int) -> None: def start_table_diff_model_progress(self, model: str) -> None: if self.table_diff_model_progress and model not in self.table_diff_model_tasks: - self.table_diff_model_tasks[model] = self.table_diff_model_progress.add_task( - f"Diffing {model}...", - view_name=model, - total=1, + self.table_diff_model_tasks[model] = ( + self.table_diff_model_progress.add_task( + f"Diffing {model}...", + view_name=model, + total=1, + ) ) def update_table_diff_progress(self, model: str) -> None: if self.table_diff_progress: - self.table_diff_progress.update(self.table_diff_model_task, refresh=True, advance=1) + self.table_diff_progress.update( + self.table_diff_model_task, refresh=True, advance=1 + ) if self.table_diff_model_progress and model in self.table_diff_model_tasks: model_task_id = self.table_diff_model_tasks[model] self.table_diff_model_progress.remove_task(model_task_id) @@ -2478,7 +2669,8 @@ def show_schema_diff(self, schema_diff: SchemaDiff) -> None: first_line = f"\n[b]Schema Diff Between '[{self.TABLE_DIFF_SOURCE_BLUE}]{source_name}[/{self.TABLE_DIFF_SOURCE_BLUE}]' and '[{self.TABLE_DIFF_TARGET_GREEN}]{target_name}[/{self.TABLE_DIFF_TARGET_GREEN}]'" if schema_diff.model_name: first_line = ( - first_line + f" environments for model '[blue]{schema_diff.model_name}[/blue]'" + first_line + + f" environments for model '[blue]{schema_diff.model_name}[/blue]'" ) tree = Tree(first_line + ":") @@ -2507,7 +2699,10 @@ def show_schema_diff(self, schema_diff: SchemaDiff) -> None: self.console.print(tree) def show_row_diff( - self, row_diff: RowDiff, show_sample: bool = True, skip_grain_check: bool = False + self, + row_diff: RowDiff, + show_sample: bool = True, + skip_grain_check: bool = False, ) -> None: if row_diff.empty: self.console.print( @@ -2556,11 +2751,15 @@ def show_row_diff( if row_diff.column_stats.shape[0] > 0: self.console.print(row_diff.column_stats.to_string(index=True), end="\n\n") else: - self.console.print(" No columns with same name and data type in both tables") + self.console.print( + " No columns with same name and data type in both tables" + ) if show_sample: sample = row_diff.joined_sample - self.console.print("\n[b][blue]COMMON ROWS[/blue] sample data differences:[/b]") + self.console.print( + "\n[b][blue]COMMON ROWS[/blue] sample data differences:[/b]" + ) if sample.shape[0] > 0: keys: list[str] = [] columns: dict[str, list[str]] = {} @@ -2590,12 +2789,16 @@ def show_row_diff( for column, [source_column, target_column] in columns.items(): # Create a table with the joined keys and comparison columns - column_table = row_diff.joined_sample[keys + [source_column, target_column]] + column_table = row_diff.joined_sample[ + keys + [source_column, target_column] + ] # Filter to retain non identical-valued rows column_table = column_table[ column_table.apply( - lambda row: not _cells_match(row[source_column], row[target_column]), + lambda row: not _cells_match( + row[source_column], row[target_column] + ), axis=1, ) ] @@ -2635,11 +2838,15 @@ def show_row_diff( self.console.print(" All joined rows match") if row_diff.s_sample.shape[0] > 0: - self.console.print(f"\n[b][yellow]{source_name} ONLY[/yellow] sample rows:[/b]") + self.console.print( + f"\n[b][yellow]{source_name} ONLY[/yellow] sample rows:[/b]" + ) self.console.print(row_diff.s_sample.to_string(index=False), end="\n\n") if row_diff.t_sample.shape[0] > 0: - self.console.print(f"\n[b][green]{target_name} ONLY[/green] sample rows:[/b]") + self.console.print( + f"\n[b][green]{target_name} ONLY[/green] sample rows:[/b]" + ) self.console.print(row_diff.t_sample.to_string(index=False), end="\n\n") def show_table_diff( @@ -2657,7 +2864,8 @@ def show_table_diff( fully_matched = [] for table_diff in table_diffs: if ( - table_diff.schema_diff().source_schema == table_diff.schema_diff().target_schema + table_diff.schema_diff().source_schema + == table_diff.schema_diff().target_schema ) and ( table_diff.row_diff( temp_schema=temp_schema, skip_grain_check=skip_grain_check @@ -2688,27 +2896,37 @@ def show_table_diff( self.show_table_diff_summary(table_diff) self.show_schema_diff(table_diff.schema_diff()) self.show_row_diff( - table_diff.row_diff(temp_schema=temp_schema, skip_grain_check=skip_grain_check), + table_diff.row_diff( + temp_schema=temp_schema, skip_grain_check=skip_grain_check + ), show_sample=show_sample, skip_grain_check=skip_grain_check, ) - def print_environments(self, environments_summary: t.List[EnvironmentSummary]) -> None: + def print_environments( + self, environments_summary: t.List[EnvironmentSummary] + ) -> None: """Prints all environment names along with expiry datetime.""" output = [ - f"{summary.name} - {time_like_to_str(summary.expiration_ts)}" - if summary.expiration_ts - else f"{summary.name} - No Expiry" + ( + f"{summary.name} - {time_like_to_str(summary.expiration_ts)}" + if summary.expiration_ts + else f"{summary.name} - No Expiry" + ) for summary in environments_summary ] output_str = "\n".join([str(len(output)), *output]) self.log_status_update(f"Number of SQLMesh environments are: {output_str}") - def show_intervals(self, snapshot_intervals: t.Dict[Snapshot, SnapshotIntervals]) -> None: + def show_intervals( + self, snapshot_intervals: t.Dict[Snapshot, SnapshotIntervals] + ) -> None: complete = Tree(f"[b]Complete Intervals[/b]") incomplete = Tree(f"[b]Missing Intervals[/b]") - for snapshot, intervals in sorted(snapshot_intervals.items(), key=lambda s: s[0].node.name): + for snapshot, intervals in sorted( + snapshot_intervals.items(), key=lambda s: s[0].node.name + ): if intervals.intervals: incomplete.add( f"{snapshot.node.name}: [{intervals.format_intervals(snapshot.node.interval_unit)}]" @@ -2722,7 +2940,9 @@ def show_intervals(self, snapshot_intervals: t.Dict[Snapshot, SnapshotIntervals] if incomplete.children: self._print(incomplete) - def print_connection_config(self, config: ConnectionConfig, title: str = "Connection") -> None: + def print_connection_config( + self, config: ConnectionConfig, title: str = "Connection" + ) -> None: tree = Tree(f"[b]{title}:[/b]") tree.add(f"Type: [bold cyan]{config.type_}[/bold cyan]") tree.add(f"Catalog: [bold cyan]{config.get_catalog()}[/bold cyan]") @@ -2747,7 +2967,9 @@ def _get_snapshot_change_category( snapshot, plan_builder.environment_naming_info, default_catalog ) response = self._prompt( - "\n".join([f"[{i + 1}] {choice}" for i, choice in enumerate(choices.values())]), + "\n".join( + [f"[{i + 1}] {choice}" for i, choice in enumerate(choices.values())] + ), show_choices=False, choices=[f"{i + 1}" for i in range(len(choices))], ) @@ -2864,8 +3086,8 @@ def _log_test_details( def _cells_match(x: t.Any, y: t.Any) -> bool: """Helper function to compare two cells and returns true if they're equal, handling array objects.""" - import pandas as pd import numpy as np + import pandas as pd # Convert array-like objects to list for consistent comparison def _normalize(val: t.Any) -> t.Any: @@ -2886,7 +3108,9 @@ def _normalize(val: t.Any) -> t.Any: return _normalize(x) == _normalize(y) -def add_to_layout_widget(target_widget: LayoutWidget, *widgets: widgets.Widget) -> LayoutWidget: +def add_to_layout_widget( + target_widget: LayoutWidget, *widgets: widgets.Widget +) -> LayoutWidget: """Helper function to add a widget to a layout widget. Args: @@ -2922,7 +3146,9 @@ def __init__( ipython = get_ipython() self.display = display or ( - ipython.user_ns.get("display", ipython_display) if ipython else ipython_display + ipython.user_ns.get("display", ipython_display) + if ipython + else ipython_display ) self.missing_dates_output = widgets.Output() self.dynamic_options_after_categorization_output = widgets.VBox() @@ -2960,13 +3186,18 @@ def _prompt_promote(self, plan_builder: PlanBuilder) -> None: button.output = output def _prompt_effective_from( - self, plan_builder: PlanBuilder, auto_apply: bool, default_catalog: t.Optional[str] + self, + plan_builder: PlanBuilder, + auto_apply: bool, + default_catalog: t.Optional[str], ) -> None: import ipywidgets as widgets prompt = widgets.VBox() - def effective_from_change_callback(change: t.Dict[str, datetime.datetime]) -> None: + def effective_from_change_callback( + change: t.Dict[str, datetime.datetime], + ) -> None: plan_builder.set_effective_from(change["new"]) self._show_options_after_categorization( plan_builder, auto_apply, default_catalog, no_prompts=False @@ -3011,7 +3242,10 @@ def going_forward_change_callback(change: t.Dict[str, bool]) -> None: self._add_to_dynamic_options(prompt) def _prompt_backfill( - self, plan_builder: PlanBuilder, auto_apply: bool, default_catalog: t.Optional[str] + self, + plan_builder: PlanBuilder, + auto_apply: bool, + default_catalog: t.Optional[str], ) -> None: import ipywidgets as widgets @@ -3024,7 +3258,10 @@ def _prompt_backfill( ) def _date_picker( - plan_builder: PlanBuilder, value: t.Any, on_change: t.Callable, disabled: bool = False + plan_builder: PlanBuilder, + value: t.Any, + on_change: t.Callable, + disabled: bool = False, ) -> widgets.DatePicker: picker = widgets.DatePicker( disabled=disabled, @@ -3053,10 +3290,13 @@ def end_change_callback(change: t.Dict[str, datetime.datetime]) -> None: widgets.HBox( [ widgets.Label( - f"Start {backfill_or_preview} Date:", layout={"width": "8rem"} + f"Start {backfill_or_preview} Date:", + layout={"width": "8rem"}, ), _date_picker( - plan_builder, to_date(plan_builder.build().start), start_change_callback + plan_builder, + to_date(plan_builder.build().start), + start_change_callback, ), ] ), @@ -3066,7 +3306,9 @@ def end_change_callback(change: t.Dict[str, datetime.datetime]) -> None: prompt, widgets.HBox( [ - widgets.Label(f"End {backfill_or_preview} Date:", layout={"width": "8rem"}), + widgets.Label( + f"End {backfill_or_preview} Date:", layout={"width": "8rem"} + ), _date_picker( plan_builder, to_date(plan_builder.build().end), @@ -3143,7 +3385,9 @@ def radio_button_selected(change: t.Dict[str, t.Any]) -> None: ) self.display(radio) - def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> None: + def log_test_results( + self, result: ModelTextTestResult, target_dialect: str + ) -> None: # We don't log the test results if no tests were ran if not result.testsRun: return @@ -3157,9 +3401,7 @@ def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> "font-family": "Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace", } - message = ( - f"Ran {result.testsRun} tests against {target_dialect} in {result.duration} seconds." - ) + message = f"Ran {result.testsRun} tests against {target_dialect} in {result.duration} seconds." if result.wasSuccessful(): success_color = {"color": "#008000"} @@ -3179,7 +3421,9 @@ def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> fail_color = {"color": "#db3737"} fail_shared_style = {**shared_style, **fail_color} header = str(h("span", {"style": fail_shared_style}, "-" * divider_length)) - message = str(h("span", {"style": fail_shared_style}, "Test Failure Summary")) + message = str( + h("span", {"style": fail_shared_style}, "Test Failure Summary") + ) fail_and_error_tests = result.get_fail_and_error_tests() failed_tests = [ str( @@ -3203,9 +3447,17 @@ def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> ) failures = "
".join(failed_tests) footer = str(h("span", {"style": fail_shared_style}, "=" * divider_length)) - error_output = widgets.Textarea(output, layout={"height": "300px", "width": "100%"}) - test_info = widgets.HTML("
".join([header, message, footer, failures, footer])) - self.display(widgets.VBox(children=[test_info, error_output], layout={"width": "100%"})) + error_output = widgets.Textarea( + output, layout={"height": "300px", "width": "100%"} + ) + test_info = widgets.HTML( + "
".join([header, message, footer, failures, footer]) + ) + self.display( + widgets.VBox( + children=[test_info, error_output], layout={"width": "100%"} + ) + ) class CaptureTerminalConsole(TerminalConsole): @@ -3217,7 +3469,9 @@ class CaptureTerminalConsole(TerminalConsole): this console interactively. """ - def __init__(self, console: t.Optional[RichConsole] = None, **kwargs: t.Any) -> None: + def __init__( + self, console: t.Optional[RichConsole] = None, **kwargs: t.Any + ) -> None: super().__init__(console=console, **kwargs) self._captured_outputs: t.List[str] = [] self._warnings: t.List[str] = [] @@ -3274,7 +3528,9 @@ def log_error(self, message: str, *args: t.Any, **kwargs: t.Any) -> None: def log_skipped_models(self, snapshot_names: t.Set[str]) -> None: if snapshot_names: self._captured_outputs.append( - "\n".join([f"SKIPPED snapshot {skipped}\n" for skipped in snapshot_names]) + "\n".join( + [f"SKIPPED snapshot {skipped}\n" for skipped in snapshot_names] + ) ) super().log_skipped_models(snapshot_names) @@ -3301,7 +3557,9 @@ class MarkdownConsole(CaptureTerminalConsole): AUDIT_PADDING = 7 def __init__(self, **kwargs: t.Any) -> None: - self.alert_block_max_content_length = int(kwargs.pop("alert_block_max_content_length", 500)) + self.alert_block_max_content_length = int( + kwargs.pop("alert_block_max_content_length", 500) + ) self.alert_block_collapsible_threshold = int( kwargs.pop("alert_block_collapsible_threshold", 200) ) @@ -3312,7 +3570,10 @@ def __init__(self, **kwargs: t.Any) -> None: self.error_capture_only = kwargs.pop("error_capture_only", False) super().__init__( - **{**kwargs, "console": RichConsole(no_color=True, width=kwargs.pop("width", None))} + **{ + **kwargs, + "console": RichConsole(no_color=True, width=kwargs.pop("width", None)), + } ) def show_environment_difference_summary( @@ -3373,7 +3634,9 @@ def show_model_difference_summary( if added_snapshots: self._print("\n**Added Models:**") self._print_models_with_threshold( - environment_naming_info, {s for s in added_snapshots if s.is_model}, default_catalog + environment_naming_info, + {s for s in added_snapshots if s.is_model}, + default_catalog, ) added_snapshot_audits = {s for s in added_snapshots if s.is_audit} @@ -3393,7 +3656,9 @@ def show_model_difference_summary( default_catalog, ) - removed_audit_snapshot_table_infos = {s for s in removed_snapshot_table_infos if s.is_audit} + removed_audit_snapshot_table_infos = { + s for s in removed_snapshot_table_infos if s.is_audit + } if removed_audit_snapshot_table_infos: self._print("\n**Removed Standalone Audits:**") for snapshot_table_info in sorted(removed_audit_snapshot_table_infos): @@ -3402,11 +3667,16 @@ def show_model_difference_summary( ) modified_snapshots = { - current_snapshot for current_snapshot, _ in context_diff.modified_snapshots.values() + current_snapshot + for current_snapshot, _ in context_diff.modified_snapshots.values() } if modified_snapshots: self._print_modified_models( - context_diff, modified_snapshots, environment_naming_info, default_catalog, no_diff + context_diff, + modified_snapshots, + environment_naming_info, + default_catalog, + no_diff, ) def _print_models_with_threshold( @@ -3430,7 +3700,9 @@ def _print_models_with_threshold( ) else: for snapshot_table_info in models: - category_str = SNAPSHOT_CHANGE_CATEGORY_STR[snapshot_table_info.change_category] + category_str = SNAPSHOT_CHANGE_CATEGORY_STR[ + snapshot_table_info.change_category + ] self._print( f"- `{snapshot_table_info.display_name(environment_naming_info, default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, dialect=self.dialect)}` ({category_str})" ) @@ -3462,7 +3734,11 @@ def _print_modified_models( ) indirectly_modified_children = sorted( - [s for s in indirectly_modified if snapshot.snapshot_id in s.parents] + [ + s + for s in indirectly_modified + if snapshot.snapshot_id in s.parents + ] ) if not no_diff: @@ -3539,7 +3815,9 @@ def _show_missing_dates(self, plan: Plan, default_catalog: t.Optional[str]) -> N for snap in snapshots: self._print(snap) - def _show_categorized_snapshots(self, plan: Plan, default_catalog: t.Optional[str]) -> None: + def _show_categorized_snapshots( + self, plan: Plan, default_catalog: t.Optional[str] + ) -> None: context_diff = plan.context_diff for snapshot in plan.categorized: if context_diff.directly_modified(snapshot.name): @@ -3548,7 +3826,9 @@ def _show_categorized_snapshots(self, plan: Plan, default_catalog: t.Optional[st f"[bold][direct]Directly Modified: {snapshot.display_name(plan.environment_naming_info, default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, dialect=self.dialect)} ({category_str})" ) indirect_tree = None - for child_sid in sorted(plan.indirectly_modified.get(snapshot.snapshot_id, set())): + for child_sid in sorted( + plan.indirectly_modified.get(snapshot.snapshot_id, set()) + ): child_snapshot = context_diff.snapshots[child_sid] if not indirect_tree: indirect_tree = Tree("[indirect]Indirectly Modified Children:") @@ -3560,7 +3840,9 @@ def _show_categorized_snapshots(self, plan: Plan, default_catalog: t.Optional[st f"[indirect]{child_snapshot.display_name(plan.environment_naming_info, default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, dialect=self.dialect)} ({child_category_str})" ) if indirect_tree: - indirect_tree = self._limit_model_names(indirect_tree, self.verbosity) + indirect_tree = self._limit_model_names( + indirect_tree, self.verbosity + ) elif context_diff.metadata_updated(snapshot.name): tree = Tree( f"[bold][metadata]Metadata Updated: {snapshot.display_name(plan.environment_naming_info, default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, dialect=self.dialect)}" @@ -3585,8 +3867,12 @@ def stop_promotion_progress(self, success: bool = True) -> None: super().stop_promotion_progress(success) self._print("\n") - def log_warning(self, short_message: str, long_message: t.Optional[str] = None) -> None: - super().log_warning(short_message, long_message, print=not self.warning_capture_only) + def log_warning( + self, short_message: str, long_message: t.Optional[str] = None + ) -> None: + super().log_warning( + short_message, long_message, print=not self.warning_capture_only + ) def log_error(self, message: str) -> None: super().log_error(message, print=not self.error_capture_only) @@ -3594,7 +3880,9 @@ def log_error(self, message: str) -> None: def log_success(self, message: str) -> None: self._print(message) - def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> None: + def log_test_results( + self, result: ModelTextTestResult, target_dialect: str + ) -> None: # We don't log the test results if no tests were ran if not result.testsRun: return @@ -3668,9 +3956,7 @@ def _render_alert_block(self, block_type: str, items: t.List[str]) -> str: item_contents += f">\n> {list_indicator}{item}\n" if len(item_contents) > self.alert_block_max_content_length: - truncation_msg = ( - "...\n>\n> Truncated. Please check the console for full information.\n" - ) + truncation_msg = "...\n>\n> Truncated. Please check the console for full information.\n" item_contents = item_contents[ 0 : self.alert_block_max_content_length - len(truncation_msg) ] @@ -3728,7 +4014,8 @@ def start_evaluation_progress( audit_only: bool = False, ) -> None: self.evaluation_model_batch_sizes = { - snapshot: len(intervals) for snapshot, intervals in batched_intervals.items() + snapshot: len(intervals) + for snapshot, intervals in batched_intervals.items() } self.evaluation_environment_naming_info = environment_naming_info self.default_catalog = default_catalog @@ -3739,7 +4026,11 @@ def start_snapshot_evaluation_progress( if not self.evaluation_batch_progress.get(snapshot.snapshot_id): display_name = snapshot.display_name( self.evaluation_environment_naming_info, - self.default_catalog if self.verbosity < Verbosity.VERY_VERBOSE else None, + ( + self.default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) self.evaluation_batch_progress[snapshot.snapshot_id] = (display_name, 0) @@ -3768,17 +4059,23 @@ def update_snapshot_evaluation_progress( total_batches = self.evaluation_model_batch_sizes[snapshot] loaded_batches += 1 - self.evaluation_batch_progress[snapshot.snapshot_id] = (view_name, loaded_batches) + self.evaluation_batch_progress[snapshot.snapshot_id] = ( + view_name, + loaded_batches, + ) finished_loading = loaded_batches == total_batches status = "Loaded" if finished_loading else "Loading" - print(f"{status} '{view_name}', Completed Batches: {loaded_batches}/{total_batches}") + print( + f"{status} '{view_name}', Completed Batches: {loaded_batches}/{total_batches}" + ) if finished_loading: total_finished_loading = len( [ s for s, total in self.evaluation_model_batch_sizes.items() - if self.evaluation_batch_progress.get(s.snapshot_id, (None, -1))[1] == total + if self.evaluation_batch_progress.get(s.snapshot_id, (None, -1))[1] + == total ] ) total = len(self.evaluation_batch_progress) @@ -3822,7 +4119,9 @@ def start_promotion_progress( self.promotion_status = (0, len(snapshots)) print(f"Virtually Updating '{environment_naming_info.name}'") - def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) -> None: + def update_promotion_progress( + self, snapshot: SnapshotInfoLike, promoted: bool + ) -> None: """Update the snapshot promotion progress.""" num_promotions, total_promotions = self.promotion_status num_promotions += 1 @@ -3969,7 +4268,9 @@ def start_promotion_progress( if snapshots: self._write(f"Starting promotion for {len(snapshots)} snapshots") - def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) -> None: + def update_promotion_progress( + self, snapshot: SnapshotInfoLike, promoted: bool + ) -> None: self._write(f"Promoting {snapshot.name}") def stop_promotion_progress(self, success: bool = True) -> None: @@ -4029,7 +4330,9 @@ def show_model_difference_summary( for modified in context_diff.modified_snapshots: self._write(f" Modified: {modified}") - def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> None: + def log_test_results( + self, result: ModelTextTestResult, target_dialect: str + ) -> None: self._write("Test Results:", result) def show_sql(self, sql: str) -> None: @@ -4041,7 +4344,9 @@ def log_status_update(self, message: str) -> None: def log_error(self, message: str) -> None: self._write(message, style="bold red") - def log_warning(self, short_message: str, long_message: t.Optional[str] = None) -> None: + def log_warning( + self, short_message: str, long_message: t.Optional[str] = None + ) -> None: logger.warning(long_message or short_message) if not self.ignore_warnings: self._write(short_message, style="bold yellow") @@ -4060,7 +4365,10 @@ def show_schema_diff(self, schema_diff: SchemaDiff) -> None: self._write(schema_diff) def show_row_diff( - self, row_diff: RowDiff, show_sample: bool = True, skip_grain_check: bool = False + self, + row_diff: RowDiff, + show_sample: bool = True, + skip_grain_check: bool = False, ) -> None: self._write(row_diff) @@ -4075,7 +4383,9 @@ def show_table_diff( self.show_table_diff_summary(table_diff) self.show_schema_diff(table_diff.schema_diff()) self.show_row_diff( - table_diff.row_diff(temp_schema=temp_schema, skip_grain_check=skip_grain_check), + table_diff.row_diff( + temp_schema=temp_schema, skip_grain_check=skip_grain_check + ), show_sample=show_sample, skip_grain_check=skip_grain_check, ) @@ -4165,9 +4475,7 @@ def _format_missing_intervals(snapshot: Snapshot, missing: SnapshotIntervals) -> return ( missing.format_intervals(snapshot.node.interval_unit) if snapshot.is_incremental - else "recreate view" - if snapshot.is_view - else "full refresh" + else "recreate view" if snapshot.is_view else "full refresh" ) @@ -4183,11 +4491,11 @@ def _format_node_error(ex: NodeExecutionFailedError) -> str: error_msg = _format_audits_errors(cause) elif not isinstance(cause, (NodeExecutionFailedError, PythonModelEvalError)): error_msg = " " + error_msg.replace("\n", "\n ") - error_msg = ( - f" {cause.__class__.__name__}:\n{error_msg}" # include error class name in msg - ) + error_msg = f" {cause.__class__.__name__}:\n{error_msg}" # include error class name in msg error_msg = error_msg.replace("\n", "\n ") - error_msg = error_msg + "\n" if not error_msg.rstrip(" ").endswith("\n") else error_msg + error_msg = ( + error_msg + "\n" if not error_msg.rstrip(" ").endswith("\n") else error_msg + ) return error_msg @@ -4216,12 +4524,18 @@ def _format_audits_errors(error: NodeAuditsErrors) -> str: for err in error.errors: audit_args_sql = [] for arg_name, arg_value in err.audit_args.items(): - audit_args_sql.append(f"{arg_name} := {arg_value.sql(dialect=err.adapter_dialect)}") - audit_args_sql_msg = ("\n".join(audit_args_sql) + "\n\n") if audit_args_sql else "" + audit_args_sql.append( + f"{arg_name} := {arg_value.sql(dialect=err.adapter_dialect)}" + ) + audit_args_sql_msg = ( + ("\n".join(audit_args_sql) + "\n\n") if audit_args_sql else "" + ) err_msg = f"'{err.audit_name}' audit error: {err.count} {'row' if err.count == 1 else 'rows'} failed" - query = "\n ".join(textwrap.wrap(err.sql(err.adapter_dialect), width=LINE_WRAP_WIDTH)) + query = "\n ".join( + textwrap.wrap(err.sql(err.adapter_dialect), width=LINE_WRAP_WIDTH) + ) msg = f"{err_msg}\n\nAudit arguments\n {audit_args_sql_msg}Audit query\n {query}\n\n" msg = msg.replace("\n", "\n ") error_messages.append(msg) @@ -4270,8 +4584,12 @@ def _create_evaluation_model_annotation( rows_processed = execution_stats.total_rows_processed if rows_processed: # 1.00 and 1.0 to 1 - rows_processed_str = metric(rows_processed).replace(".00", "").replace(".0", "") - execution_stats_str += f"{rows_processed_str} row{'s' if rows_processed > 1 else ''}" + rows_processed_str = ( + metric(rows_processed).replace(".00", "").replace(".0", "") + ) + execution_stats_str += ( + f"{rows_processed_str} row{'s' if rows_processed > 1 else ''}" + ) bytes_processed = execution_stats.total_bytes_processed execution_stats_str += ( @@ -4314,7 +4632,9 @@ def _calculate_interval_str_len( interval_str_len, len( _create_evaluation_model_annotation( - snapshot, _format_evaluation_model_interval(snapshot, interval), execution_stats + snapshot, + _format_evaluation_model_interval(snapshot, interval), + execution_stats, ) ), ) @@ -4330,7 +4650,8 @@ def _calculate_audit_str_len(snapshot: Snapshot, audit_padding: int = 0) -> int: if snapshot.is_audit: # +1 for "1" audit count, +1 for red X audit_str_len = max( - audit_str_len, audit_base_str_len + (2 if not snapshot.audit.blocking else 1) + audit_str_len, + audit_base_str_len + (2 if not snapshot.audit.blocking else 1), ) if snapshot.is_model and snapshot.model.audits: num_audits = len(snapshot.model.audits_with_args) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index c3abff1d94..918b68b2be 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -40,13 +40,13 @@ import time import traceback import typing as t +from datetime import datetime from functools import cached_property from io import StringIO from itertools import chain from pathlib import Path from shutil import rmtree from types import MappingProxyType -from datetime import datetime from sqlglot import Dialect, exp from sqlglot.helper import first @@ -56,89 +56,56 @@ from sqlmesh.core import constants as c from sqlmesh.core.analytics import python_api_analytics from sqlmesh.core.audit import Audit, ModelAudit, StandaloneAudit -from sqlmesh.core.config import ( - CategorizerConfig, - Config, - load_configs, -) +from sqlmesh.core.config import CategorizerConfig, Config, load_configs from sqlmesh.core.config.connection import ConnectionConfig from sqlmesh.core.config.loader import C +from sqlmesh.core.config.model import ModelDefaultsConfig from sqlmesh.core.config.root import RegexKeyDict from sqlmesh.core.console import get_console from sqlmesh.core.context_diff import ContextDiff -from sqlmesh.core.dialect import ( - format_model_expressions, - is_meta_expression, - normalize_model_name, - pandas_to_sql, - parse, - parse_one, -) +from sqlmesh.core.dialect import (format_model_expressions, is_meta_expression, + normalize_model_name, pandas_to_sql, parse, + parse_one) from sqlmesh.core.engine_adapter import EngineAdapter -from sqlmesh.core.environment import Environment, EnvironmentNamingInfo, EnvironmentStatements -from sqlmesh.core.loader import Loader +from sqlmesh.core.environment import (Environment, EnvironmentNamingInfo, + EnvironmentStatements) +from sqlmesh.core.janitor import (cleanup_expired_views, + delete_expired_snapshots) from sqlmesh.core.linter.definition import AnnotatedRuleViolation, Linter from sqlmesh.core.linter.rules import BUILTIN_RULES +from sqlmesh.core.loader import Loader from sqlmesh.core.macros import ExecutableOrMacro, macro from sqlmesh.core.metric import Metric, rewrite from sqlmesh.core.model import Model, update_model_schemas -from sqlmesh.core.config.model import ModelDefaultsConfig -from sqlmesh.core.notification_target import ( - NotificationEvent, - NotificationTarget, - NotificationTargetManager, -) -from sqlmesh.core.plan import Plan, PlanBuilder, SnapshotIntervals, PlanExplainer +from sqlmesh.core.notification_target import (NotificationEvent, + NotificationTarget, + NotificationTargetManager) +from sqlmesh.core.plan import (Plan, PlanBuilder, PlanExplainer, + SnapshotIntervals) from sqlmesh.core.plan.definition import UserProvidedFlags from sqlmesh.core.reference import ReferenceGraph -from sqlmesh.core.scheduler import Scheduler, CompletionStatus +from sqlmesh.core.scheduler import CompletionStatus, Scheduler from sqlmesh.core.schema_loader import create_external_models_file -from sqlmesh.core.selector import Selector, NativeSelector -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Snapshot, - SnapshotEvaluator, - SnapshotFingerprint, - missing_intervals, - to_table_mapping, -) +from sqlmesh.core.selector import NativeSelector, Selector +from sqlmesh.core.snapshot import (DeployabilityIndex, Snapshot, + SnapshotEvaluator, SnapshotFingerprint, + missing_intervals, to_table_mapping) from sqlmesh.core.snapshot.definition import get_next_model_interval_start -from sqlmesh.core.state_sync import ( - CachingStateSync, - StateReader, - StateSync, -) -from sqlmesh.core.janitor import cleanup_expired_views, delete_expired_snapshots +from sqlmesh.core.state_sync import CachingStateSync, StateReader, StateSync from sqlmesh.core.table_diff import TableDiff -from sqlmesh.core.test import ( - ModelTextTestResult, - ModelTestMetadata, - generate_test, - run_tests, - filter_tests_by_patterns, -) +from sqlmesh.core.test import (ModelTestMetadata, ModelTextTestResult, + filter_tests_by_patterns, generate_test, + run_tests) from sqlmesh.core.user import User from sqlmesh.utils import CorrelationId, UniqueKeyDict, Verbosity from sqlmesh.utils.concurrency import concurrent_apply_to_values -from sqlmesh.utils.dag import DAG -from sqlmesh.utils.date import ( - TimeLike, - to_timestamp, - format_tz_datetime, - now_timestamp, - now, - to_datetime, - make_exclusive, -) -from sqlmesh.utils.errors import ( - CircuitBreakerError, - ConfigError, - PlanError, - SQLMeshError, - UncategorizedPlanError, - LinterError, -) from sqlmesh.utils.config import print_config +from sqlmesh.utils.dag import DAG +from sqlmesh.utils.date import (TimeLike, format_tz_datetime, make_exclusive, + now, now_timestamp, to_datetime, to_timestamp) +from sqlmesh.utils.errors import (CircuitBreakerError, ConfigError, + LinterError, PlanError, SQLMeshError, + UncategorizedPlanError) from sqlmesh.utils.jinja import JinjaMacroRegistry from sqlmesh.utils.windows import IS_WINDOWS, fix_windows_path @@ -146,15 +113,11 @@ import pandas as pd from typing_extensions import Literal - from sqlmesh.core.engine_adapter._typing import ( - BigframeSession, - DF, - PySparkDataFrame, - PySparkSession, - SnowparkSession, - ) + from sqlmesh.core.engine_adapter._typing import (DF, BigframeSession, + PySparkDataFrame, + PySparkSession, + SnowparkSession) from sqlmesh.core.snapshot import Node - from sqlmesh.core.snapshot.definition import Intervals ModelOrSnapshot = t.Union[str, Model, Snapshot] @@ -216,7 +179,9 @@ def resolve_table(self, model_name: str) -> str: Returns: The physical table name. """ - model_name = normalize_model_name(model_name, self.default_catalog, self.default_dialect) + model_name = normalize_model_name( + model_name, self.default_catalog, self.default_dialect + ) if model_name not in self._model_tables: model_name_list = "\n".join(list(self._model_tables)) @@ -259,7 +224,9 @@ def fetch_pyspark_df( Returns: A PySpark dataframe. """ - return self.engine_adapter.fetch_pyspark_df(query, quote_identifiers=quote_identifiers) + return self.engine_adapter.fetch_pyspark_df( + query, quote_identifiers=quote_identifiers + ) class ExecutionContext(BaseContext): @@ -324,11 +291,15 @@ def is_restatement(self) -> t.Optional[bool]: def parent_intervals(self) -> t.Optional[Intervals]: return self._parent_intervals - def var(self, var_name: str, default: t.Optional[t.Any] = None) -> t.Optional[t.Any]: + def var( + self, var_name: str, default: t.Optional[t.Any] = None + ) -> t.Optional[t.Any]: """Returns a variable value.""" return self._variables.get(var_name.lower(), default) - def blueprint_var(self, var_name: str, default: t.Optional[t.Any] = None) -> t.Optional[t.Any]: + def blueprint_var( + self, var_name: str, default: t.Optional[t.Any] = None + ) -> t.Optional[t.Any]: """Returns a blueprint variable value.""" return self._blueprint_variables.get(var_name.lower(), default) @@ -394,7 +365,9 @@ def __init__( self.configs = ( config if isinstance(config, dict) - else load_configs(config, self.CONFIG_TYPE, paths, **(config_loader_kwargs or {})) + else load_configs( + config, self.CONFIG_TYPE, paths, **(config_loader_kwargs or {}) + ) ) self._projects = {config.project for config in self.configs.values()} self.dag: DAG[str] = DAG() @@ -404,8 +377,12 @@ def __init__( "standaloneaudits" ) self._model_test_metadata: t.List[ModelTestMetadata] = [] - self._model_test_metadata_path_index: t.Dict[Path, t.List[ModelTestMetadata]] = {} - self._model_test_metadata_fully_qualified_name_index: t.Dict[str, ModelTestMetadata] = {} + self._model_test_metadata_path_index: t.Dict[ + Path, t.List[ModelTestMetadata] + ] = {} + self._model_test_metadata_fully_qualified_name_index: t.Dict[ + str, ModelTestMetadata + ] = {} self._models_with_tests: t.Set[str] = set() self._macros: UniqueKeyDict[str, ExecutableOrMacro] = UniqueKeyDict("macros") @@ -420,7 +397,9 @@ def __init__( self._load_state: bool = load_state self._selector_cls = selector or NativeSelector - self.path, self.config = t.cast(t.Tuple[Path, C], next(iter(self.configs.items()))) + self.path, self.config = t.cast( + t.Tuple[Path, C], next(iter(self.configs.items())) + ) self._all_dialects: t.Set[str] = {self.config.dialect or ""} @@ -430,11 +409,15 @@ def __init__( self.gateway = gateway self._scheduler = self.config.get_scheduler(self.gateway) self.environment_ttl = self.config.environment_ttl - self.pinned_environments = Environment.sanitize_names(self.config.pinned_environments) + 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).lower() - gw_model_defaults = self.config.get_gateway(self.selected_gateway).model_defaults + gw_model_defaults = self.config.get_gateway( + self.selected_gateway + ).model_defaults if gw_model_defaults: # Merge global model defaults with the selected gateway's, if it's overriden global_defaults = self.config.model_defaults.model_dump(exclude_unset=True) @@ -470,7 +453,9 @@ def __init__( self._state_sync: t.Optional[StateSync] = None # Should we dedupe notification_targets? If so how? - self.notification_targets = (notification_targets or []) + self.config.notification_targets + self.notification_targets = ( + notification_targets or [] + ) + self.config.notification_targets self.users = (users or []) + self.config.users self.users = list({user.username: user for user in self.users}.values()) self._register_notification_targets() @@ -554,7 +539,10 @@ def upsert_model(self, model: t.Union[str, Model], **kwargs: t.Any) -> Model: { model.fqn: model, # bust the fingerprint cache for all downstream models - **{fqn: self._models[fqn].copy() for fqn in self.dag.downstream(model.fqn)}, + **{ + fqn: self._models[fqn].copy() + for fqn in self.dag.downstream(model.fqn) + }, } ) @@ -590,14 +578,18 @@ def scheduler( stored_environment = self.state_sync.get_environment(environment) if stored_environment is None: raise ConfigError(f"Environment '{environment}' was not found.") - snapshots = self.state_sync.get_snapshots(stored_environment.snapshots).values() + snapshots = self.state_sync.get_snapshots( + stored_environment.snapshots + ).values() else: snapshots = self.snapshots.values() if not snapshots: raise ConfigError("No models were found") - return self.create_scheduler(snapshots, snapshot_evaluator or self.snapshot_evaluator) + return self.create_scheduler( + snapshots, snapshot_evaluator or self.snapshot_evaluator + ) def create_scheduler( self, snapshots: t.Iterable[Snapshot], snapshot_evaluator: SnapshotEvaluator @@ -692,7 +684,9 @@ def load(self, update_schemas: bool = True) -> GenericContext[C]: if self._load_state and any(self._projects): prod = self.state_reader.get_environment(c.PROD) if prod: - existing_statements = self.state_reader.get_environment_statements(c.PROD) + existing_statements = self.state_reader.get_environment_statements( + c.PROD + ) for stmt in existing_statements: if stmt.project and stmt.project not in self._projects: self._environment_statements.append(stmt) @@ -703,11 +697,17 @@ def load(self, update_schemas: bool = True) -> GenericContext[C]: prod = self.state_reader.get_environment(c.PROD) if prod: - for snapshot in self.state_reader.get_snapshots(prod.snapshots).values(): + for snapshot in self.state_reader.get_snapshots( + prod.snapshots + ).values(): if snapshot.node.project in self._projects: uncached.add(snapshot.name) else: - local_store = self._standalone_audits if snapshot.is_audit else self._models + local_store = ( + self._standalone_audits + if snapshot.is_audit + else self._models + ) if snapshot.name in local_store: uncached.add(snapshot.name) else: @@ -728,7 +728,9 @@ def load(self, update_schemas: bool = True) -> GenericContext[C]: # model will get mutated (schema changes) but the object is the same as the remote cache if any(dep in uncached for dep in model.depends_on): uncached.add(fqn) - self._models.update({fqn: model.copy(update={"mapping_schema": {}})}) + self._models.update( + {fqn: model.copy(update={"mapping_schema": {}})} + ) continue update_model_schemas( @@ -919,7 +921,9 @@ def run_janitor( if self.console.start_cleanup(ignore_ttl): try: - self._run_janitor(ignore_ttl, force_delete=force_delete, environment=environment) + self._run_janitor( + ignore_ttl, force_delete=force_delete, environment=environment + ) success = True finally: self.console.stop_cleanup(success=success) @@ -951,7 +955,10 @@ def destroy(self) -> bool: else: adapter = self.engine_adapter - if environment.suffix_target.is_schema or environment.suffix_target.is_catalog: + if ( + environment.suffix_target.is_schema + or environment.suffix_target.is_catalog + ): schema = snapshot.qualified_view_name.schema_for_environment( environment.naming_info, dialect=adapter.dialect ) @@ -973,7 +980,9 @@ def destroy(self) -> bool: table_name = snapshot.table_name() tables_to_delete.add(table_name) - if self.console.start_destroy(schemas_to_delete, views_to_delete, tables_to_delete): + if self.console.start_destroy( + schemas_to_delete, views_to_delete, tables_to_delete + ): try: success = self._destroy() finally: @@ -1034,7 +1043,9 @@ def get_model( return None @t.overload - def get_snapshot(self, node_or_snapshot: NodeOrSnapshot) -> t.Optional[Snapshot]: ... + def get_snapshot( + self, node_or_snapshot: NodeOrSnapshot + ) -> t.Optional[Snapshot]: ... @t.overload def get_snapshot( @@ -1154,7 +1165,9 @@ def render( if expand and not isinstance(expand, bool): expand = { normalize_model_name( - x, default_catalog=self.default_catalog, dialect=self.default_dialect + x, + default_catalog=self.default_catalog, + dialect=self.default_dialect, ) for x in expand } @@ -1309,7 +1322,9 @@ def _format( append_newline: t.Optional[bool] = None, **kwargs: t.Any, ) -> str: - expressions = parse(before, default_dialect=self.config_for_node(target).dialect) + expressions = parse( + before, default_dialect=self.config_for_node(target).dialect + ) if transpile and is_meta_expression(expressions[0]): for prop in expressions[0].expressions: if prop.name.lower() == "dialect": @@ -1325,7 +1340,9 @@ def _format( expressions, transpile or target.dialect, rewrite_casts=( - rewrite_casts if rewrite_casts is not None else not format_config.no_rewrite_casts + rewrite_casts + if rewrite_casts is not None + else not format_config.no_rewrite_casts ), **{**format_config.generator_options, **kwargs}, ) @@ -1469,7 +1486,9 @@ def plan( auto_apply if auto_apply is not None else self.config.plan.auto_apply, self.default_catalog, no_diff=no_diff if no_diff is not None else self.config.plan.no_diff, - no_prompts=no_prompts if no_prompts is not None else self.config.plan.no_prompts, + no_prompts=( + no_prompts if no_prompts is not None else self.config.plan.no_prompts + ), ) return plan @@ -1559,22 +1578,30 @@ def plan_builder( "execution_time": execution_time, "create_from": create_from, "skip_tests": skip_tests, - "restate_models": list(restate_models) if restate_models is not None else None, + "restate_models": ( + list(restate_models) if restate_models is not None else None + ), "no_gaps": no_gaps, "skip_backfill": skip_backfill, "empty_backfill": empty_backfill, "forward_only": forward_only, - "allow_destructive_models": list(allow_destructive_models) - if allow_destructive_models is not None - else None, - "allow_additive_models": list(allow_additive_models) - if allow_additive_models is not None - else None, + "allow_destructive_models": ( + list(allow_destructive_models) + if allow_destructive_models is not None + else None + ), + "allow_additive_models": ( + list(allow_additive_models) + if allow_additive_models is not None + else None + ), "no_auto_categorization": no_auto_categorization, "effective_from": effective_from, "include_unmodified": include_unmodified, "select_models": list(select_models) if select_models is not None else None, - "backfill_models": list(backfill_models) if backfill_models is not None else None, + "backfill_models": ( + list(backfill_models) if backfill_models is not None else None + ), "enable_preview": enable_preview, "preview_start": preview_start, "preview_min_intervals": preview_min_intervals, @@ -1617,7 +1644,9 @@ def plan_builder( self._run_plan_tests(skip_tests=skip_tests) environment_ttl = ( - self.environment_ttl if environment not in self.pinned_environments else None + self.environment_ttl + if environment not in self.pinned_environments + else None ) model_selector = self._new_selector() @@ -1630,7 +1659,9 @@ def plan_builder( expanded_destructive_models = None if allow_additive_models: - expanded_additive_models = model_selector.expand_model_selections(allow_additive_models) + expanded_additive_models = model_selector.expand_model_selections( + allow_additive_models + ) else: expanded_additive_models = None @@ -1666,10 +1697,14 @@ def plan_builder( expanded_restate_models = None if restate_models is not None: - expanded_restate_models = model_selector.expand_model_selections(restate_models) + expanded_restate_models = model_selector.expand_model_selections( + restate_models + ) if (restate_models is not None and not expanded_restate_models) or ( - backfill_models is not None and not backfill_models and not selected_deletion_fqns + backfill_models is not None + and not backfill_models + and not selected_deletion_fqns ): raise PlanError( "Selector did not return any models. Please check your model selection and try again." @@ -1678,7 +1713,9 @@ def plan_builder( if always_include_local_changes is None: # default behaviour - if restatements are detected; we operate entirely out of state and ignore local changes force_no_diff = restate_models is not None or ( - backfill_models is not None and not backfill_models and not selected_deletion_fqns + backfill_models is not None + and not backfill_models + and not selected_deletion_fqns ) else: force_no_diff = not always_include_local_changes @@ -1725,7 +1762,9 @@ def plan_builder( execution_time or now(), ) - execution_time_ts = to_timestamp(execution_time) if execution_time is not None else None + execution_time_ts = ( + to_timestamp(execution_time) if execution_time is not None else None + ) if ( execution_time_ts is not None and end is None @@ -1796,7 +1835,9 @@ def plan_builder( default_start=default_start, default_end=default_end, enable_preview=( - enable_preview if enable_preview is not None else self._plan_preview_enabled + enable_preview + if enable_preview is not None + else self._plan_preview_enabled ), preview_start=preview_start, preview_min_intervals=preview_min_intervals or 0, @@ -1808,7 +1849,9 @@ def plan_builder( user_provided_flags=user_provided_flags, selected_models={ dbt_unique_id - for model in model_selector.expand_model_selections(select_models or "*") + for model in model_selector.expand_model_selections( + select_models or "*" + ) if (dbt_unique_id := snapshots[model].node.dbt_unique_id) }, explain=explain or False, @@ -1836,7 +1879,9 @@ def apply( ): return if plan.uncategorized: - raise UncategorizedPlanError("Can't apply a plan with uncategorized changes.") + raise UncategorizedPlanError( + "Can't apply a plan with uncategorized changes." + ) if plan.explain: explainer = PlanExplainer( @@ -1985,7 +2030,9 @@ def table_diff( raise SQLMeshError(f"Could not find environment '{target}'") criteria = ", ".join(f"'{c}'" for c in select_models) try: - selected_models = self._new_selector().expand_model_selections(select_models) + selected_models = self._new_selector().expand_model_selections( + select_models + ) if not selected_models: self.console.log_status_update( f"No models matched the selection criteria: {criteria}" @@ -1994,7 +2041,9 @@ def table_diff( raise SQLMeshError(e) models_to_diff: t.List[ - t.Tuple[Model, EngineAdapter, str, str, t.Optional[t.List[str] | exp.Expr]] + t.Tuple[ + Model, EngineAdapter, str, str, t.Optional[t.List[str] | exp.Expr] + ] ] = [] models_without_grain: t.List[Model] = [] source_snapshots_to_name = { @@ -2011,7 +2060,9 @@ def table_diff( target_snapshot = target_snapshots_to_name.get(model.fqn) if target_snapshot and source_snapshot: - if (source_snapshot.fingerprint != target_snapshot.fingerprint) and ( + if ( + source_snapshot.fingerprint != target_snapshot.fingerprint + ) and ( (source_snapshot.version != target_snapshot.version) or source_snapshot.is_forward_only ): @@ -2027,11 +2078,14 @@ def table_diff( if not model_on: models_without_grain.append(model) else: - models_to_diff.append((model, adapter, source, target, model_on)) + models_to_diff.append( + (model, adapter, source, target, model_on) + ) if models_without_grain: model_names = "\n".join( - f"─ {model.name} \n at '{model._path}'" for model in models_without_grain + f"─ {model.name} \n at '{model._path}'" + for model in models_without_grain ) message = ( "SQLMesh doesn't know how to join the tables for the following models:\n" @@ -2099,7 +2153,9 @@ def table_diff( ] if show: - self.console.show_table_diff(table_diffs, show_sample, skip_grain_check, temp_schema) + self.console.show_table_diff( + table_diffs, show_sample, skip_grain_check, temp_schema + ) return table_diffs @@ -2140,7 +2196,9 @@ def _model_diff( if show: # Trigger row_diff in parallel execution so it's available for ordered display later - table_diff.row_diff(temp_schema=temp_schema, skip_grain_check=skip_grain_check) + table_diff.row_diff( + temp_schema=temp_schema, skip_grain_check=skip_grain_check + ) self.console.update_table_diff_progress(model.name) @@ -2232,7 +2290,9 @@ def get_dag( ) @python_api_analytics - def render_dag(self, path: str, select_models: t.Optional[t.Collection[str]] = None) -> None: + def render_dag( + self, path: str, select_models: t.Optional[t.Collection[str]] = None + ) -> None: """Render the dag as HTML and save it to a file. Args: @@ -2396,7 +2456,9 @@ def audit( f"{audit_id} ❌ [red]FAIL [{audit_result.count}][/red]." ) else: - self.console.log_status_update(f"{audit_id} ✅ [green]PASS[/green].") + self.console.log_status_update( + f"{audit_id} ✅ [green]PASS[/green]." + ) self.console.log_status_update( f"\nFinished with {len(errors)} audit error{'' if len(errors) == 1 else 's'} " @@ -2458,7 +2520,9 @@ def check_intervals( if not env: raise SQLMeshError(f"Environment '{environment}' was not found.") - snapshots = {k.name: v for k, v in self.state_sync.get_snapshots(env.snapshots).items()} + snapshots = { + k.name: v for k, v in self.state_sync.get_snapshots(env.snapshots).items() + } missing = { k.name: v @@ -2483,9 +2547,11 @@ def check_intervals( results[snapshot] = SnapshotIntervals( snapshot.snapshot_id, - intervals - if no_signals - else snapshot.check_ready_intervals(intervals, execution_context), + ( + intervals + if no_signals + else snapshot.check_ready_intervals(intervals, execution_context) + ), ) return results @@ -2534,10 +2600,14 @@ def create_external_models(self, strict: bool = False) -> None: deprecated_yaml = path / c.EXTERNAL_MODELS_DEPRECATED_YAML external_models_yaml = ( - path / c.EXTERNAL_MODELS_YAML if not deprecated_yaml.exists() else deprecated_yaml + path / c.EXTERNAL_MODELS_YAML + if not deprecated_yaml.exists() + else deprecated_yaml ) - external_models_gateway: t.Optional[str] = self.gateway or self.config.default_gateway + external_models_gateway: t.Optional[str] = ( + self.gateway or self.config.default_gateway + ) if not external_models_gateway: # can happen if there was no --gateway defined and the default_gateway is '' # which means that the single gateway syntax is being used which means there is @@ -2569,25 +2639,35 @@ def print_info( ) -> None: """Prints information about connections, models, macros, etc. to the console.""" self.console.log_status_update(f"Models: {len(self.models)}") - self.console.log_status_update(f"Macros: {len(self._macros) - len(macro.get_registry())}") + self.console.log_status_update( + f"Macros: {len(self._macros) - len(macro.get_registry())}" + ) if skip_connection: return if verbosity >= Verbosity.VERBOSE: self.console.log_status_update("") - print_config(self.config.get_connection(self.gateway), self.console, "Connection") print_config( - self.config.get_test_connection(self.gateway), self.console, "Test Connection" + self.config.get_connection(self.gateway), self.console, "Connection" ) print_config( - self.config.get_state_connection(self.gateway), self.console, "State Connection" + self.config.get_test_connection(self.gateway), + self.console, + "Test Connection", + ) + print_config( + self.config.get_state_connection(self.gateway), + self.console, + "State Connection", ) self._try_connection("data warehouse", self.engine_adapter.ping) state_connection = self.config.get_state_connection(self.gateway) if state_connection: - self._try_connection("state backend", state_connection.connection_validator()) + self._try_connection( + "state backend", state_connection.connection_validator() + ) @python_api_analytics def print_environment_names(self) -> None: @@ -2620,7 +2700,9 @@ def _run( no_auto_upstream: bool, snapshot_evaluator: t.Optional[SnapshotEvaluator] = None, ) -> CompletionStatus: - scheduler = self.scheduler(environment=environment, snapshot_evaluator=snapshot_evaluator) + scheduler = self.scheduler( + environment=environment, snapshot_evaluator=snapshot_evaluator + ) snapshots = scheduler.snapshots if select_models is not None: @@ -2643,11 +2725,19 @@ def _run( if completion_status.is_nothing_to_do: next_run_ready_msg = "" - next_ready_interval_start = get_next_model_interval_start(snapshots.values()) + next_ready_interval_start = get_next_model_interval_start( + snapshots.values() + ) if next_ready_interval_start: utc_time = format_tz_datetime(next_ready_interval_start) - local_time = format_tz_datetime(next_ready_interval_start, use_local_timezone=True) - time_msg = local_time if local_time == utc_time else f"{local_time} ({utc_time})" + local_time = format_tz_datetime( + next_ready_interval_start, use_local_timezone=True + ) + time_msg = ( + local_time + if local_time == utc_time + else f"{local_time} ({utc_time})" + ) next_run_ready_msg = f"\n\nNext run will be ready at {time_msg}." self.console.log_status_update( @@ -2656,7 +2746,9 @@ def _run( return completion_status - def _apply(self, plan: Plan, circuit_breaker: t.Optional[t.Callable[[], bool]]) -> None: + def _apply( + self, plan: Plan, circuit_breaker: t.Optional[t.Callable[[], bool]] + ) -> None: self._scheduler.create_plan_evaluator(self).evaluate( plan.to_evaluatable(), circuit_breaker=circuit_breaker ) @@ -2752,7 +2844,9 @@ def export_state( self.console.stop_state_export(success=False, output_file=output_file) raise - def import_state(self, input_file: Path, clear: bool = False, confirm: bool = True) -> None: + def import_state( + self, input_file: Path, clear: bool = False, confirm: bool = True + ) -> None: from sqlmesh.core.state_sync.export_import import import_state if self.console.start_state_import( @@ -2781,7 +2875,9 @@ def _run_tests( result = self.test(stream=test_output_io, verbosity=verbosity) return result, test_output_io.getvalue() - def _run_plan_tests(self, skip_tests: bool = False) -> t.Optional[ModelTextTestResult]: + def _run_plan_tests( + self, skip_tests: bool = False + ) -> t.Optional[ModelTextTestResult]: if not skip_tests: result = self.test() if not result.wasSuccessful(): @@ -2830,7 +2926,8 @@ def _warn_if_virtual_catalog_rematerialization(self, plan: "Plan") -> None: max_display = 10 model_lines = "\n".join( - f" - {new_name} (was: {old_name})" for new_name, old_name in affected[:max_display] + f" - {new_name} (was: {old_name})" + for new_name, old_name in affected[:max_display] ) if len(affected) > max_display: model_lines += f"\n ... and {len(affected) - max_display} more" @@ -2881,7 +2978,9 @@ def cache_dir(self) -> Path: @cached_property def engine_adapters(self) -> t.Dict[str, EngineAdapter]: """Returns all the engine adapters for the gateways defined in the configurations.""" - adapters: t.Dict[str, EngineAdapter] = {self.selected_gateway: self.engine_adapter} + adapters: t.Dict[str, EngineAdapter] = { + self.selected_gateway: self.engine_adapter + } for config in self.configs.values(): for gateway_name in config.gateways: if gateway_name not in adapters: @@ -2937,7 +3036,9 @@ def _get_engine_adapter(self, gateway: t.Optional[str] = None) -> EngineAdapter: if gateway: if adapter := self.engine_adapters.get(gateway): return adapter - raise SQLMeshError(f"Gateway '{gateway}' not found in the available engine adapters.") + raise SQLMeshError( + f"Gateway '{gateway}' not found in the available engine adapters." + ) return self.engine_adapter def _snapshots( @@ -2955,7 +3056,8 @@ def _snapshots( if unrestorable_snapshots: for snapshot in unrestorable_snapshots: logger.info( - "Found a unrestorable snapshot %s. Restamping the model...", snapshot.name + "Found a unrestorable snapshot %s. Restamping the model...", + snapshot.name, ) node = nodes[snapshot.name] nodes[snapshot.name] = node.copy( @@ -2968,7 +3070,10 @@ def _snapshots( # Keep the original model instance to preserve the query cache. snapshot.node = snapshots[snapshot.name].node - return {name: stored_snapshots.get(s.snapshot_id, s) for name, s in snapshots.items()} + return { + name: stored_snapshots.get(s.snapshot_id, s) + for name, s in snapshots.items() + } def _context_diff( self, @@ -3002,7 +3107,9 @@ def _context_diff( def _destroy(self) -> bool: # Invalidate all environments, including prod for environment in self.state_reader.get_environments(): - self.state_sync.invalidate_environment(name=environment.name, protect_prod=False) + self.state_sync.invalidate_environment( + name=environment.name, protect_prod=False + ) self.console.log_success(f"Environment '{environment.name}' invalidated.") # Run janitor to clean up all objects @@ -3098,7 +3205,9 @@ def _cleanup_environments( # we want to retry on the next janitor pass if drops failed, unless # force_delete is set in which case we purge state records regardless if not failures or force_delete: - self.state_sync.delete_expired_environments(current_ts=current_ts, name=name) + self.state_sync.delete_expired_environments( + current_ts=current_ts, name=name + ) return failures def _cleanup_adapters_for_environment( @@ -3120,7 +3229,8 @@ def _cleanup_adapters_for_environment( gateway = ( snapshot.model_gateway - if environment.gateway_managed and snapshot.model_gateway in engine_adapters + if environment.gateway_managed + and snapshot.model_gateway in engine_adapters else self.selected_gateway ) adapter = engine_adapters.get(gateway, default_adapter) @@ -3132,7 +3242,9 @@ def _cleanup_adapters_for_environment( for gateway, catalogs in catalogs_by_gateway.items(): if len(catalogs) > 1: - catalogs_description = ", ".join(f"'{catalog}'" for catalog in sorted(catalogs)) + catalogs_description = ", ".join( + f"'{catalog}'" for catalog in sorted(catalogs) + ) return ( default_adapter, engine_adapters, @@ -3145,7 +3257,9 @@ def _cleanup_adapters_for_environment( cleanup_engine_adapters = engine_adapters.copy() cleanup_default_adapter = default_adapter for gateway, catalogs in catalogs_by_gateway.items(): - cleanup_adapter = engine_adapters.get(gateway, default_adapter).with_settings() + cleanup_adapter = engine_adapters.get( + gateway, default_adapter + ).with_settings() cleanup_adapter.inject_virtual_catalog(gateway) # inject_virtual_catalog() may initialize adapter-specific state in addition to # _default_catalog. Override only the cleanup clone with the catalog persisted in the @@ -3157,11 +3271,15 @@ def _cleanup_adapters_for_environment( return cleanup_default_adapter, cleanup_engine_adapters, None - def _try_connection(self, connection_name: str, validator: t.Callable[[], None]) -> None: + def _try_connection( + self, connection_name: str, validator: t.Callable[[], None] + ) -> None: connection_name = connection_name.capitalize() try: validator() - self.console.log_status_update(f"{connection_name} connection [green]succeeded[/green]") + self.console.log_status_update( + f"{connection_name} connection [green]succeeded[/green]" + ) except Exception as ex: self.console.log_error(f"{connection_name} connection failed. {ex}") @@ -3169,7 +3287,9 @@ def _new_state_sync(self) -> StateSync: return self._provided_state_sync or self._scheduler.create_state_sync(self) def _new_selector( - self, models: t.Optional[UniqueKeyDict[str, Model]] = None, dag: t.Optional[DAG[str]] = None + self, + models: t.Optional[UniqueKeyDict[str, Model]] = None, + dag: t.Optional[DAG[str]] = None, ) -> Selector: return self._selector_cls( self.state_reader, @@ -3194,7 +3314,9 @@ def _register_notification_targets(self) -> None: for user in self.users } self.notification_target_manager = NotificationTargetManager( - event_notifications, user_notification_targets, username=self.config.username + event_notifications, + user_notification_targets, + username=self.config.username, ) def _load_materializations(self) -> None: @@ -3237,7 +3359,9 @@ def _nodes_to_snapshots(self, nodes: t.Dict[str, Node]) -> t.Dict[str, Snapshot] if node.project in self._projects: config = self.config_for_node(node) kwargs["ttl"] = config.snapshot_ttl - kwargs["table_naming_convention"] = config.physical_table_naming_convention + kwargs["table_naming_convention"] = ( + config.physical_table_naming_convention + ) snapshot = Snapshot.from_node( node, @@ -3251,7 +3375,9 @@ def _nodes_to_snapshots(self, nodes: t.Dict[str, Node]) -> t.Dict[str, Snapshot] def _node_or_snapshot_to_fqn(self, node_or_snapshot: NodeOrSnapshot) -> str: if isinstance(node_or_snapshot, Snapshot): return node_or_snapshot.name - if isinstance(node_or_snapshot, str) and not self.standalone_audits.get(node_or_snapshot): + if isinstance(node_or_snapshot, str) and not self.standalone_audits.get( + node_or_snapshot + ): return normalize_model_name( node_or_snapshot, dialect=self.default_dialect, @@ -3291,7 +3417,9 @@ def _get_plan_default_start_end( default_end = to_timestamp(max(non_seed_interval_ends.values())) default_start: t.Optional[int] = None # Infer the default start by finding the smallest interval start that corresponds to the default end. - for model_name in backfill_models or modified_model_names or max_interval_end_per_model: + for model_name in ( + backfill_models or modified_model_names or max_interval_end_per_model + ): if model_name not in snapshots: continue node = snapshots[model_name].node @@ -3366,7 +3494,10 @@ def _calculate_start_override_per_model( # this works because topological ordering guarantees that they've already been visited # and we always set a start override min_child_start = min( - [start_overrides[immediate_child_fqn] for immediate_child_fqn in graph[model_fqn]], + [ + start_overrides[immediate_child_fqn] + for immediate_child_fqn in graph[model_fqn] + ], default=plan_start_dt, ) @@ -3481,7 +3612,9 @@ def select_tests( else: test_path = Path(test) if test_path in self._model_test_metadata_path_index: - filtered_tests.extend(self._model_test_metadata_path_index[test_path]) + filtered_tests.extend( + self._model_test_metadata_path_index[test_path] + ) test_meta = filtered_tests diff --git a/sqlmesh/core/context_diff.py b/sqlmesh/core/context_diff.py index 047e58609a..a11e477d5c 100644 --- a/sqlmesh/core/context_diff.py +++ b/sqlmesh/core/context_diff.py @@ -16,6 +16,7 @@ import typing as t from difflib import ndiff, unified_diff from functools import cached_property + from sqlmesh.core import constants as c from sqlmesh.core.console import get_console from sqlmesh.core.macros import RuntimeStage @@ -33,8 +34,8 @@ if t.TYPE_CHECKING: from sqlmesh.core.state_sync import StateReader -from sqlmesh.utils.metaprogramming import Executable # noqa from sqlmesh.core.environment import EnvironmentStatements +from sqlmesh.utils.metaprogramming import Executable # noqa IGNORED_PACKAGES = {"sqlmesh", "sqlglot", "sqlglotc"} @@ -132,7 +133,9 @@ def create( existing_env = state_reader.get_environment(environment) create_from_env_exists = False - recreate_environment = always_recreate_environment and not environment == create_from + recreate_environment = ( + always_recreate_environment and not environment == create_from + ) if existing_env is None or existing_env.expired or recreate_environment: env = state_reader.get_environment(create_from.lower()) @@ -148,7 +151,9 @@ def create( else: env = existing_env is_new_environment = False - previously_promoted_snapshot_ids = {s.snapshot_id for s in env.promoted_snapshots} + previously_promoted_snapshot_ids = { + s.snapshot_id for s in env.promoted_snapshots + } environment_snapshot_infos = [] if env: @@ -158,7 +163,8 @@ def create( else env.finalized_or_current_snapshots ) remote_snapshot_name_to_info = { - snapshot_info.name: snapshot_info for snapshot_info in environment_snapshot_infos + snapshot_info.name: snapshot_info + for snapshot_info in environment_snapshot_infos } removed = { snapshot_table_info.snapshot_id: snapshot_table_info @@ -174,7 +180,8 @@ def create( snapshot.name: remote_snapshot_name_to_info[snapshot.name] for snapshot in snapshots.values() if snapshot.snapshot_id not in added - and snapshot.fingerprint != remote_snapshot_name_to_info[snapshot.name].fingerprint + and snapshot.fingerprint + != remote_snapshot_name_to_info[snapshot.name].fingerprint } stored = state_reader.get_snapshots( @@ -187,10 +194,15 @@ def create( for snapshot in snapshots.values(): s_id = snapshot.snapshot_id - modified_snapshot_info = modified_snapshot_name_to_snapshot_info.get(snapshot.name) + modified_snapshot_info = modified_snapshot_name_to_snapshot_info.get( + snapshot.name + ) existing_snapshot = stored.get(s_id) - if modified_snapshot_info and snapshot.node_type != modified_snapshot_info.node_type: + if ( + modified_snapshot_info + and snapshot.node_type != modified_snapshot_info.node_type + ): added.add(snapshot.snapshot_id) removed[modified_snapshot_info.snapshot_id] = modified_snapshot_info modified_snapshot_name_to_snapshot_info.pop(snapshot.name) @@ -235,7 +247,8 @@ def create( environment=environment, is_new_environment=is_new_environment, is_unfinalized_environment=bool(env and not env.finalized_ts), - normalize_environment_name=is_new_environment or bool(env and env.normalize_name), + normalize_environment_name=is_new_environment + or bool(env and env.normalize_name), create_from=create_from, create_from_env_exists=create_from_env_exists, added=added, @@ -245,13 +258,17 @@ def create( new_snapshots=new_snapshots, previous_plan_id=previous_plan_id, previously_promoted_snapshot_ids=previously_promoted_snapshot_ids, - previous_finalized_snapshots=env.previous_finalized_snapshots if env else None, + previous_finalized_snapshots=( + env.previous_finalized_snapshots if env else None + ), previous_requirements=env.requirements if env else {}, requirements=requirements, 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, + previous_gateway_managed_virtual_layer=( + env.gateway_managed if env else False + ), gateway_managed_virtual_layer=gateway_managed_virtual_layer, ) @@ -268,7 +285,9 @@ def create_no_diff(cls, environment: str, state_reader: StateReader) -> ContextD """ env = state_reader.get_environment(environment.lower()) if not env: - raise SQLMeshError(f"Environment '{environment}' must exist for this operation.") + raise SQLMeshError( + f"Environment '{environment}' must exist for this operation." + ) environment_statements = state_reader.get_environment_statements(environment) snapshots = state_reader.get_snapshots(env.snapshots) @@ -286,7 +305,9 @@ def create_no_diff(cls, environment: str, state_reader: StateReader) -> ContextD snapshots=snapshots, new_snapshots={}, previous_plan_id=env.plan_id, - previously_promoted_snapshot_ids={s.snapshot_id for s in env.promoted_snapshots}, + previously_promoted_snapshot_ids={ + s.snapshot_id for s in env.promoted_snapshots + }, previous_finalized_snapshots=env.previous_finalized_snapshots, previous_requirements=env.requirements, requirements=env.requirements, @@ -304,7 +325,8 @@ 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 + or self.previous_gateway_managed_virtual_layer + != self.gateway_managed_virtual_layer ) @property @@ -313,9 +335,9 @@ def has_requirement_changes(self) -> bool: @property def has_environment_statements_changes(self) -> bool: - return sorted(self.environment_statements, key=lambda s: s.project or "") != sorted( - self.previous_environment_statements, key=lambda s: s.project or "" - ) + return sorted( + self.environment_statements, key=lambda s: s.project or "" + ) != sorted(self.previous_environment_statements, key=lambda s: s.project or "") @property def has_snapshot_changes(self) -> bool: @@ -367,7 +389,9 @@ def requirements_diff(self) -> str: def environment_statements_diff( self, include_python_env: bool = False ) -> t.List[t.Tuple[str, str]]: - def extract_statements(statements: t.List[EnvironmentStatements], attr: str) -> t.List[str]: + def extract_statements( + statements: t.List[EnvironmentStatements], attr: str + ) -> t.List[str]: return [ string for statement in statements @@ -380,7 +404,9 @@ def extract_statements(statements: t.List[EnvironmentStatements], attr: str) -> ] def compute_diff(attribute: str) -> t.Optional[t.Tuple[str, str]]: - previous = extract_statements(self.previous_environment_statements, attribute) + previous = extract_statements( + self.previous_environment_statements, attribute + ) current = extract_statements(self.environment_statements, attribute) if previous == current: diff --git a/sqlmesh/core/dialect.py b/sqlmesh/core/dialect.py index e4ab522198..6aa419c910 100644 --- a/sqlmesh/core/dialect.py +++ b/sqlmesh/core/dialect.py @@ -10,28 +10,29 @@ from enum import Enum, auto from functools import lru_cache -from sqlglot import Dialect, Generator, ParseError, Parser, Tokenizer, TokenType, exp -from sqlglot.dialects.dialect import DialectType -from sqlglot.dialects import DuckDB, Snowflake, TSQL import sqlglot.dialects.athena as athena import sqlglot.generators.athena as athena_generators -from sqlglot.parsers.athena import AthenaTrinoParser +from sqlglot import (Dialect, Generator, ParseError, Parser, Tokenizer, + TokenType, exp) +from sqlglot.dialects import TSQL, DuckDB, Snowflake +from sqlglot.dialects.dialect import DialectType from sqlglot.helper import seq_get from sqlglot.optimizer.normalize_identifiers import normalize_identifiers from sqlglot.optimizer.qualify_columns import quote_identifiers from sqlglot.optimizer.qualify_tables import qualify_tables from sqlglot.optimizer.scope import traverse_scope +from sqlglot.parsers.athena import AthenaTrinoParser from sqlglot.schema import MappingSchema from sqlglot.tokens import Token -from sqlmesh.core.constants import LIQUID_CLUSTERING_KEYWORDS, MAX_MODEL_DEFINITION_SIZE +from sqlmesh.core.constants import (LIQUID_CLUSTERING_KEYWORDS, + MAX_MODEL_DEFINITION_SIZE) from sqlmesh.utils import get_source_columns_to_types -from sqlmesh.utils.errors import SQLMeshError, ConfigError +from sqlmesh.utils.errors import ConfigError, SQLMeshError from sqlmesh.utils.pandas import columns_to_types_from_df if t.TYPE_CHECKING: import pandas as pd - from sqlglot._typing import E @@ -166,7 +167,11 @@ def _parse_id_var( any_token: bool = True, tokens: t.Optional[t.Collection[TokenType]] = None, ) -> t.Optional[exp.Expr]: - if self._prev and self._prev.text == SQLMESH_MACRO_PREFIX and self._match(TokenType.L_BRACE): + if ( + self._prev + and self._prev.text == SQLMESH_MACRO_PREFIX + and self._match(TokenType.L_BRACE) + ): identifier = self.__parse_id_var(any_token=any_token, tokens=tokens) # type: ignore if not self._match(TokenType.R_BRACE): self.raise_error("Expecting }") @@ -209,7 +214,9 @@ def _parse_id_var( else: self.raise_error("Expecting }") - identifier = self.expression(exp.Identifier(this=this, quoted=identifier.quoted)) + identifier = self.expression( + exp.Identifier(this=this, quoted=identifier.quoted) + ) return identifier @@ -220,7 +227,11 @@ def _parse_macro(self: Parser, keyword_macro: str = "") -> t.Optional[exp.Expr]: comments = self._prev.comments index = self._index - field = self._parse_primary() or self._parse_function(functions={}) or self._parse_id_var() + field = ( + self._parse_primary() + or self._parse_function(functions={}) + or self._parse_id_var() + ) def _build_macro(field: t.Optional[exp.Expr]) -> t.Optional[exp.Expr]: if isinstance(field, exp.Func): @@ -239,9 +250,14 @@ def _build_macro(field: t.Optional[exp.Expr]) -> t.Optional[exp.Expr]: comments=comments, ) if macro_name == "SQL": - into = field.expressions[1].this.lower() if len(field.expressions) > 1 else None + into = ( + field.expressions[1].this.lower() + if len(field.expressions) > 1 + else None + ) return self.expression( - MacroSQL(this=field.expressions[0], into=into), comments=comments + MacroSQL(this=field.expressions[0], into=into), + comments=comments, ) else: field = self.expression( @@ -339,7 +355,9 @@ def _parse_join( def _warn_unsupported(self: Parser) -> None: from sqlmesh.core.console import get_console - sql = self._find_sql(self._tokens[0], self._tokens[-1])[: self.error_message_context] + sql = self._find_sql(self._tokens[0], self._tokens[-1])[ + : self.error_message_context + ] get_console().log_warning( f"'{sql}' could not be semantically understood as it contains unsupported syntax, SQLMesh will treat the command as is. Note that any references to the model's " @@ -385,7 +403,9 @@ def _parse_where(self: Parser, skip_where_token: bool = False) -> t.Optional[exp return macro -def _parse_group(self: Parser, skip_group_by_token: bool = False) -> t.Optional[exp.Expr]: +def _parse_group( + self: Parser, skip_group_by_token: bool = False +) -> t.Optional[exp.Expr]: macro = _parse_matching_macro(self, "GROUP_BY") if not macro: return self.__parse_group(skip_group_by_token=skip_group_by_token) # type: ignore @@ -394,7 +414,9 @@ def _parse_group(self: Parser, skip_group_by_token: bool = False) -> t.Optional[ return macro -def _parse_having(self: Parser, skip_having_token: bool = False) -> t.Optional[exp.Expr]: +def _parse_having( + self: Parser, skip_having_token: bool = False +) -> t.Optional[exp.Expr]: macro = _parse_matching_macro(self, "HAVING") if not macro: return self.__parse_having(skip_having_token=skip_having_token) # type: ignore @@ -464,7 +486,9 @@ def _parse_props(self: Parser) -> t.Optional[exp.Expr]: elif name == "merge_filter": value = self._parse_conjunction() elif self._match(TokenType.L_PAREN): - value = self.expression(exp.Tuple(expressions=self._parse_csv(self._parse_equality))) + value = self.expression( + exp.Tuple(expressions=self._parse_csv(self._parse_equality)) + ) self._match_r_paren() else: value = self._parse_bracket(self._parse_field(any_token=True)) @@ -621,7 +645,9 @@ def _parse_if(self: Parser) -> t.Optional[exp.Expr]: return exp.Anonymous(this="IF", expressions=[cond, stmt]) -def _create_parser(expression_type: t.Type[exp.Expr], table_keys: t.List[str]) -> t.Callable: +def _create_parser( + expression_type: t.Type[exp.Expr], table_keys: t.List[str] +) -> t.Callable: def parse(self: Parser) -> t.Optional[exp.Expr]: from sqlmesh.core.model.kind import ModelKindName @@ -629,7 +655,10 @@ def parse(self: Parser) -> t.Optional[exp.Expr]: while True: prev_property = seq_get(expressions, -1) - if not self._match(TokenType.COMMA, expression=prev_property) and expressions: + if ( + not self._match(TokenType.COMMA, expression=prev_property) + and expressions + ): break key_expression = self._parse_id_var(any_token=True) @@ -656,7 +685,9 @@ def parse(self: Parser) -> t.Optional[exp.Expr]: elif key == "columns": value = self._parse_schema() elif key == "kind": - field = _parse_macro_or_clause(self, lambda: self._parse_id_var(any_token=True)) + field = _parse_macro_or_clause( + self, lambda: self._parse_id_var(any_token=True) + ) if not field or isinstance(field, (MacroVar, MacroFunc)): value = field @@ -680,11 +711,15 @@ def parse(self: Parser) -> t.Optional[exp.Expr]: ModelKindName.SCD_TYPE_2_BY_COLUMN, ModelKindName.CUSTOM, ) and self._match(TokenType.L_PAREN, advance=False): - props = self._parse_wrapped_csv(functools.partial(_parse_props, self)) + props = self._parse_wrapped_csv( + functools.partial(_parse_props, self) + ) else: props = None - value = self.expression(ModelKind(this=kind.value, expressions=props)) + value = self.expression( + ModelKind(this=kind.value, expressions=props) + ) elif key == "expression": value = self._parse_conjunction() elif key == "partitioned_by": @@ -710,7 +745,9 @@ def parse(self: Parser) -> t.Optional[exp.Expr]: # Unwrap Paren wrapping a bare column to match partitioned_by normalisation: # clustered_by (a) → stored as Column(a), not Paren(Column(a)). # Preserve parens around function expressions: (TO_DATE(col)) stays as-is. - if isinstance(parsed, exp.Paren) and isinstance(parsed.this, exp.Column): + if isinstance(parsed, exp.Paren) and isinstance( + parsed.this, exp.Column + ): value = parsed.unnest() else: value = parsed @@ -754,15 +791,19 @@ def _props_sql(self: Generator, expressions: t.List[exp.Expr]) -> str: def _on_virtual_update_sql(self: Generator, expressions: t.List[exp.Expr]) -> str: statements = "\n".join( - self.sql(expression) - if isinstance(expression, JinjaStatement) - else f"{self.sql(expression)};" + ( + self.sql(expression) + if isinstance(expression, JinjaStatement) + else f"{self.sql(expression)};" + ) for expression in expressions ) return f"{ON_VIRTUAL_UPDATE_BEGIN};\n{statements}\n{ON_VIRTUAL_UPDATE_END};" -def _sqlmesh_ddl_sql(self: Generator, expression: Model | Audit | Metric, name: str) -> str: +def _sqlmesh_ddl_sql( + self: Generator, expression: Model | Audit | Metric, name: str +) -> str: return "\n".join([f"{name} (", _props_sql(self, expression.expressions), ")"]) @@ -874,7 +915,9 @@ def cast_to_colon(node: exp.Expr) -> exp.Expr: ): this = node.this - if not isinstance(this, (exp.Binary, exp.Unary)) or isinstance(this, exp.Paren): + if not isinstance(this, (exp.Binary, exp.Unary)) or isinstance( + this, exp.Paren + ): cast = DColonCast(this=this, to=node.to) cast.comments = node.comments node = cast @@ -1052,9 +1095,11 @@ def parse( chunks.append( ( [], - ChunkType.VIRTUAL_STATEMENT - if virtual and tokens[pos] != ON_VIRTUAL_UPDATE_END - else ChunkType.SQL, + ( + ChunkType.VIRTUAL_STATEMENT + if virtual and tokens[pos] != ON_VIRTUAL_UPDATE_END + else ChunkType.SQL + ), ) ) elif _is_jinja_query_begin(tokens, pos): @@ -1074,9 +1119,13 @@ def parse( parser = dialect.parser() expressions: t.List[exp.Expr] = [] - def parse_sql_chunk(chunk: t.List[Token], meta_sql: bool = True) -> t.List[exp.Expr]: + def parse_sql_chunk( + chunk: t.List[Token], meta_sql: bool = True + ) -> t.List[exp.Expr]: parsed_expressions: t.List[t.Optional[exp.Expr]] = ( - parser.parse(chunk, sql) if into is None else parser.parse_into(into, chunk, sql) + parser.parse(chunk, sql) + if into is None + else parser.parse_into(into, chunk, sql) ) expressions = [] for expression in parsed_expressions: @@ -1089,7 +1138,9 @@ def parse_sql_chunk(chunk: t.List[Token], meta_sql: bool = True) -> t.List[exp.E def parse_jinja_chunk(chunk: t.List[Token], meta_sql: bool = True) -> exp.Expr: start, *_, end = chunk segment = sql[start.end + 2 : end.start - 1] - factory = jinja_query if chunk_type == ChunkType.JINJA_QUERY else jinja_statement + factory = ( + jinja_query if chunk_type == ChunkType.JINJA_QUERY else jinja_statement + ) expression = factory(segment.strip()) if meta_sql: expression.meta["sql"] = sql[start.start : end.end + 1] @@ -1103,7 +1154,8 @@ def parse_virtual_statement( start = chunks[pos][0][0].start while ( - chunks[pos - 1][0] == [] or chunks[pos - 1][0][-1].text.upper() != ON_VIRTUAL_UPDATE_END + chunks[pos - 1][0] == [] + or chunks[pos - 1][0][-1].text.upper() != ON_VIRTUAL_UPDATE_END ): chunk, chunk_type = chunks[pos] if chunk_type == ChunkType.JINJA_STATEMENT: @@ -1111,7 +1163,10 @@ def parse_virtual_statement( else: virtual_update_statements.extend( parse_sql_chunk( - chunk[int(chunk[0].text.upper() == ON_VIRTUAL_UPDATE_BEGIN) : -1], False + chunk[ + int(chunk[0].text.upper() == ON_VIRTUAL_UPDATE_BEGIN) : -1 + ], + False, ), ) pos += 1 @@ -1166,7 +1221,9 @@ def extend_sqlglot() -> None: tokenizer.VAR_SINGLE_TOKENS.update(SQLMESH_MACRO_PREFIX) for parser in parsers: - parser.FUNCTIONS.update({"JINJA": Jinja.from_arg_list, "METRIC": MetricAgg.from_arg_list}) + parser.FUNCTIONS.update( + {"JINJA": Jinja.from_arg_list, "METRIC": MetricAgg.from_arg_list} + ) parser.PLACEHOLDER_PARSERS.update({TokenType.PARAMETER: _parse_macro}) parser.QUERY_MODIFIER_PARSERS.update( {TokenType.PARAMETER: lambda self: _parse_body_macro(self)} @@ -1183,7 +1240,9 @@ def extend_sqlglot() -> None: JinjaStatement: lambda self, e: ( f"{JINJA_STATEMENT_BEGIN};\n{e.name}\n{JINJA_END};" ), - VirtualUpdateStatement: lambda self, e: _on_virtual_update_sql(self, e), + VirtualUpdateStatement: lambda self, e: _on_virtual_update_sql( + self, e + ), MacroDef: lambda self, e: f"@DEF({self.sql(e.this)}, {self.sql(e.expression)})", MacroFunc: _macro_func_sql, MacroStrReplace: lambda self, e: f"@{self.sql(e.this)}", @@ -1192,7 +1251,9 @@ def extend_sqlglot() -> None: Metric: lambda self, e: _sqlmesh_ddl_sql(self, e, "METRIC"), Model: lambda self, e: _sqlmesh_ddl_sql(self, e, "MODEL"), ModelKind: _model_kind_sql, - PythonCode: lambda self, e: self.expressions(e, sep="\n", indent=False), + PythonCode: lambda self, e: self.expressions( + e, sep="\n", indent=False + ), StagedFilePath: lambda self, e: self.table_sql(e), exp.Whens: _whens_sql, } @@ -1275,13 +1336,18 @@ def select_from_values_for_batch_range( source_columns: t.Optional[t.List[str]] = None, ) -> exp.Select: source_columns = source_columns or list(target_columns_to_types) - source_columns_to_types = get_source_columns_to_types(target_columns_to_types, source_columns) + source_columns_to_types = get_source_columns_to_types( + target_columns_to_types, source_columns + ) if not values: # Ensures we don't generate an empty VALUES clause & forces a zero-row output where = exp.false() expressions = [ - tuple(exp.cast(exp.null(), to=kind) for kind in source_columns_to_types.values()) + tuple( + exp.cast(exp.null(), to=kind) + for kind in source_columns_to_types.values() + ) ] else: where = None @@ -1304,14 +1370,19 @@ def select_from_values_for_batch_range( casted_columns = [ exp.alias_( exp.cast( - exp.column(column) if column in source_columns_to_types else exp.Null(), to=kind + exp.column(column) if column in source_columns_to_types else exp.Null(), + to=kind, ), column, copy=False, ) for column, kind in target_columns_to_types.items() ] - return exp.select(*casted_columns).from_(values_exp, copy=False).where(where, copy=False) + return ( + exp.select(*casted_columns) + .from_(values_exp, copy=False) + .where(where, copy=False) + ) def pandas_to_sql( @@ -1358,7 +1429,9 @@ def normalize_model_name( dialect: DialectType = None, ) -> str: if isinstance(table, exp.Column): - table = exp.table_(table.this, db=table.args.get("table"), catalog=table.args.get("db")) + table = exp.table_( + table.this, db=table.args.get("table"), catalog=table.args.get("db") + ) else: # We are relying on sqlglot's flexible parsing here to accept quotes from other dialects. # Ex: I have a a normalized name of '"my_table"' but the dialect is spark and therefore we should @@ -1391,7 +1464,9 @@ def find_tables( """ if TABLES_META not in expression.meta: expression.meta[TABLES_META] = { - normalize_model_name(table, default_catalog=default_catalog, dialect=dialect) + normalize_model_name( + table, default_catalog=default_catalog, dialect=dialect + ) for scope in traverse_scope(expression) for table in scope.tables if table.name and table.name not in scope.cte_sources @@ -1432,7 +1507,9 @@ def _transform_value(value: t.Any, dtype: exp.DataType) -> t.Any: and len(value) == len(dtype.expressions) ): expressions = [] - for (field_name, field_value), field_type in zip(value.items(), dtype.expressions): + for (field_name, field_value), field_type in zip( + value.items(), dtype.expressions + ): if isinstance(field_type, exp.ColumnDef): field_type = field_type.kind else: @@ -1460,7 +1537,8 @@ def to_schema(sql_path: str | exp.Table, dialect: DialectType = None) -> exp.Tab if isinstance(sql_path, exp.Table) and sql_path.this is None: return sql_path table = exp.to_table( - sql_path.copy() if isinstance(sql_path, exp.Table) else sql_path, dialect=dialect + sql_path.copy() if isinstance(sql_path, exp.Table) else sql_path, + dialect=dialect, ) table.set("catalog", table.args.get("db")) table.set("db", table.args.get("this")) @@ -1497,7 +1575,8 @@ def normalize_mapping_schema(schema: t.Dict, dialect: DialectType) -> MappingSch def _unquote_schema(schema: t.Dict) -> t.Dict: """SQLGlot schema expects unquoted normalized keys.""" return { - k.strip('"'): _unquote_schema(v) if isinstance(v, dict) else v for k, v in schema.items() + k.strip('"'): _unquote_schema(v) if isinstance(v, dict) else v + for k, v in schema.items() } @@ -1564,7 +1643,10 @@ def extract_function_calls(func_calls: t.Any, allow_tuples: bool = False) -> t.A """Used for extracting function calls for signals or audits.""" if isinstance(func_calls, (exp.Tuple, exp.Array)): - return [extract_func_call(i, allow_tuples=allow_tuples) for i in func_calls.expressions] + return [ + extract_func_call(i, allow_tuples=allow_tuples) + for i in func_calls.expressions + ] if isinstance(func_calls, exp.Paren): return [extract_func_call(func_calls.this, allow_tuples=allow_tuples)] if isinstance(func_calls, exp.Expr): @@ -1578,7 +1660,9 @@ def extract_function_calls(func_calls: t.Any, allow_tuples: bool = False) -> t.A elif isinstance(entry, (tuple, list)): name, args = entry else: - raise ConfigError(f"Audit must be a dictionary or named tuple. Got {entry}.") + raise ConfigError( + f"Audit must be a dictionary or named tuple. Got {entry}." + ) function_calls.append( ( @@ -1599,17 +1683,28 @@ def is_meta_expression(v: t.Any) -> bool: return isinstance(v, (Audit, Metric, Model)) -def replace_merge_table_aliases(expression: exp.Expr, dialect: t.Optional[str] = None) -> exp.Expr: +def replace_merge_table_aliases( + expression: exp.Expr, dialect: t.Optional[str] = None +) -> exp.Expr: """ Resolves references from the "source" and "target" tables (or their DBT equivalents) with the corresponding SQLMesh merge aliases (MERGE_SOURCE_ALIAS and MERGE_TARGET_ALIAS) """ - from sqlmesh.core.engine_adapter.base import MERGE_SOURCE_ALIAS, MERGE_TARGET_ALIAS + from sqlmesh.core.engine_adapter.base import (MERGE_SOURCE_ALIAS, + MERGE_TARGET_ALIAS) if isinstance(expression, exp.Column) and (first_part := expression.parts[0]): - if first_part.this.lower() in ("target", "dbt_internal_dest", "__merge_target__"): + if first_part.this.lower() in ( + "target", + "dbt_internal_dest", + "__merge_target__", + ): first_part.replace(exp.to_identifier(MERGE_TARGET_ALIAS, quoted=True)) - elif first_part.this.lower() in ("source", "dbt_internal_source", "__merge_source__"): + elif first_part.this.lower() in ( + "source", + "dbt_internal_source", + "__merge_source__", + ): first_part.replace(exp.to_identifier(MERGE_SOURCE_ALIAS, quoted=True)) return expression diff --git a/sqlmesh/core/engine_adapter/__init__.py b/sqlmesh/core/engine_adapter/__init__.py index cb9db5ea77..8297ad35f7 100644 --- a/sqlmesh/core/engine_adapter/__init__.py +++ b/sqlmesh/core/engine_adapter/__init__.py @@ -2,25 +2,23 @@ import typing as t -from sqlmesh.core.engine_adapter.base import ( - EngineAdapter, - EngineAdapterWithIndexSupport, -) +from sqlmesh.core.engine_adapter.athena import AthenaEngineAdapter +from sqlmesh.core.engine_adapter.base import (EngineAdapter, + EngineAdapterWithIndexSupport) from sqlmesh.core.engine_adapter.bigquery import BigQueryEngineAdapter from sqlmesh.core.engine_adapter.clickhouse import ClickhouseEngineAdapter from sqlmesh.core.engine_adapter.databricks import DatabricksEngineAdapter from sqlmesh.core.engine_adapter.duckdb import DuckDBEngineAdapter +from sqlmesh.core.engine_adapter.fabric import FabricEngineAdapter from sqlmesh.core.engine_adapter.mssql import MSSQLEngineAdapter from sqlmesh.core.engine_adapter.mysql import MySQLEngineAdapter from sqlmesh.core.engine_adapter.postgres import PostgresEngineAdapter from sqlmesh.core.engine_adapter.redshift import RedshiftEngineAdapter +from sqlmesh.core.engine_adapter.risingwave import RisingwaveEngineAdapter from sqlmesh.core.engine_adapter.snowflake import SnowflakeEngineAdapter from sqlmesh.core.engine_adapter.spark import SparkEngineAdapter from sqlmesh.core.engine_adapter.starrocks import StarRocksEngineAdapter from sqlmesh.core.engine_adapter.trino import TrinoEngineAdapter -from sqlmesh.core.engine_adapter.athena import AthenaEngineAdapter -from sqlmesh.core.engine_adapter.risingwave import RisingwaveEngineAdapter -from sqlmesh.core.engine_adapter.fabric import FabricEngineAdapter DIALECT_TO_ENGINE_ADAPTER = { "hive": SparkEngineAdapter, diff --git a/sqlmesh/core/engine_adapter/_typing.py b/sqlmesh/core/engine_adapter/_typing.py index 77bcf2c015..e212dfe190 100644 --- a/sqlmesh/core/engine_adapter/_typing.py +++ b/sqlmesh/core/engine_adapter/_typing.py @@ -8,14 +8,18 @@ import pandas as pd import pyspark import pyspark.sql.connect.dataframe - from bigframes.session import Session as BigframeSession # noqa from bigframes.dataframe import DataFrame as BigframeDataFrame + from bigframes.session import Session as BigframeSession # noqa snowpark = optional_import("snowflake.snowpark") Query = exp.Query - PySparkSession = t.Union[pyspark.sql.SparkSession, pyspark.sql.connect.dataframe.SparkSession] - PySparkDataFrame = t.Union[pyspark.sql.DataFrame, pyspark.sql.connect.dataframe.DataFrame] + PySparkSession = t.Union[ + pyspark.sql.SparkSession, pyspark.sql.connect.dataframe.SparkSession + ] + PySparkDataFrame = t.Union[ + pyspark.sql.DataFrame, pyspark.sql.connect.dataframe.DataFrame + ] # snowpark is not available on python 3.12 from snowflake.snowpark import Session as SnowparkSession # noqa diff --git a/sqlmesh/core/engine_adapter/athena.py b/sqlmesh/core/engine_adapter/athena.py index 338381549b..6cdfb431d1 100644 --- a/sqlmesh/core/engine_adapter/athena.py +++ b/sqlmesh/core/engine_adapter/athena.py @@ -1,24 +1,25 @@ from __future__ import annotations -from functools import lru_cache -import typing as t + import logging +import posixpath +import typing as t +from functools import lru_cache + from sqlglot import exp + from sqlmesh.core.dialect import to_schema -from sqlmesh.utils.aws import validate_s3_uri, parse_s3_uri -from sqlmesh.core.engine_adapter.mixins import PandasNativeFetchDFSupportMixin, RowDiffMixin +from sqlmesh.core.engine_adapter.mixins import ( + PandasNativeFetchDFSupportMixin, RowDiffMixin) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + InsertOverwriteStrategy, + SourceQuery) from sqlmesh.core.engine_adapter.trino import TrinoEngineAdapter from sqlmesh.core.node import IntervalUnit -import posixpath +from sqlmesh.utils.aws import parse_s3_uri, validate_s3_uri from sqlmesh.utils.errors import SQLMeshError -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - DataObject, - DataObjectType, - CommentCreationTable, - CommentCreationView, - SourceQuery, - InsertOverwriteStrategy, -) if t.TYPE_CHECKING: from sqlmesh.core._typing import SchemaName, TableName @@ -49,7 +50,10 @@ class AthenaEngineAdapter(PandasNativeFetchDFSupportMixin, RowDiffMixin): SUPPORTED_DROP_CASCADE_OBJECT_KINDS = ["DATABASE", "SCHEMA"] def __init__( - self, *args: t.Any, s3_warehouse_location: t.Optional[str] = None, **kwargs: t.Any + self, + *args: t.Any, + s3_warehouse_location: t.Optional[str] = None, + **kwargs: t.Any, ): # Need to pass s3_warehouse_location to the superclass so that it goes into _extra_config # which means that EngineAdapter.with_settings() keeps this property when it makes a clone @@ -74,7 +78,9 @@ def s3_warehouse_location_or_raise(self) -> str: if location := self.s3_warehouse_location: return location - raise SQLMeshError("s3_warehouse_location was expected to be populated; it isnt") + raise SQLMeshError( + "s3_warehouse_location was expected to be populated; it isnt" + ) @property def catalog_support(self) -> CatalogSupport: @@ -144,7 +150,10 @@ def columns( query = ( exp.select("column_name", "data_type") .from_("information_schema.columns") - .where(exp.column("table_schema").eq(table.db), exp.column("table_name").eq(table.name)) + .where( + exp.column("table_schema").eq(table.db), + exp.column("table_name").eq(table.name), + ) .order_by("ordinal_position") ) result = self.fetchdf(query, quote_identifiers=True) @@ -161,7 +170,9 @@ def _create_schema( properties: t.List[exp.Expr], kind: str, ) -> None: - if location := self._table_location(table_properties=None, table=exp.to_table(schema_name)): + if location := self._table_location( + table_properties=None, table=exp.to_table(schema_name) + ): # don't add extra LocationProperty's if one already exists if not any(p for p in properties if isinstance(p, exp.LocationProperty)): properties.append(location) @@ -217,7 +228,8 @@ def _build_create_table_exp( filtered_expressions = [ e for e in table_name_or_schema.expressions - if isinstance(e, exp.ColumnDef) and e.this.name not in partitioned_by_column_names + if isinstance(e, exp.ColumnDef) + and e.this.name not in partitioned_by_column_names ] table_name_or_schema.args["expressions"] = filtered_expressions @@ -259,11 +271,15 @@ def _build_table_properties_exp( if table_format: properties.append( - exp.Property(this=exp.var("table_type"), value=exp.Literal.string(table_format)) + exp.Property( + this=exp.var("table_type"), value=exp.Literal.string(table_format) + ) ) if table_description: - properties.append(exp.SchemaCommentProperty(this=exp.Literal.string(table_description))) + properties.append( + exp.SchemaCommentProperty(this=exp.Literal.string(table_description)) + ) if partitioned_by: schema_expressions: t.List[exp.Expr] = [] @@ -274,13 +290,17 @@ def _build_table_properties_exp( for match_name, match_dtype in self._find_matching_columns( partitioned_by, target_columns_to_types ): - column_def = exp.ColumnDef(this=exp.to_identifier(match_name), kind=match_dtype) + column_def = exp.ColumnDef( + this=exp.to_identifier(match_name), kind=match_dtype + ) schema_expressions.append(column_def) else: schema_expressions = partitioned_by properties.append( - exp.PartitionedByProperty(this=exp.Schema(expressions=schema_expressions)) + exp.PartitionedByProperty( + this=exp.Schema(expressions=schema_expressions) + ) ) if clustered_by: @@ -289,7 +309,9 @@ def _build_table_properties_exp( # defines `clustered_by` as a List[str] with no way of indicating the number of buckets # # Athena's concept of CLUSTER BY is more like Iceberg's `bucket(, col)` partition transform - logging.warning("clustered_by is not supported in the Athena adapter at this time") + logging.warning( + "clustered_by is not supported in the Athena adapter at this time" + ) if storage_format: if is_iceberg: @@ -299,14 +321,18 @@ def _build_table_properties_exp( # STORED AS PARQUET properties.append(exp.FileFormatProperty(this=storage_format)) - if table and (location := self._table_location_or_raise(table_properties, table)): + if table and ( + location := self._table_location_or_raise(table_properties, table) + ): properties.append(location) if is_iceberg and expression: # To make a CTAS expression persist as iceberg, alongside setting `table_type=iceberg`, you also need to set is_external=false # Note that SQLGlot does the right thing with LocationProperty and writes it as `location` (Iceberg) instead of `external_location` (Hive) # ref: https://docs.aws.amazon.com/athena/latest/ug/create-table-as.html#ctas-table-properties - properties.append(exp.Property(this=exp.var("is_external"), value="false")) + properties.append( + exp.Property(this=exp.var("is_external"), value="false") + ) for name, value in table_properties.items(): properties.append(exp.Property(this=exp.var(name), value=value)) @@ -316,7 +342,9 @@ def _build_table_properties_exp( return None - def drop_table(self, table_name: TableName, exists: bool = True, **kwargs: t.Any) -> None: + def drop_table( + self, table_name: TableName, exists: bool = True, **kwargs: t.Any + ) -> None: table = exp.to_table(table_name) if self._query_table_type(table) == "hive": @@ -364,7 +392,9 @@ def _query_table_type_or_raise(self, table: exp.Table) -> TableType: """ # Note: SHOW TBLPROPERTIES gets parsed by SQLGlot as an exp.Command anyway so we just use a string here # This also means we need to use dialect="hive" instead of dialect="athena" so that the identifiers get the correct quoting (backticks) - for row in self.fetchall(f"SHOW TBLPROPERTIES {table.sql(dialect='hive', identify=True)}"): + for row in self.fetchall( + f"SHOW TBLPROPERTIES {table.sql(dialect='hive', identify=True)}" + ): # This query returns a single column with values like 'EXTERNAL\tTRUE' row_lower = row[0].lower() if "external" in row_lower and "true" in row_lower: @@ -415,17 +445,23 @@ def _table_location( else: return None - full_uri = validate_s3_uri(posixpath.join(base_uri, table.text("this") or ""), base=True) + full_uri = validate_s3_uri( + posixpath.join(base_uri, table.text("this") or ""), base=True + ) return exp.LocationProperty(this=exp.Literal.string(full_uri)) def _find_matching_columns( - self, partitioned_by: t.List[exp.Expr], columns_to_types: t.Dict[str, exp.DataType] + self, + partitioned_by: t.List[exp.Expr], + columns_to_types: t.Dict[str, exp.DataType], ) -> t.List[t.Tuple[str, exp.DataType]]: matches = [] for col in partitioned_by: # TODO: do we care about normalization? key = col.name - if isinstance(col, exp.Column) and (match_dtype := columns_to_types.get(key)): + if isinstance(col, exp.Column) and ( + match_dtype := columns_to_types.get(key) + ): matches.append((key, match_dtype)) return matches @@ -486,7 +522,9 @@ def _insert_overwrite_by_time_partition( **kwargs, ) - def _clear_partition_data(self, table: exp.Table, where: t.Optional[exp.Condition]) -> None: + def _clear_partition_data( + self, table: exp.Table, where: t.Optional[exp.Condition] + ) -> None: if partitions_to_drop := self._list_partitions(table, where): for _, s3_location in partitions_to_drop: logger.debug( @@ -517,16 +555,23 @@ def _list_partitions( if limit: query = query.limit(limit) - partition_values = [list(r) for r in self.fetchall(query, quote_identifiers=True)] + partition_values = [ + list(r) for r in self.fetchall(query, quote_identifiers=True) + ] if partition_values: response = self._glue_client.batch_get_partition( DatabaseName=table.db, TableName=table.name, - PartitionsToGet=[{"Values": [str(v) for v in lst]} for lst in partition_values], + PartitionsToGet=[ + {"Values": [str(v) for v in lst]} for lst in partition_values + ], ) return sorted( - [(p["Values"], p["StorageDescriptor"]["Location"]) for p in response["Partitions"]] + [ + (p["Values"], p["StorageDescriptor"]["Location"]) + for p in response["Partitions"] + ] ) return [] @@ -535,7 +580,11 @@ def _query_table_s3_location(self, table: exp.Table) -> str: response = self._glue_client.get_table(DatabaseName=table.db, Name=table.name) # Athena wont let you create a table without a location, so *theoretically* this should never be empty - if location := response.get("Table", {}).get("StorageDescriptor", {}).get("Location", None): + if ( + location := response.get("Table", {}) + .get("StorageDescriptor", {}) + .get("Location", None) + ): return location raise SQLMeshError(f"Table {table} has no location set in the metastore!") @@ -598,7 +647,9 @@ def _clear_s3_location(self, s3_uri: str) -> None: keys_to_delete.append(keys) for chunk in keys_to_delete: - s3.delete_objects(Bucket=bucket, Delete={"Objects": [{"Key": k} for k in chunk]}) + s3.delete_objects( + Bucket=bucket, Delete={"Objects": [{"Key": k} for k in chunk]} + ) @property def _glue_client(self) -> t.Any: diff --git a/sqlmesh/core/engine_adapter/base.py b/sqlmesh/core/engine_adapter/base.py index bd435db76f..b6c6896dbf 100644 --- a/sqlmesh/core/engine_adapter/base.py +++ b/sqlmesh/core/engine_adapter/base.py @@ -21,55 +21,38 @@ from sqlglot.helper import ensure_list, seq_get from sqlglot.optimizer.qualify_columns import quote_identifiers -from sqlmesh.core.dialect import ( - add_table, - schema_, - select_from_values_for_batch_range, - to_schema, -) -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, - EngineRunMode, - InsertOverwriteStrategy, - SourceQuery, - set_catalog, -) +from sqlmesh.core.dialect import (add_table, schema_, + select_from_values_for_batch_range, + to_schema) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + EngineRunMode, + InsertOverwriteStrategy, + SourceQuery, set_catalog) from sqlmesh.core.model.kind import TimeColumn from sqlmesh.core.schema_diff import SchemaDiffer, TableAlterOperation from sqlmesh.core.snapshot.execution_tracker import QueryExecutionTracker -from sqlmesh.utils import ( - CorrelationId, - columns_to_types_all_known, - random_id, - get_source_columns_to_types, -) -from sqlmesh.utils.connection_pool import ConnectionPool, create_connection_pool +from sqlmesh.utils import (CorrelationId, columns_to_types_all_known, + get_source_columns_to_types, random_id) +from sqlmesh.utils.connection_pool import (ConnectionPool, + create_connection_pool) from sqlmesh.utils.date import TimeLike, make_inclusive, to_time_column -from sqlmesh.utils.errors import ( - MissingDefaultCatalogError, - SQLMeshError, - UnsupportedCatalogOperationError, -) +from sqlmesh.utils.errors import (MissingDefaultCatalogError, SQLMeshError, + UnsupportedCatalogOperationError) from sqlmesh.utils.pandas import columns_to_types_from_df if t.TYPE_CHECKING: import pandas as pd from sqlmesh.core._typing import SchemaName, SessionProperties, TableName - from sqlmesh.core.engine_adapter._typing import ( - DF, - BigframeSession, - GrantsConfig, - PySparkDataFrame, - PySparkSession, - Query, - QueryOrDF, - SnowparkSession, - ) + from sqlmesh.core.engine_adapter._typing import (DF, BigframeSession, + GrantsConfig, + PySparkDataFrame, + PySparkSession, Query, + QueryOrDF, + SnowparkSession) from sqlmesh.core.node import IntervalUnit logger = logging.getLogger(__name__) @@ -173,7 +156,9 @@ def __init__( def with_settings(self, **kwargs: t.Any) -> EngineAdapter: extra_kwargs = { "null_connection": True, - "execute_log_level": kwargs.pop("execute_log_level", self._execute_log_level), + "execute_log_level": kwargs.pop( + "execute_log_level", self._execute_log_level + ), "correlation_id": kwargs.pop("correlation_id", self.correlation_id), "query_execution_tracker": kwargs.pop( "query_execution_tracker", self._query_execution_tracker @@ -272,9 +257,11 @@ def _casted_columns( return [ exp.alias_( exp.cast( - exp.column(column, quoted=True) - if column in source_columns_lookup - else exp.Null(), + ( + exp.column(column, quoted=True) + if column in source_columns_lookup + else exp.Null() + ), to=kind, ), column, @@ -316,13 +303,17 @@ def _get_source_queries( if source_columns: source_columns_lookup = set(source_columns) if not target_columns_to_types: - raise SQLMeshError("columns_to_types must be set if source_columns is set") + raise SQLMeshError( + "columns_to_types must be set if source_columns is set" + ) if not set(target_columns_to_types).issubset(source_columns_lookup): select_columns = [ - exp.column(c, quoted=True) - if c in source_columns_lookup - else exp.cast(exp.Null(), target_columns_to_types[c], copy=False).as_( - c, copy=False, quoted=True + ( + exp.column(c, quoted=True) + if c in source_columns_lookup + else exp.cast( + exp.Null(), target_columns_to_types[c], copy=False + ).as_(c, copy=False, quoted=True) ) for c in target_columns_to_types ] @@ -432,7 +423,9 @@ def _columns_to_types( import pandas as pd if not target_columns_to_types and isinstance(query_or_df, pd.DataFrame): - target_columns_to_types = columns_to_types_from_df(t.cast(pd.DataFrame, query_or_df)) + target_columns_to_types = columns_to_types_from_df( + t.cast(pd.DataFrame, query_or_df) + ) if not source_columns and target_columns_to_types: source_columns = list(target_columns_to_types) # source columns should only contain columns that are defined in the target. If there are extras then @@ -513,14 +506,18 @@ def replace_query( target_data_object = self.get_data_object(target_table) table_exists = target_data_object is not None - if self.drop_data_object_on_type_mismatch(target_data_object, DataObjectType.TABLE): + if self.drop_data_object_on_type_mismatch( + target_data_object, DataObjectType.TABLE + ): table_exists = False - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - target_table=target_table, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + target_table=target_table, + source_columns=source_columns, + ) ) if not target_columns_to_types and table_exists: target_columns_to_types = self.columns(target_table) @@ -574,7 +571,8 @@ def replace_query( lambda node: ( # type: ignore temp_table # type: ignore if isinstance(node, exp.Table) - and quote_identifiers(node) == quote_identifiers(target_table) + and quote_identifiers(node) + == quote_identifiers(target_table) else node ) ) @@ -706,7 +704,9 @@ def create_managed_table( column_descriptions: Optional column descriptions from model query. kwargs: Optional create table properties. """ - raise NotImplementedError(f"Engine does not support managed tables: {type(self)}") + raise NotImplementedError( + f"Engine does not support managed tables: {type(self)}" + ) def ctas( self, @@ -730,11 +730,13 @@ def ctas( column_descriptions: Optional column descriptions from model query. kwargs: Optional create table properties. """ - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - target_table=table_name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + target_table=table_name, + source_columns=source_columns, + ) ) return self._create_table_from_source_queries( table_name, @@ -877,7 +879,9 @@ def _build_column_defs( column, column_descriptions=column_descriptions, engine_supports_schema_comments=engine_supports_schema_comments, - col_type=None if is_view else kind, # don't include column data type for views + col_type=( + None if is_view else kind + ), # don't include column data type for views ) for column, kind in target_columns_to_types.items() ] @@ -895,7 +899,9 @@ def _build_column_def( kind=col_type, constraints=( self._build_col_comment_exp(col_name, column_descriptions) - if engine_supports_schema_comments and self.comments_enabled and column_descriptions + if engine_supports_schema_comments + and self.comments_enabled + and column_descriptions else None ), ) @@ -943,8 +949,9 @@ def _create_table_from_source_queries( # types, and for evaluation methods like `LogicalReplaceQueryMixin.replace_query()` # calls and SCD Type 2 model calls. schema = None - target_columns_to_types_known = target_columns_to_types and columns_to_types_all_known( + target_columns_to_types_known = ( target_columns_to_types + and columns_to_types_all_known(target_columns_to_types) ) if ( column_descriptions @@ -1014,7 +1021,8 @@ def _create_table( target_columns_to_types=target_columns_to_types, table_description=( table_description - if self.COMMENT_CREATION_TABLE.supports_schema_def and self.comments_enabled + if self.COMMENT_CREATION_TABLE.supports_schema_def + and self.comments_enabled else None ), table_kind=table_kind, @@ -1084,7 +1092,9 @@ def create_table_like( target_table_name: The name of the table to create. Can be fully qualified or just table name. source_table_name: The name of the table to base the new table on. """ - self._create_table_like(target_table_name, source_table_name, exists=exists, **kwargs) + self._create_table_like( + target_table_name, source_table_name, exists=exists, **kwargs + ) self._clear_data_object_cache(target_table_name) def clone_table( @@ -1123,7 +1133,9 @@ def clone_table( ) self._clear_data_object_cache(target_table_name) - def drop_data_object(self, data_object: DataObject, ignore_if_not_exists: bool = True) -> None: + def drop_data_object( + self, data_object: DataObject, ignore_if_not_exists: bool = True + ) -> None: """Drops a data object of arbitrary type. Args: @@ -1131,10 +1143,14 @@ def drop_data_object(self, data_object: DataObject, ignore_if_not_exists: bool = ignore_if_not_exists: If True, no error will be raised if the data object does not exist. """ if data_object.type.is_view: - self.drop_view(data_object.to_table(), ignore_if_not_exists=ignore_if_not_exists) + self.drop_view( + data_object.to_table(), ignore_if_not_exists=ignore_if_not_exists + ) elif data_object.type.is_materialized_view: self.drop_view( - data_object.to_table(), ignore_if_not_exists=ignore_if_not_exists, materialized=True + data_object.to_table(), + ignore_if_not_exists=ignore_if_not_exists, + materialized=True, ) elif data_object.type.is_table: self.drop_table(data_object.to_table(), exists=ignore_if_not_exists) @@ -1145,7 +1161,9 @@ def drop_data_object(self, data_object: DataObject, ignore_if_not_exists: bool = f"Can't drop data object '{data_object.to_table().sql(dialect=self.dialect)}' of type '{data_object.type.value}'" ) - def drop_table(self, table_name: TableName, exists: bool = True, **kwargs: t.Any) -> None: + def drop_table( + self, table_name: TableName, exists: bool = True, **kwargs: t.Any + ) -> None: """Drops a table. Args: @@ -1161,7 +1179,9 @@ def drop_managed_table(self, table_name: TableName, exists: bool = True) -> None table_name: The name of the table to drop. exists: If exists, defaults to True. """ - raise NotImplementedError(f"Engine does not support managed tables: {type(self)}") + raise NotImplementedError( + f"Engine does not support managed tables: {type(self)}" + ) def _drop_object( self, @@ -1186,7 +1206,9 @@ def _drop_object( if cascade and kind.upper() in self.SUPPORTED_DROP_CASCADE_OBJECT_KINDS: drop_args["cascade"] = cascade - self.execute(exp.Drop(this=exp.to_table(name), kind=kind, exists=exists, **drop_args)) + self.execute( + exp.Drop(this=exp.to_table(name), kind=kind, exists=exists, **drop_args) + ) self._clear_data_object_cache(name) def get_alter_operations( @@ -1220,7 +1242,8 @@ def alter_table( """ with self.transaction(): for alter_expression in [ - x.expression if isinstance(x, TableAlterOperation) else x for x in alter_expressions + x.expression if isinstance(x, TableAlterOperation) else x + for x in alter_expressions ]: self.execute(alter_expression) @@ -1258,7 +1281,9 @@ def create_view( import pandas as pd if materialized_properties and not materialized: - raise SQLMeshError("Materialized properties are only supported for materialized views") + raise SQLMeshError( + "Materialized properties are only supported for materialized views" + ) query_or_df = self._native_df_to_pandas_df(query_or_df) @@ -1281,12 +1306,14 @@ def create_view( batch_end=len(values), ) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - batch_size=0, - target_table=view_name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + batch_size=0, + target_table=view_name, + source_columns=source_columns, + ) ) if len(source_queries) != 1: raise SQLMeshError("Only one source query is supported for creating views") @@ -1313,7 +1340,9 @@ def create_view( if materialized and self.SUPPORTS_MATERIALIZED_VIEWS: properties.append("expressions", exp.MaterializedProperty()) - if not self.SUPPORTS_MATERIALIZED_VIEW_SCHEMA and isinstance(schema, exp.Schema): + if not self.SUPPORTS_MATERIALIZED_VIEW_SCHEMA and isinstance( + schema, exp.Schema + ): schema = schema.this if not self.SUPPORTS_VIEW_SCHEMA and isinstance(schema, exp.Schema): @@ -1331,7 +1360,9 @@ def create_view( ) is not None ): - materialized_properties["catalog_name"] = exp.to_table(view_name).catalog + materialized_properties["catalog_name"] = exp.to_table( + view_name + ).catalog properties.append("expressions", partitioned_by_prop) if ( clustered_by @@ -1348,7 +1379,8 @@ def create_view( view_properties, ( table_description - if self.COMMENT_CREATION_VIEW.supports_schema_def and self.comments_enabled + if self.COMMENT_CREATION_VIEW.supports_schema_def + and self.comments_enabled else None ), physical_cluster=create_kwargs.pop("physical_cluster", None), @@ -1357,7 +1389,9 @@ def create_view( for view_property in create_view_properties.expressions: # Small hack to make sure SECURE goes at the beginning before materialized as required by Snowflake if isinstance(view_property, exp.SecureProperty): - properties.set("expressions", view_property, index=0, overwrite=False) + properties.set( + "expressions", view_property, index=0, overwrite=False + ) else: properties.append("expressions", view_property) @@ -1367,7 +1401,11 @@ def create_view( if replace: self.drop_data_object_on_type_mismatch( self.get_data_object(view_name), - DataObjectType.VIEW if not materialized else DataObjectType.MATERIALIZED_VIEW, + ( + DataObjectType.VIEW + if not materialized + else DataObjectType.MATERIALIZED_VIEW + ), ) with source_queries[0] as query: @@ -1405,7 +1443,9 @@ def create_view( ) and self.comments_enabled ): - self._create_column_comments(view_name, column_descriptions, "VIEW", materialized) + self._create_column_comments( + view_name, column_descriptions, "VIEW", materialized + ) @set_catalog() def create_schema( @@ -1481,7 +1521,9 @@ def drop_view( ) def create_catalog(self, catalog_name: str | exp.Identifier) -> None: - return self._create_catalog(exp.parse_identifier(catalog_name, dialect=self.dialect)) + return self._create_catalog( + exp.parse_identifier(catalog_name, dialect=self.dialect) + ) def _create_catalog(self, catalog_name: exp.Identifier) -> None: raise SQLMeshError( @@ -1489,7 +1531,9 @@ def _create_catalog(self, catalog_name: exp.Identifier) -> None: ) def drop_catalog(self, catalog_name: str | exp.Identifier) -> None: - return self._drop_catalog(exp.parse_identifier(catalog_name, dialect=self.dialect)) + return self._drop_catalog( + exp.parse_identifier(catalog_name, dialect=self.dialect) + ) def _drop_catalog(self, catalog_name: exp.Identifier) -> None: raise SQLMeshError( @@ -1504,17 +1548,24 @@ def columns( describe_output = self.cursor.fetchall() return { # Note: MySQL returns the column type as bytes. - column_name: exp.DataType.build(_decoded_str(column_type), dialect=self.dialect) + column_name: exp.DataType.build( + _decoded_str(column_type), dialect=self.dialect + ) for column_name, column_type, *_ in itertools.takewhile( lambda t: not t[0].startswith("#"), describe_output, ) - if column_name and column_name.strip() and column_type and column_type.strip() + if column_name + and column_name.strip() + and column_type + and column_type.strip() } def table_exists(self, table_name: TableName) -> bool: table = exp.to_table(table_name) - data_object_cache_key = _get_data_object_cache_key(table.catalog, table.db, table.name) + data_object_cache_key = _get_data_object_cache_key( + table.catalog, table.db, table.name + ) if data_object_cache_key in self._data_object_cache: logger.debug("Table existence cache hit: %s", data_object_cache_key) return self._data_object_cache[data_object_cache_key] is not None @@ -1536,11 +1587,13 @@ def insert_append( track_rows_processed: bool = True, source_columns: t.Optional[t.List[str]] = None, ) -> None: - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - target_table=table_name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + target_table=table_name, + source_columns=source_columns, + ) ) self._insert_append_source_queries( table_name, source_queries, target_columns_to_types, track_rows_processed @@ -1554,7 +1607,9 @@ def _insert_append_source_queries( track_rows_processed: bool = True, ) -> None: with self.transaction(condition=len(source_queries) > 0): - target_columns_to_types = target_columns_to_types or self.columns(table_name) + target_columns_to_types = target_columns_to_types or self.columns( + table_name + ) for source_query in source_queries: with source_query as query: self._insert_append_query( @@ -1589,14 +1644,18 @@ def insert_overwrite_by_partition( ) -> None: if self.INSERT_OVERWRITE_STRATEGY.is_insert_overwrite: target_table = exp.to_table(table_name) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - target_table=target_table, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + target_table=target_table, + source_columns=source_columns, + ) ) self._insert_overwrite_by_condition( - table_name, source_queries, target_columns_to_types=target_columns_to_types + table_name, + source_queries, + target_columns_to_types=target_columns_to_types, ) else: self._replace_by_key( @@ -1614,19 +1673,25 @@ def insert_overwrite_by_time_partition( query_or_df: QueryOrDF, start: TimeLike, end: TimeLike, - time_formatter: t.Callable[[TimeLike, t.Optional[t.Dict[str, exp.DataType]]], exp.Expr], + time_formatter: t.Callable[ + [TimeLike, t.Optional[t.Dict[str, exp.DataType]]], exp.Expr + ], time_column: TimeColumn | exp.Expr | str, target_columns_to_types: t.Optional[t.Dict[str, exp.DataType]] = None, source_columns: t.Optional[t.List[str]] = None, **kwargs: t.Any, ) -> None: - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - target_table=table_name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + target_table=table_name, + source_columns=source_columns, + ) ) - if not target_columns_to_types or not columns_to_types_all_known(target_columns_to_types): + if not target_columns_to_types or not columns_to_types_all_known( + target_columns_to_types + ): target_columns_to_types = self.columns(table_name) low, high = [ time_formatter(dt, target_columns_to_types) @@ -1635,7 +1700,11 @@ def insert_overwrite_by_time_partition( if isinstance(time_column, TimeColumn): time_column = time_column.column where = exp.Between( - this=exp.to_column(time_column) if isinstance(time_column, str) else time_column, + this=( + exp.to_column(time_column) + if isinstance(time_column, str) + else time_column + ), low=low, high=high, ) @@ -1687,9 +1756,12 @@ def _insert_overwrite_by_condition( insert_overwrite_strategy_override or self.INSERT_OVERWRITE_STRATEGY ) with self.transaction( - condition=len(source_queries) > 0 or insert_overwrite_strategy.is_delete_insert + condition=len(source_queries) > 0 + or insert_overwrite_strategy.is_delete_insert ): - target_columns_to_types = target_columns_to_types or self.columns(table_name) + target_columns_to_types = target_columns_to_types or self.columns( + table_name + ) for i, source_query in enumerate(source_queries): with source_query as query: query = self._order_projections_and_filter( @@ -1725,7 +1797,10 @@ def _insert_overwrite_by_condition( query=query, on=exp.false(), whens=exp.Whens( - expressions=[when_not_matched_by_source, when_not_matched_by_target] + expressions=[ + when_not_matched_by_source, + when_not_matched_by_target, + ] ), ) else: @@ -1758,12 +1833,15 @@ def _merge( on: exp.Expr, whens: exp.Whens, ) -> None: - this = exp.alias_(exp.to_table(target_table), alias=MERGE_TARGET_ALIAS, table=True) + this = exp.alias_( + exp.to_table(target_table), alias=MERGE_TARGET_ALIAS, table=True + ) using = exp.alias_( exp.Subquery(this=query), alias=MERGE_SOURCE_ALIAS, copy=False, table=True ) self.execute( - exp.Merge(this=this, using=using, on=on, whens=whens), track_rows_processed=True + exp.Merge(this=this, using=using, on=on, whens=whens), + track_rows_processed=True, ) def scd_type_2_by_time( @@ -1862,7 +1940,9 @@ def remove_managed_columns( cols_to_types: t.Dict[str, exp.DataType], ) -> t.Dict[str, exp.DataType]: return { - k: v for k, v in cols_to_types.items() if k not in {valid_from_name, valid_to_name} + k: v + for k, v in cols_to_types.items() + if k not in {valid_from_name, valid_to_name} } valid_from_name = valid_from_col.name @@ -1875,20 +1955,27 @@ def remove_managed_columns( ): target_columns_to_types = self.columns(target_table) unmanaged_columns_to_types = ( - remove_managed_columns(target_columns_to_types) if target_columns_to_types else None + remove_managed_columns(target_columns_to_types) + if target_columns_to_types + else None ) - source_queries, unmanaged_columns_to_types = self._get_source_queries_and_columns_to_types( - source_table, - unmanaged_columns_to_types, - target_table=target_table, - batch_size=0, - source_columns=source_columns, + source_queries, unmanaged_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + source_table, + unmanaged_columns_to_types, + target_table=target_table, + batch_size=0, + source_columns=source_columns, + ) ) updated_at_name = updated_at_col.name if updated_at_col else None if not target_columns_to_types: - raise SQLMeshError(f"Could not get columns_to_types. Does {target_table} exist?") - unmanaged_columns_to_types = unmanaged_columns_to_types or remove_managed_columns( - target_columns_to_types + raise SQLMeshError( + f"Could not get columns_to_types. Does {target_table} exist?" + ) + unmanaged_columns_to_types = ( + unmanaged_columns_to_types + or remove_managed_columns(target_columns_to_types) ) if not unique_key: raise SQLMeshError("unique_key must be provided for SCD Type 2") @@ -1930,7 +2017,9 @@ def remove_managed_columns( execution_ts = ( exp.cast(execution_time, time_data_type, dialect=self.dialect) if isinstance(execution_time, exp.Column) - else to_time_column(execution_time, time_data_type, self.dialect, nullable=True) + else to_time_column( + execution_time, time_data_type, self.dialect, nullable=True + ) ) if updated_at_as_valid_from: if not updated_at_col: @@ -1950,7 +2039,9 @@ def remove_managed_columns( insert_valid_from_start = execution_ts if check_columns else updated_at_col # type: ignore # joined._exists IS NULL is saying "if the row is deleted" delete_check = ( - exp.column("_exists", "joined").is_(exp.Null()) if invalidate_hard_deletes else None + exp.column("_exists", "joined").is_(exp.Null()) + if invalidate_hard_deletes + else None ) prefixed_valid_to_col = valid_to_col.copy() prefixed_valid_to_col.this.set("this", f"t_{prefixed_valid_to_col.name}") @@ -1969,8 +2060,12 @@ def remove_managed_columns( row_check_conditions.extend( [ col_qualified.neq(t_col), - exp.and_(t_col.is_(exp.Null()), col_qualified.is_(exp.Null()).not_()), - exp.and_(t_col.is_(exp.Null()).not_(), col_qualified.is_(exp.Null())), + exp.and_( + t_col.is_(exp.Null()), col_qualified.is_(exp.Null()).not_() + ), + exp.and_( + t_col.is_(exp.Null()).not_(), col_qualified.is_(exp.Null()) + ), ] ) row_value_check = exp.or_(*row_check_conditions) @@ -2012,7 +2107,9 @@ def remove_managed_columns( updated_at_col_qualified = updated_at_col.copy() updated_at_col_qualified.set("table", exp.to_identifier("joined")) prefixed_updated_at_col = updated_at_col_qualified.copy() - prefixed_updated_at_col.this.set("this", f"t_{updated_at_col_qualified.name}") + prefixed_updated_at_col.this.set( + "this", f"t_{updated_at_col_qualified.name}" + ) updated_row_filter = updated_at_col_qualified > prefixed_updated_at_col valid_to_case_stmt_builder = exp.Case().when( @@ -2022,9 +2119,9 @@ def remove_managed_columns( valid_to_case_stmt_builder = valid_to_case_stmt_builder.when( delete_check, execution_ts ) - valid_to_case_stmt = valid_to_case_stmt_builder.else_(prefixed_valid_to_col).as_( - valid_to_col.this - ) + valid_to_case_stmt = valid_to_case_stmt_builder.else_( + prefixed_valid_to_col + ).as_(valid_to_col.this) valid_from_case_stmt = ( exp.Case() @@ -2035,7 +2132,8 @@ def remove_managed_columns( ), exp.Case() .when( - exp.column(valid_to_col.this, "latest_deleted") > updated_at_col, + exp.column(valid_to_col.this, "latest_deleted") + > updated_at_col, exp.column(valid_to_col.this, "latest_deleted"), ) .else_(updated_at_col), @@ -2044,9 +2142,9 @@ def remove_managed_columns( .else_(prefixed_valid_from_col) ).as_(valid_from_col.this) - existing_rows_query = exp.select(*table_columns, exp.true().as_("_exists")).from_( - target_table - ) + existing_rows_query = exp.select( + *table_columns, exp.true().as_("_exists") + ).from_(target_table) if truncate: existing_rows_query = existing_rows_query.limit(0) @@ -2096,7 +2194,9 @@ def remove_managed_columns( # Deleted records which can be used to determine `valid_from` for undeleted source records .with_( "deleted", - exp.select(*[exp.column(col, "static") for col in target_columns_to_types]) + exp.select( + *[exp.column(col, "static") for col in target_columns_to_types] + ) .from_("static") .join( "latest", @@ -2130,7 +2230,9 @@ def remove_managed_columns( exp.select( exp.column("_exists", table="source").as_("_exists"), *( - exp.column(col, table="latest").as_(prefixed_columns_to_types[i].this) + exp.column(col, table="latest").as_( + prefixed_columns_to_types[i].this + ) for i, col in enumerate(target_columns_to_types) ), *( @@ -2168,7 +2270,9 @@ def remove_managed_columns( "source", on=exp.and_( *[ - add_table(key, "latest").eq(add_table(key, "source")) + add_table(key, "latest").eq( + add_table(key, "source") + ) for key in unique_key ] ), @@ -2185,7 +2289,9 @@ def remove_managed_columns( *( exp.func( "COALESCE", - exp.column(prefixed_unmanaged_columns[i].this, table="joined"), + exp.column( + prefixed_unmanaged_columns[i].this, table="joined" + ), exp.column(col, table="joined"), ).as_(col) for i, col in enumerate(unmanaged_columns_to_types) @@ -2213,9 +2319,9 @@ def remove_managed_columns( exp.select( *unmanaged_columns_to_types, insert_valid_from_start.as_(valid_from_col.this), # type: ignore - to_time_column(exp.null(), time_data_type, self.dialect, nullable=True).as_( - valid_to_col.this - ), + to_time_column( + exp.null(), time_data_type, self.dialect, nullable=True + ).as_(valid_to_col.this), ) .from_("joined") .where(updated_row_filter), @@ -2242,16 +2348,20 @@ def merge( source_columns: t.Optional[t.List[str]] = None, **kwargs: t.Any, ) -> None: - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - source_table, - target_columns_to_types, - target_table=target_table, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + source_table, + target_columns_to_types, + target_table=target_table, + source_columns=source_columns, + ) ) target_columns_to_types = target_columns_to_types or self.columns(target_table) on = exp.and_( *( - add_table(part, MERGE_TARGET_ALIAS).eq(add_table(part, MERGE_SOURCE_ALIAS)) + add_table(part, MERGE_TARGET_ALIAS).eq( + add_table(part, MERGE_SOURCE_ALIAS) + ) for part in unique_key ) ) @@ -2286,7 +2396,8 @@ def merge( ), expression=exp.Tuple( expressions=[ - exp.column(col, MERGE_SOURCE_ALIAS) for col in target_columns_to_types + exp.column(col, MERGE_SOURCE_ALIAS) + for col in target_columns_to_types ] ), ), @@ -2375,7 +2486,9 @@ def get_data_objects( object_names_list = list(missing_names) batches = [ object_names_list[i : i + self.DATA_OBJECT_FILTER_BATCH_SIZE] - for i in range(0, len(object_names_list), self.DATA_OBJECT_FILTER_BATCH_SIZE) + for i in range( + 0, len(object_names_list), self.DATA_OBJECT_FILTER_BATCH_SIZE + ) ] fetched_objects = [] @@ -2405,7 +2518,9 @@ def get_data_objects( fetched_objects = self._get_data_objects(schema_name) if safe_to_cache: for obj in fetched_objects: - cache_key = _get_data_object_cache_key(obj.catalog, obj.schema_name, obj.name) + cache_key = _get_data_object_cache_key( + obj.catalog, obj.schema_name, obj.name + ) self._data_object_cache[cache_key] = obj return fetched_objects @@ -2477,7 +2592,9 @@ def fetch_pyspark_df( self, query: t.Union[exp.Expr, str], quote_identifiers: bool = False ) -> PySparkDataFrame: """Fetches a PySpark DataFrame from the cursor""" - raise NotImplementedError(f"Engine does not support PySpark DataFrames: {type(self)}") + raise NotImplementedError( + f"Engine does not support PySpark DataFrames: {type(self)}" + ) @property def wap_enabled(self) -> bool: @@ -2540,8 +2657,12 @@ def sync_grants_config( raise NotImplementedError(f"Engine does not support grants: {type(self)}") current_grants = self._get_current_grants_config(table) - new_grants, revoked_grants = self._diff_grants_configs(grants_config, current_grants) - revoke_exprs = self._revoke_grants_config_expr(table, revoked_grants, table_type) + new_grants, revoked_grants = self._diff_grants_configs( + grants_config, current_grants + ) + revoke_exprs = self._revoke_grants_config_expr( + table, revoked_grants, table_type + ) grant_exprs = self._apply_grants_config_expr(table, new_grants, table_type) dcl_exprs = revoke_exprs + grant_exprs @@ -2612,7 +2733,9 @@ def execute( ) -> None: """Execute a sql query.""" to_sql_kwargs = ( - {"unsupported_level": ErrorLevel.IGNORE} if ignore_unsupported_errors else {} + {"unsupported_level": ErrorLevel.IGNORE} + if ignore_unsupported_errors + else {} ) with self.transaction(): for e in ensure_list(expressions): @@ -2655,12 +2778,19 @@ def _log_sql( logger.log(self._execute_log_level, "Executing SQL: %s", sql_to_log) def _record_execution_stats( - self, sql: str, rowcount: t.Optional[int] = None, bytes_processed: t.Optional[int] = None + self, + sql: str, + rowcount: t.Optional[int] = None, + bytes_processed: t.Optional[int] = None, ) -> None: if self._query_execution_tracker: - self._query_execution_tracker.record_execution(sql, rowcount, bytes_processed) + self._query_execution_tracker.record_execution( + sql, rowcount, bytes_processed + ) - def _execute(self, sql: str, track_rows_processed: bool = False, **kwargs: t.Any) -> None: + def _execute( + self, sql: str, track_rows_processed: bool = False, **kwargs: t.Any + ) -> None: self.cursor.execute(sql, **kwargs) if ( @@ -2700,14 +2830,21 @@ def temp_table( """ name = exp.to_table(name) # ensure that we use default catalog if none is not specified - if isinstance(name, exp.Table) and not name.catalog and name.db and self.default_catalog: + if ( + isinstance(name, exp.Table) + and not name.catalog + and name.db + and self.default_catalog + ): name.set("catalog", exp.parse_identifier(self.default_catalog)) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types=target_columns_to_types, - target_table=name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types=target_columns_to_types, + target_table=name, + source_columns=source_columns, + ) ) with self.transaction(): @@ -2808,7 +2945,9 @@ def _build_table_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) @@ -2832,7 +2971,9 @@ def _build_view_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) @@ -2870,7 +3011,9 @@ def _to_sql(self, expression: exp.Expr, quote: bool = True, **kwargs: t.Any) -> return expression.sql(**sql_gen_kwargs, copy=False) # type: ignore - def _clear_data_object_cache(self, table_name: t.Optional[TableName] = None) -> None: + def _clear_data_object_cache( + self, table_name: t.Optional[TableName] = None + ) -> None: """Clears the cache entry for the given table name, or clears the entire cache if table_name is None.""" if table_name is None: logger.debug("Clearing entire data object cache") @@ -2897,7 +3040,10 @@ def _get_temp_table( """ table = t.cast(exp.Table, exp.to_table(table).copy()) table.set( - "this", exp.to_identifier(f"__temp_{table.name}_{random_id(short=True)}", quoted=quoted) + "this", + exp.to_identifier( + f"__temp_{table.name}_{random_id(short=True)}", quoted=quoted + ), ) if table_only: @@ -2914,7 +3060,9 @@ def _order_projections_and_filter( coerce_types: bool = False, ) -> Query: if not isinstance(query, exp.Query) or ( - not where and not coerce_types and query.named_selects == list(target_columns_to_types) + not where + and not coerce_types + and query.named_selects == list(target_columns_to_types) ): return query @@ -2930,7 +3078,9 @@ def _order_projections_and_filter( for i, (col, col_tpe) in enumerate(target_columns_to_types.items()) ] - query = exp.select(*select_exprs).from_(query.subquery("_subquery", copy=False), copy=False) + query = exp.select(*select_exprs).from_( + query.subquery("_subquery", copy=False), copy=False + ) if where: query = query.where(where, copy=False) @@ -2998,7 +3148,9 @@ def _replace_by_key( try: delete_query = exp.select(key_exp).from_(temp_table) - insert_query = self._select_columns(target_columns_to_types).from_(temp_table) + insert_query = self._select_columns(target_columns_to_types).from_( + temp_table + ) if not is_unique_key: delete_query = delete_query.distinct() else: @@ -3036,7 +3188,9 @@ def _create_table_comment( table = exp.to_table(table_name) try: - self.execute(self._build_create_comment_table_exp(table, table_comment, table_kind)) + self.execute( + self._build_create_comment_table_exp(table, table_comment, table_kind) + ) except Exception: logger.warning( f"Table comment for '{table.alias_or_name}' not registered - this may be due to limited permissions", @@ -3044,12 +3198,18 @@ def _create_table_comment( ) def _build_create_comment_column_exp( - self, table: exp.Table, column_name: str, column_comment: str, table_kind: str = "TABLE" + self, + table: exp.Table, + column_name: str, + column_comment: str, + table_kind: str = "TABLE", ) -> exp.Comment | str: return exp.Comment( this=exp.column(column_name, *reversed(table.parts)), # type: ignore kind="COLUMN", - expression=exp.Literal.string(self._truncate_column_comment(column_comment)), + expression=exp.Literal.string( + self._truncate_column_comment(column_comment) + ), ) def _create_column_comments( @@ -3063,7 +3223,11 @@ def _create_column_comments( for col, comment in column_comments.items(): try: - self.execute(self._build_create_comment_column_exp(table, col, comment, table_kind)) + self.execute( + self._build_create_comment_column_exp( + table, col, comment, table_kind + ) + ) except Exception: logger.warning( f"Column comments for column '{col}' in table '{table.alias_or_name}' not registered - this may be due to limited permissions", @@ -3077,7 +3241,9 @@ def _create_table_like( exists: bool, **kwargs: t.Any, ) -> None: - self.create_table(target_table_name, self.columns(source_table_name), exists=exists) + self.create_table( + target_table_name, self.columns(source_table_name), exists=exists + ) def _rename_table( self, @@ -3110,9 +3276,11 @@ def _select_columns( ) -> exp.Select: return exp.select( *( - exp.column(c, quoted=True) - if c in (source_columns or columns) - else exp.alias_(exp.Null(), c, quoted=True) + ( + exp.column(c, quoted=True) + if c in (source_columns or columns) + else exp.alias_(exp.Null(), c, quoted=True) + ) for c in columns ) ) @@ -3160,7 +3328,9 @@ def _diff_grants_configs( def _diffs(config1: GrantsConfig, config2: GrantsConfig) -> GrantsConfig: diffs: GrantsConfig = {} - cf_config2 = {k.casefold(): {g.casefold() for g in v} for k, v in config2.items()} + cf_config2 = { + k.casefold(): {g.casefold() for g in v} for k, v in config2.items() + } for key, grantees in config1.items(): cf_key = key.casefold() @@ -3261,7 +3431,9 @@ def _decoded_str(value: t.Union[str, bytes]) -> str: return value -def _get_data_object_cache_key(catalog: t.Optional[str], schema_name: str, object_name: str) -> str: +def _get_data_object_cache_key( + catalog: t.Optional[str], schema_name: str, object_name: str +) -> str: """Returns a cache key for a data object based on its fully qualified name.""" catalog = f"{catalog}." if catalog else "" return f"{catalog}{schema_name}.{object_name}" diff --git a/sqlmesh/core/engine_adapter/base_postgres.py b/sqlmesh/core/engine_adapter/base_postgres.py index e2347b1263..9c8bb171c4 100644 --- a/sqlmesh/core/engine_adapter/base_postgres.py +++ b/sqlmesh/core/engine_adapter/base_postgres.py @@ -1,19 +1,17 @@ from __future__ import annotations -import typing as t import logging +import typing as t from sqlglot import exp from sqlmesh.core.dialect import to_schema -from sqlmesh.core.engine_adapter.base import EngineAdapter, _get_data_object_cache_key -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, -) +from sqlmesh.core.engine_adapter.base import (EngineAdapter, + _get_data_object_cache_key) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType) from sqlmesh.utils.errors import SQLMeshError if t.TYPE_CHECKING: @@ -80,7 +78,9 @@ def table_exists(self, table_name: TableName) -> bool: Reference: https://github.com/aws/amazon-redshift-python-driver/blob/master/redshift_connector/cursor.py#L528-L553 """ table = exp.to_table(table_name) - data_object_cache_key = _get_data_object_cache_key(table.catalog, table.db, table.name) + data_object_cache_key = _get_data_object_cache_key( + table.catalog, table.db, table.name + ) if data_object_cache_key in self._data_object_cache: logger.debug("Table existence cache hit: %s", data_object_cache_key) return self._data_object_cache[data_object_cache_key] is not None diff --git a/sqlmesh/core/engine_adapter/bigquery.py b/sqlmesh/core/engine_adapter/bigquery.py index d136445114..75ef3a2bd9 100644 --- a/sqlmesh/core/engine_adapter/bigquery.py +++ b/sqlmesh/core/engine_adapter/bigquery.py @@ -9,23 +9,17 @@ from sqlmesh.core.dialect import to_schema from sqlmesh.core.engine_adapter.base import _get_data_object_cache_key -from sqlmesh.core.engine_adapter.mixins import ( - ClusteredByMixin, - GrantsFromInfoSchemaMixin, - RowDiffMixin, - TableAlterClusterByOperation, -) -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - DataObject, - DataObjectType, - SourceQuery, - set_catalog, - InsertOverwriteStrategy, -) +from sqlmesh.core.engine_adapter.mixins import (ClusteredByMixin, + GrantsFromInfoSchemaMixin, + RowDiffMixin, + TableAlterClusterByOperation) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, DataObject, + DataObjectType, + InsertOverwriteStrategy, + SourceQuery, set_catalog) from sqlmesh.core.node import IntervalUnit -from sqlmesh.core.schema_diff import TableAlterOperation, NestedSupport -from sqlmesh.utils import optional_import, get_source_columns_to_types +from sqlmesh.core.schema_diff import NestedSupport, TableAlterOperation +from sqlmesh.utils import get_source_columns_to_types, optional_import from sqlmesh.utils.date import to_datetime from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.pandas import columns_to_types_from_dtypes @@ -41,7 +35,8 @@ from google.cloud.bigquery.table import Table as BigQueryTable from sqlmesh.core._typing import SchemaName, SessionProperties, TableName - from sqlmesh.core.engine_adapter._typing import BigframeSession, DCL, DF, GrantsConfig, Query + from sqlmesh.core.engine_adapter._typing import (DCL, DF, BigframeSession, + GrantsConfig, Query) from sqlmesh.core.engine_adapter.base import QueryOrDF @@ -141,7 +136,9 @@ def _job_params(self) -> t.Dict[str, t.Any]: ), } if self._extra_config.get("maximum_bytes_billed") is not None: - params["maximum_bytes_billed"] = self._extra_config.get("maximum_bytes_billed") + params["maximum_bytes_billed"] = self._extra_config.get( + "maximum_bytes_billed" + ) if self._extra_config.get("reservation") is not None: params["reservation"] = self._extra_config.get("reservation") if self.correlation_id: @@ -187,14 +184,18 @@ def query_factory() -> Query: elif not self.table_exists(temp_table): # Make mypy happy assert isinstance(ordered_df, pd.DataFrame) - self._db_call(self.client.create_table, table=temp_bq_table, exists_ok=False) + self._db_call( + self.client.create_table, table=temp_bq_table, exists_ok=False + ) result = self.__load_pandas_to_table( temp_bq_table, ordered_df, source_columns_to_types, replace=False ) if result.errors: raise SQLMeshError(result.errors) return exp.select( - *self._casted_columns(target_columns_to_types, source_columns=source_columns) + *self._casted_columns( + target_columns_to_types, source_columns=source_columns + ) ).from_(temp_table) return [ @@ -257,7 +258,9 @@ def _begin_session(self, properties: SessionProperties) -> None: ) if parsed_query_label: - query_label_str = ",".join([":".join(label) for label in parsed_query_label]) + query_label_str = ",".join( + [":".join(label) for label in parsed_query_label] + ) query = f'SET @@query_label = "{query_label_str}";SELECT 1;' else: query = "SELECT 1;" @@ -302,7 +305,9 @@ def create_schema( warn_on_error=False, ) except Exception as e: - is_already_exists_error = isinstance(e, Conflict) and "Already Exists:" in str(e) + is_already_exists_error = isinstance( + e, Conflict + ) and "Already Exists:" in str(e) if is_already_exists_error and ignore_if_exists: return if not warn_on_error: @@ -341,7 +346,9 @@ def dtype_to_sql( assert struct_type fields = ", ".join( f"{struct_field.name} {dtype_to_sql(struct_field.type, nested_field)}" - for struct_field, nested_field in zip(struct_type.fields, field.fields) + for struct_field, nested_field in zip( + struct_type.fields, field.fields + ) ) return f"STRUCT<{fields}>" if kind.name == "TYPE_KIND_UNSPECIFIED": @@ -361,7 +368,8 @@ def create_mapping_schema( ) -> t.Dict[str, exp.DataType]: return { field.name: exp.DataType.build( - dtype_to_sql(field.to_standard_sql().type, field), dialect=self.dialect + dtype_to_sql(field.to_standard_sql().type, field), + dialect=self.dialect, ) for field in schema } @@ -381,11 +389,15 @@ def create_mapping_schema( if include_pseudo_columns: if bq_table.time_partitioning and not bq_table.time_partitioning.field: - columns["_PARTITIONTIME"] = exp.DataType.build("TIMESTAMP", dialect="bigquery") + columns["_PARTITIONTIME"] = exp.DataType.build( + "TIMESTAMP", dialect="bigquery" + ) if bq_table.time_partitioning.type_ == "DAY": columns["_PARTITIONDATE"] = exp.DataType.build("DATE") if bq_table.table_id.endswith("*"): - columns["_TABLE_SUFFIX"] = exp.DataType.build("STRING", dialect="bigquery") + columns["_TABLE_SUFFIX"] = exp.DataType.build( + "STRING", dialect="bigquery" + ) if ( bq_table.external_data_configuration is not None and bq_table.external_data_configuration.source_format @@ -398,7 +410,9 @@ def create_mapping_schema( "DATASTORE_BACKUP", ) ): - columns["_FILE_NAME"] = exp.DataType.build("STRING", dialect="bigquery") + columns["_FILE_NAME"] = exp.DataType.build( + "STRING", dialect="bigquery" + ) return columns @@ -425,10 +439,14 @@ def alter_table( for op in cluster_by_operations: self._update_clustering_key(op) - nested_fields, non_nested_expressions = self._split_alter_expressions(alter_statements) + nested_fields, non_nested_expressions = self._split_alter_expressions( + alter_statements + ) if nested_fields: - self._update_table_schema_nested_fields(nested_fields, alter_statements[0].this) + self._update_table_schema_nested_fields( + nested_fields, alter_statements[0].this + ) if non_nested_expressions: super().alter_table(non_nested_expressions) @@ -487,10 +505,14 @@ def _split_alter_expressions( and isinstance(action.this, exp.Dot) and isinstance(action.kind, exp.DataType) ): - root_field, *leaf_fields = action.this.this.sql(dialect=self.dialect).split(".") + root_field, *leaf_fields = action.this.this.sql( + dialect=self.dialect + ).split(".") new_field = action.this.expression.sql(dialect=self.dialect) data_type = action.kind.sql(dialect=self.dialect) - nested_fields_to_add[root_field].append((new_field, data_type, leaf_fields)) + nested_fields_to_add[root_field].append( + (new_field, data_type, leaf_fields) + ) else: non_nested_expressions.append(alter_expression) @@ -581,7 +603,9 @@ def __load_pandas_to_table( """ from google.cloud import bigquery - job_config = bigquery.job.LoadJobConfig(schema=self.__get_bq_schema(columns_to_types)) + job_config = bigquery.job.LoadJobConfig( + schema=self.__get_bq_schema(columns_to_types) + ) if replace: job_config.write_disposition = bigquery.WriteDisposition.WRITE_TRUNCATE logger.info(f"Loading dataframe to BigQuery. Table Path: {table.path}") @@ -594,14 +618,19 @@ def __load_pandas_to_table( return result def __db_load_table_from_dataframe( - self, df: pd.DataFrame, table: bigquery.Table, job_config: bigquery.LoadJobConfig + self, + df: pd.DataFrame, + table: bigquery.Table, + job_config: bigquery.LoadJobConfig, ) -> BigQueryQueryResult: job = self.client.load_table_from_dataframe( dataframe=df, destination=table, job_config=job_config ) return self._db_call(job.result) - def __get_bq_schemafield(self, name: str, tpe: exp.DataType) -> bigquery.SchemaField: + def __get_bq_schemafield( + self, name: str, tpe: exp.DataType + ) -> bigquery.SchemaField: from google.cloud import bigquery mode = "NULLABLE" @@ -621,7 +650,9 @@ def __get_bq_schemafield(self, name: str, tpe: exp.DataType) -> bigquery.SchemaF raise ValueError( f"cannot convert unknown type to BQ schema field {inner_field}" ) - fields.append(self.__get_bq_schemafield(name=inner_name, tpe=inner_type)) + fields.append( + self.__get_bq_schemafield(name=inner_name, tpe=inner_type) + ) else: raise ValueError(f"unexpected nested expression {inner_field}") @@ -743,18 +774,26 @@ def insert_overwrite_by_partition( f"DECLARE _sqlmesh_target_partitions_ ARRAY<{partition_type_sql}> DEFAULT ({select_array_agg_partitions});" ) - where = t.cast(exp.Condition, partition_exp).isin(unnest="_sqlmesh_target_partitions_") + where = t.cast(exp.Condition, partition_exp).isin( + unnest="_sqlmesh_target_partitions_" + ) self._insert_overwrite_by_condition( table_name, - [SourceQuery(query_factory=lambda: exp.select("*").from_(temp_table_name))], + [ + SourceQuery( + query_factory=lambda: exp.select("*").from_(temp_table_name) + ) + ], target_columns_to_types, where=where, ) def table_exists(self, table_name: TableName) -> bool: table = exp.to_table(table_name) - data_object_cache_key = _get_data_object_cache_key(table.catalog, table.db, table.name) + data_object_cache_key = _get_data_object_cache_key( + table.catalog, table.db, table.name + ) if data_object_cache_key in self._data_object_cache: logger.debug("Table existence cache hit: %s", data_object_cache_key) return self._data_object_cache[data_object_cache_key] is not None @@ -781,9 +820,7 @@ def get_table_last_modified_ts(self, table_names: t.List[TableName]) -> t.List[i results = [] for dataset, tables in datasets_to_tables.items(): - query = ( - f"SELECT TIMESTAMP_MILLIS(last_modified_time) FROM `{dataset}.__TABLES__` WHERE " - ) + query = f"SELECT TIMESTAMP_MILLIS(last_modified_time) FROM `{dataset}.__TABLES__` WHERE " for i, table_name in enumerate(tables): query += f"TABLE_ID = '{table_name}'" if i < len(tables) - 1: @@ -833,7 +870,9 @@ def _create_column_comments( # Traverse the fields with nested fields down to leaf level for idx, name in enumerate(field_names): - if field := next((field for field in fields if field["name"] == name), None): + if field := next( + (field for field in fields if field["name"] == name), None + ): if idx == last_index: field["description"] = self._truncate_comment( comment, self.MAX_COLUMN_COMMENT_LENGTH @@ -880,7 +919,9 @@ def _build_partitioned_by_exp( and partition_interval_unit is not None and not partition_interval_unit.is_minute ): - column_type: t.Optional[exp.DataType] = (target_columns_to_types or {}).get(this.name) + column_type: t.Optional[exp.DataType] = (target_columns_to_types or {}).get( + this.name + ) if column_type == exp.DataType.build( "date", dialect=self.dialect @@ -931,7 +972,9 @@ def _build_table_properties_exp( ): properties.append(partitioned_by_prop) - if clustered_by and (clustered_by_exp := self._build_clustered_by_exp(clustered_by)): + if clustered_by and ( + clustered_by_exp := self._build_clustered_by_exp(clustered_by) + ): properties.append(clustered_by_exp) if table_description: @@ -941,7 +984,9 @@ def _build_table_properties_exp( ), ) - properties.extend(self._table_or_view_properties_to_expressions(table_properties)) + properties.extend( + self._table_or_view_properties_to_expressions(table_properties) + ) if properties: return exp.Properties(expressions=properties) @@ -975,13 +1020,17 @@ def _build_struct_with_descriptions( else: column = column_def column_expressions.append(column) - return exp.DataType(this=col_type.this, expressions=column_expressions, nested=True) + return exp.DataType( + this=col_type.this, expressions=column_expressions, nested=True + ) # Recursively build column definitions for BigQuery's RECORDs (struct) and REPEATED RECORDs (array of struct) if isinstance(col_type, exp.DataType) and col_type.expressions: expressions = col_type.expressions if col_type.is_type(exp.DataType.Type.STRUCT): - col_type = _build_struct_with_descriptions(col_type, nested_names + [col_name]) + col_type = _build_struct_with_descriptions( + col_type, nested_names + [col_name] + ) elif col_type.is_type(exp.DataType.Type.ARRAY) and expressions[0].is_type( exp.DataType.Type.STRUCT ): @@ -1002,7 +1051,9 @@ def _build_struct_with_descriptions( self._build_col_comment_exp( ".".join(nested_names + [col_name]), column_descriptions ) - if engine_supports_schema_comments and self.comments_enabled and column_descriptions + if engine_supports_schema_comments + and self.comments_enabled + and column_descriptions else None ), ) @@ -1041,7 +1092,9 @@ def _build_view_properties_exp( ), ) - properties.extend(self._table_or_view_properties_to_expressions(view_properties)) + properties.extend( + self._table_or_view_properties_to_expressions(view_properties) + ) if properties: return exp.Properties(expressions=properties) @@ -1055,10 +1108,16 @@ def _build_create_comment_table_exp( truncated_comment = self._truncate_table_comment(table_comment) comment_sql = exp.Literal.string(truncated_comment).sql(dialect=self.dialect) - return f"ALTER {table_kind} {table_sql} SET OPTIONS(description = {comment_sql})" + return ( + f"ALTER {table_kind} {table_sql} SET OPTIONS(description = {comment_sql})" + ) def _build_create_comment_column_exp( - self, table: exp.Table, column_name: str, column_comment: str, table_kind: str = "TABLE" + self, + table: exp.Table, + column_name: str, + column_comment: str, + table_kind: str = "TABLE", ) -> exp.Comment | str: table_sql = table.sql(dialect=self.dialect, identify=True) column_sql = exp.column(column_name).sql(dialect=self.dialect, identify=True) @@ -1079,7 +1138,9 @@ def create_state_table( target_columns_to_types, ) - def _db_call(self, func: t.Callable[..., t.Any], *args: t.Any, **kwargs: t.Any) -> t.Any: + def _db_call( + self, func: t.Callable[..., t.Any], *args: t.Any, **kwargs: t.Any + ) -> t.Any: return func( retry=self.__retry, *args, @@ -1109,7 +1170,9 @@ def _execute( ) # Create job config - job_config = QueryJobConfig(**self._job_params, connection_properties=connection_properties) + job_config = QueryJobConfig( + **self._job_params, connection_properties=connection_properties + ) self._query_job = self._db_call( self.client.query, @@ -1171,10 +1234,17 @@ def _get_data_objects( exp.column("table_name").as_("name"), exp.column("table_schema").as_("schema_name"), exp.case() - .when(exp.column("table_type").eq("BASE TABLE"), exp.Literal.string("TABLE")) + .when( + exp.column("table_type").eq("BASE TABLE"), + exp.Literal.string("TABLE"), + ) .when(exp.column("table_type").eq("CLONE"), exp.Literal.string("TABLE")) - .when(exp.column("table_type").eq("EXTERNAL"), exp.Literal.string("TABLE")) - .when(exp.column("table_type").eq("SNAPSHOT"), exp.Literal.string("TABLE")) + .when( + exp.column("table_type").eq("EXTERNAL"), exp.Literal.string("TABLE") + ) + .when( + exp.column("table_type").eq("SNAPSHOT"), exp.Literal.string("TABLE") + ) .when(exp.column("table_type").eq("VIEW"), exp.Literal.string("VIEW")) .when( exp.column("table_type").eq("MATERIALIZED VIEW"), @@ -1201,12 +1271,15 @@ def _get_data_objects( dialect=self.dialect, ) ) - .where(exp.column("clustering_ordinal_position").is_(exp.not_(exp.null()))) + .where( + exp.column("clustering_ordinal_position").is_(exp.not_(exp.null())) + ) .group_by("1", "2", "3"), ) .from_( exp.to_table( - f"`{catalog}`.`{schema.db}`.INFORMATION_SCHEMA.TABLES", dialect=self.dialect + f"`{catalog}`.`{schema.db}`.INFORMATION_SCHEMA.TABLES", + dialect=self.dialect, ) ) .join( @@ -1243,12 +1316,16 @@ def _update_clustering_key(self, operation: TableAlterClusterByOperation) -> Non cluster_key_expressions = getattr(operation, "cluster_key_expressions", []) bq_table = self._get_table(operation.target_table) - rendered_columns = [c.sql(dialect=self.dialect) for c in cluster_key_expressions] + rendered_columns = [ + c.sql(dialect=self.dialect) for c in cluster_key_expressions + ] bq_table.clustering_fields = ( rendered_columns or None ) # causes a drop of the key if cluster_by is empty or None - self._db_call(self.client.update_table, table=bq_table, fields=["clustering_fields"]) + self._db_call( + self.client.update_table, table=bq_table, fields=["clustering_fields"] + ) if cluster_key_expressions: # BigQuery only applies new clustering going forward, so this rewrites the columns to apply the new clustering to historical data @@ -1297,7 +1374,9 @@ def _columns_to_types( # using dry_run=True attempts to prevent the DataFrame from being materialized just to read the column types from it dtypes = query_or_df.to_pandas(dry_run=True).columnDtypes target_columns_to_types = columns_to_types_from_dtypes(dtypes.items()) - return target_columns_to_types, list(source_columns or target_columns_to_types) + return target_columns_to_types, list( + source_columns or target_columns_to_types + ) return super()._columns_to_types( query_or_df, target_columns_to_types, source_columns=source_columns @@ -1340,7 +1419,9 @@ def _get_current_schema(self) -> str: raise NotImplementedError("BigQuery does not support current schema") def _get_bq_dataset_location(self, project: str, dataset: str) -> str: - return self._db_call(self.client.get_dataset, dataset_ref=f"{project}.{dataset}").location + return self._db_call( + self.client.get_dataset, dataset_ref=f"{project}.{dataset}" + ).location def _get_grant_expression(self, table: exp.Table) -> exp.Expr: if not table.db: @@ -1424,9 +1505,13 @@ def normalize_principal(p: str) -> str: if not principals: continue - noramlized_principals = [exp.Literal.string(normalize_principal(p)) for p in principals] + noramlized_principals = [ + exp.Literal.string(normalize_principal(p)) for p in principals + ] args: t.Dict[str, t.Any] = { - "privileges": [exp.GrantPrivilege(this=exp.to_identifier(privilege, quoted=True))], + "privileges": [ + exp.GrantPrivilege(this=exp.to_identifier(privilege, quoted=True)) + ], "securable": table.copy(), "principals": noramlized_principals, } @@ -1479,7 +1564,9 @@ def should_retry(self, error: BaseException) -> bool: return False self.error_count += 1 if self._is_retryable(error) and self.error_count <= self.num_retries: - logger.info(f"Retry Num {self.error_count} of {self.num_retries}. Error: {repr(error)}") + logger.info( + f"Retry Num {self.error_count} of {self.num_retries}. Error: {repr(error)}" + ) return True return False @@ -1513,7 +1600,9 @@ def select_partitions_expr( data_type = data_type.sql(dialect="bigquery") data_type = data_type.upper() - parse_fun = f"PARSE_{data_type}" if data_type in ("DATE", "DATETIME", "TIMESTAMP") else None + parse_fun = ( + f"PARSE_{data_type}" if data_type in ("DATE", "DATETIME", "TIMESTAMP") else None + ) if parse_fun: granularity = granularity or "day" parse_format = GRANULARITY_TO_PARTITION_FORMAT[granularity.lower()] @@ -1524,7 +1613,9 @@ def select_partitions_expr( dialect="bigquery", ) else: - partition_expr = exp.cast(exp.column("partition_id"), "INT64", dialect="bigquery") + partition_expr = exp.cast( + exp.column("partition_id"), "INT64", dialect="bigquery" + ) return ( exp.select(exp.func(agg_func, partition_expr)) diff --git a/sqlmesh/core/engine_adapter/clickhouse.py b/sqlmesh/core/engine_adapter/clickhouse.py index d1f67e0564..0ef9fd667a 100644 --- a/sqlmesh/core/engine_adapter/clickhouse.py +++ b/sqlmesh/core/engine_adapter/clickhouse.py @@ -1,21 +1,20 @@ from __future__ import annotations -import typing as t import logging import re +import typing as t + from sqlglot import exp, maybe_parse + from sqlmesh.core.dialect import to_schema -from sqlmesh.core.engine_adapter.mixins import LogicalMergeMixin from sqlmesh.core.engine_adapter.base import EngineAdapterWithIndexSupport -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - DataObject, - DataObjectType, - EngineRunMode, - SourceQuery, - CommentCreationView, - InsertOverwriteStrategy, -) +from sqlmesh.core.engine_adapter.mixins import LogicalMergeMixin +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationView, + DataObject, DataObjectType, + EngineRunMode, + InsertOverwriteStrategy, + SourceQuery) from sqlmesh.core.schema_diff import TableAlterOperation from sqlmesh.utils import get_source_columns_to_types @@ -24,7 +23,6 @@ from sqlmesh.core._typing import SchemaName, TableName from sqlmesh.core.engine_adapter._typing import DF, Query, QueryOrDF - from sqlmesh.core.node import IntervalUnit @@ -98,7 +96,11 @@ def _fetch_native_df( ) -> pd.DataFrame: """Fetches a Pandas DataFrame from the cursor""" return self.cursor.client.query_df( - self._to_sql(query, quote=quote_identifiers) if isinstance(query, exp.Expr) else query, + ( + self._to_sql(query, quote=quote_identifiers) + if isinstance(query, exp.Expr) + else query + ), use_extended_dtypes=True, ) @@ -129,11 +131,13 @@ def query_factory() -> Query: ) ordered_df = df[list(source_columns_to_types)] - self.cursor.client.insert_df(temp_table.sql(dialect=self.dialect), df=ordered_df) + self.cursor.client.insert_df( + temp_table.sql(dialect=self.dialect), df=ordered_df + ) - return exp.select(*self._casted_columns(target_columns_to_types, source_columns)).from_( - temp_table - ) + return exp.select( + *self._casted_columns(target_columns_to_types, source_columns) + ).from_(temp_table) return [ SourceQuery( @@ -351,7 +355,10 @@ def _insert_overwrite_by_condition( self.alter_table( [ self._build_alter_partition_exp( - target_table, temp_table, partitions_to_replace, partitions_to_drop + target_table, + temp_table, + partitions_to_replace, + partitions_to_drop, ) ] ) @@ -420,7 +427,9 @@ def _build_alter_partition_exp( "actions", exp.ReplacePartition( expression=exp.Partition( - expressions=[exp.PartitionId(this=exp.Literal.string(str(partition)))] + expressions=[ + exp.PartitionId(this=exp.Literal.string(str(partition))) + ] ), source=temp_table, ), @@ -432,7 +441,9 @@ def _build_alter_partition_exp( exp.DropPartition( expressions=[ exp.Partition( - expressions=[exp.PartitionId(this=exp.Literal.string(str(partition)))] + expressions=[ + exp.PartitionId(this=exp.Literal.string(str(partition))) + ] ) ], source=temp_table, @@ -450,11 +461,13 @@ def _replace_by_key( is_unique_key: bool, source_columns: t.Optional[t.List[str]] = None, ) -> None: - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - source_table, - target_columns_to_types, - target_table=target_table, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + source_table, + target_columns_to_types, + target_table=target_table, + source_columns=source_columns, + ) ) key_exp = ( @@ -481,15 +494,20 @@ def insert_overwrite_by_partition( source_columns: t.Optional[t.List[str]] = None, ) -> None: table_name = self._strip_virtual_catalog(table_name) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - target_table=table_name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + target_table=table_name, + source_columns=source_columns, + ) ) self._insert_overwrite_by_condition( - table_name, source_queries, target_columns_to_types, keep_existing_partition_rows=False + table_name, + source_queries, + target_columns_to_types, + keep_existing_partition_rows=False, ) def _create_table_like( @@ -656,10 +674,15 @@ def _exchange_tables( old_table_name: TableName, new_table_name: TableName, ) -> None: - from clickhouse_connect.driver.exceptions import DatabaseError # type: ignore + from clickhouse_connect.driver.exceptions import \ + DatabaseError # type: ignore - old_table_sql = exp.to_table(old_table_name).sql(dialect=self.dialect, identify=True) - new_table_sql = exp.to_table(new_table_name).sql(dialect=self.dialect, identify=True) + old_table_sql = exp.to_table(old_table_name).sql( + dialect=self.dialect, identify=True + ) + new_table_sql = exp.to_table(new_table_name).sql( + dialect=self.dialect, identify=True + ) try: self.execute( @@ -683,15 +706,23 @@ def _rename_table( old_table_name: TableName, new_table_name: TableName, ) -> None: - old_table_sql = exp.to_table(old_table_name).sql(dialect=self.dialect, identify=True) - new_table_sql = exp.to_table(new_table_name).sql(dialect=self.dialect, identify=True) + old_table_sql = exp.to_table(old_table_name).sql( + dialect=self.dialect, identify=True + ) + new_table_sql = exp.to_table(new_table_name).sql( + dialect=self.dialect, identify=True + ) - self.execute(f"RENAME TABLE {old_table_sql} TO {new_table_sql}{self._on_cluster_sql()}") + self.execute( + f"RENAME TABLE {old_table_sql} TO {new_table_sql}{self._on_cluster_sql()}" + ) def delete_from(self, table_name: TableName, where: t.Union[str, exp.Expr]) -> None: delete_expr = exp.delete(self._strip_virtual_catalog(table_name), where) if self.engine_run_mode.is_cluster: - delete_expr.set("cluster", exp.OnCluster(this=exp.to_identifier(self.cluster))) + delete_expr.set( + "cluster", exp.OnCluster(this=exp.to_identifier(self.cluster)) + ) self.execute(delete_expr) def alter_table( @@ -703,9 +734,12 @@ def alter_table( """ with self.transaction(): for alter_expression in [ - x.expression if isinstance(x, TableAlterOperation) else x for x in alter_expressions + x.expression if isinstance(x, TableAlterOperation) else x + for x in alter_expressions ]: - if self._default_catalog and isinstance(alter_expression.this, exp.Table): + if self._default_catalog and isinstance( + alter_expression.this, exp.Table + ): if alter_expression.this.catalog == self._default_catalog: alter_expression.this.set("catalog", None) if self.engine_run_mode.is_cluster: @@ -737,9 +771,11 @@ def _drop_object( exists=exists, kind=kind, cascade=cascade, - cluster=exp.OnCluster(this=exp.to_identifier(self.cluster)) - if self.engine_run_mode.is_cluster - else None, + cluster=( + exp.OnCluster(this=exp.to_identifier(self.cluster)) + if self.engine_run_mode.is_cluster + else None + ), **drop_args, ) @@ -804,7 +840,8 @@ def use_server_nulls_for_unmatched_after_join( if inject_setting: query.append( - "settings", exp.var(setting_name).eq(exp.Literal.number(setting_value)) + "settings", + exp.var(setting_name).eq(exp.Literal.number(setting_value)), ) return query @@ -820,9 +857,11 @@ def _build_settings_property( expressions=[ exp.EQ( this=exp.var(key.lower()), - expression=value - if isinstance(value, exp.Expr) - else exp.Literal(this=value, is_string=isinstance(value, str)), + expression=( + value + if isinstance(value, exp.Expr) + else exp.Literal(this=value, is_string=isinstance(value, str)) + ), ) for key, value in settings.items() ] @@ -854,10 +893,13 @@ def _build_table_properties_exp( # copy of table_properties so we can pop items off below then consume the rest later table_properties_copy = { - k.upper(): v for k, v in (table_properties.copy() if table_properties else {}).items() + k.upper(): v + for k, v in (table_properties.copy() if table_properties else {}).items() } - mergetree_engine = bool(re.search(self.ORDER_BY_TABLE_ENGINE_REGEX, table_engine)) + mergetree_engine = bool( + re.search(self.ORDER_BY_TABLE_ENGINE_REGEX, table_engine) + ) ordered_by_raw = table_properties_copy.pop("ORDER_BY", None) if mergetree_engine: ordered_by_exprs = [] @@ -871,7 +913,9 @@ def _build_table_properties_exp( if not ordered_by_vals: ordered_by_vals = ( - ordered_by_raw if isinstance(ordered_by_raw, list) else [ordered_by_raw] + ordered_by_raw + if isinstance(ordered_by_raw, list) + else [ordered_by_raw] ) for col in ordered_by_vals: @@ -885,7 +929,9 @@ def _build_table_properties_exp( ) ) - properties.append(exp.Order(expressions=[exp.Tuple(expressions=ordered_by_exprs)])) + properties.append( + exp.Order(expressions=[exp.Tuple(expressions=ordered_by_exprs)]) + ) primary_key = table_properties_copy.pop("PRIMARY_KEY", None) if mergetree_engine and primary_key: @@ -896,7 +942,9 @@ def _build_table_properties_exp( primary_key_vals = [primary_key.this] if not primary_key_vals: - primary_key_vals = primary_key if isinstance(primary_key, list) else [primary_key] + primary_key_vals = ( + primary_key if isinstance(primary_key, list) else [primary_key] + ) properties.append( exp.PrimaryKey( @@ -910,12 +958,15 @@ def _build_table_properties_exp( ttl = table_properties_copy.pop("TTL", None) if ttl: properties.append( - exp.MergeTreeTTL(expressions=[ttl if isinstance(ttl, exp.Expr) else exp.var(ttl)]) + exp.MergeTreeTTL( + expressions=[ttl if isinstance(ttl, exp.Expr) else exp.var(ttl)] + ) ) if ( partitioned_by - and (partitioned_by_prop := self._build_partitioned_by_exp(partitioned_by)) is not None + and (partitioned_by_prop := self._build_partitioned_by_exp(partitioned_by)) + is not None ): properties.append(partitioned_by_prop) @@ -931,7 +982,9 @@ def _build_table_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) @@ -960,7 +1013,9 @@ def _build_view_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) @@ -996,6 +1051,6 @@ def _build_create_comment_column_exp( def _on_cluster_sql(self) -> str: if self.engine_run_mode.is_cluster: - cluster_name = exp.to_identifier(self.cluster, quoted=True).sql(dialect=self.dialect) # type: ignore + cluster_name = exp.to_identifier(self.cluster, quoted=True).sql(dialect=self.dialect) # type: ignore return f" ON CLUSTER {cluster_name} " return "" diff --git a/sqlmesh/core/engine_adapter/databricks.py b/sqlmesh/core/engine_adapter/databricks.py index 098825bc2d..9ce8ac97a5 100644 --- a/sqlmesh/core/engine_adapter/databricks.py +++ b/sqlmesh/core/engine_adapter/databricks.py @@ -9,23 +9,21 @@ from sqlmesh.core.constants import LIQUID_CLUSTERING_KEYWORDS from sqlmesh.core.dialect import to_schema from sqlmesh.core.engine_adapter.mixins import GrantsFromInfoSchemaMixin -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - DataObject, - DataObjectType, - InsertOverwriteStrategy, - SourceQuery, -) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, DataObject, + DataObjectType, + InsertOverwriteStrategy, + SourceQuery) from sqlmesh.core.engine_adapter.spark import SparkEngineAdapter from sqlmesh.core.node import IntervalUnit from sqlmesh.core.schema_diff import NestedSupport -from sqlmesh.engines.spark.db_api.spark_session import connection, SparkSessionConnection -from sqlmesh.utils.errors import SQLMeshError, MissingDefaultCatalogError +from sqlmesh.engines.spark.db_api.spark_session import (SparkSessionConnection, + connection) +from sqlmesh.utils.errors import MissingDefaultCatalogError, SQLMeshError if t.TYPE_CHECKING: import pandas as pd - from sqlmesh.core._typing import SchemaName, TableName, SessionProperties + from sqlmesh.core._typing import SchemaName, SessionProperties, TableName from sqlmesh.core.engine_adapter._typing import DF, PySparkSession, Query logger = logging.getLogger(__name__) @@ -38,7 +36,9 @@ def _query_tags( return None if not isinstance(query_tags, (exp.Map, exp.VarMap)): - raise SQLMeshError("Invalid value for `session_properties.query_tags`. Must be a map.") + raise SQLMeshError( + "Invalid value for `session_properties.query_tags`. Must be a map." + ) keys = query_tags.args.get("keys") values = query_tags.args.get("values") @@ -114,7 +114,9 @@ def can_access_databricks_connect(cls, disable_databricks_connect: bool) -> bool @property def _use_spark_session(self) -> bool: - if self.can_access_spark_session(bool(self._extra_config.get("disable_spark_session"))): + if self.can_access_spark_session( + bool(self._extra_config.get("disable_spark_session")) + ): return True if self.can_access_databricks_connect( @@ -138,8 +140,9 @@ def is_spark_session_connection(self) -> bool: @property def _is_databricks_sql_connector_connection(self) -> bool: - return not self.is_spark_session_connection and not self._connection_pool.get_attribute( - "use_spark_engine_adapter" + return ( + not self.is_spark_session_connection + and not self._connection_pool.get_attribute("use_spark_engine_adapter") ) def _set_spark_engine_adapter_if_needed(self) -> None: @@ -157,11 +160,15 @@ def _set_spark_engine_adapter_if_needed(self) -> None: if self._extra_config.get("databricks_connect_use_serverless"): connect_kwargs["serverless"] = True else: - connect_kwargs["cluster_id"] = self._extra_config["databricks_connect_cluster_id"] + connect_kwargs["cluster_id"] = self._extra_config[ + "databricks_connect_cluster_id" + ] catalog = self._extra_config.get("catalog") spark = ( - DatabricksSession.builder.remote(**connect_kwargs).userAgent("sqlmesh").getOrCreate() + DatabricksSession.builder.remote(**connect_kwargs) + .userAgent("sqlmesh") + .getOrCreate() ) self._spark_engine_adapter = SparkEngineAdapter( partial(connection, spark=spark, catalog=catalog), @@ -225,13 +232,17 @@ def _begin_session(self, properties: SessionProperties) -> t.Any: """Begin a new session.""" # Align the different possible connectors to a single catalog self.set_current_catalog(self.default_catalog) # type: ignore - self._connection_pool.set_attribute("query_tags", _query_tags(properties.get("query_tags"))) + self._connection_pool.set_attribute( + "query_tags", _query_tags(properties.get("query_tags")) + ) def _end_session(self) -> None: self._connection_pool.set_attribute("query_tags", None) self._connection_pool.set_attribute("use_spark_engine_adapter", False) - def _execute(self, sql: str, track_rows_processed: bool = False, **kwargs: t.Any) -> None: + def _execute( + self, sql: str, track_rows_processed: bool = False, **kwargs: t.Any + ) -> None: query_tags = self._connection_pool.get_attribute("query_tags") if ( query_tags @@ -252,7 +263,11 @@ def _df_to_source_queries( ) -> t.List[SourceQuery]: if not self._use_spark_session: return super(SparkEngineAdapter, self)._df_to_source_queries( - df, target_columns_to_types, batch_size, target_table, source_columns=source_columns + df, + target_columns_to_types, + batch_size, + target_table, + source_columns=source_columns, ) pyspark_df = self._ensure_pyspark_df( df, target_columns_to_types, source_columns=source_columns @@ -262,7 +277,9 @@ def query_factory() -> Query: temp_table = self._get_temp_table(target_table or "spark", table_only=True) pyspark_df.createOrReplaceTempView(temp_table.sql(dialect=self.dialect)) self._connection_pool.set_attribute("use_spark_engine_adapter", True) - return exp.select(*self._select_columns(target_columns_to_types)).from_(temp_table) + return exp.select(*self._select_columns(target_columns_to_types)).from_( + temp_table + ) return [SourceQuery(query_factory=query_factory)] @@ -297,7 +314,8 @@ def get_current_catalog(self) -> t.Optional[str]: sql_connector_catalog = None if self._spark_engine_adapter: from py4j.protocol import Py4JError - from pyspark.errors.exceptions.connect import SparkConnectGrpcException + from pyspark.errors.exceptions.connect import \ + SparkConnectGrpcException try: # Note: Spark 3.4+ Only API @@ -318,7 +336,8 @@ def get_current_catalog(self) -> t.Optional[str]: def set_current_catalog(self, catalog_name: str) -> None: def _set_spark_session_current_catalog(spark: PySparkSession) -> None: from py4j.protocol import Py4JError - from pyspark.errors.exceptions.connect import SparkConnectGrpcException + from pyspark.errors.exceptions.connect import \ + SparkConnectGrpcException try: # Note: Spark 3.4+ Only API @@ -352,7 +371,8 @@ def _get_data_objects( exp.case(exp.column("table_type")) .when(exp.Literal.string("VIEW"), exp.Literal.string("view")) .when( - exp.Literal.string("MATERIALIZED_VIEW"), exp.Literal.string("materialized_view") + exp.Literal.string("MATERIALIZED_VIEW"), + exp.Literal.string("materialized_view"), ) .else_(exp.Literal.string("table")) .as_("type"), @@ -445,13 +465,17 @@ def _build_table_properties_exp( if clustered_by: if len(clustered_by) == 1 and isinstance(clustered_by[0], exp.Var): if clustered_by[0].name.upper() not in LIQUID_CLUSTERING_KEYWORDS: - raise ValueError(f"Unexpected bare Var in clustered_by: {clustered_by[0]!r}") + raise ValueError( + f"Unexpected bare Var in clustered_by: {clustered_by[0]!r}" + ) # exp.Cluster with a bare Var generates: CLUSTER BY AUTO (no parens) clustered_by_exp = exp.Cluster(expressions=[clustered_by[0].copy()]) else: # Databricks expects column expressions wrapped in a tuple clustered_by_exp = exp.Cluster( - expressions=[exp.Tuple(expressions=[c.copy() for c in clustered_by])] + expressions=[ + exp.Tuple(expressions=[c.copy() for c in clustered_by]) + ] ) expressions = properties.expressions if properties else [] expressions.append(clustered_by_exp) @@ -497,4 +521,6 @@ def columns( self.execute(query.sql(dialect=self.dialect)) result = self.cursor.fetchall() - return {row[0]: exp.DataType.build(row[1], dialect=self.dialect) for row in result} + return { + row[0]: exp.DataType.build(row[1], dialect=self.dialect) for row in result + } diff --git a/sqlmesh/core/engine_adapter/duckdb.py b/sqlmesh/core/engine_adapter/duckdb.py index ebfcaa7901..cf63b56381 100644 --- a/sqlmesh/core/engine_adapter/duckdb.py +++ b/sqlmesh/core/engine_adapter/duckdb.py @@ -1,31 +1,29 @@ from __future__ import annotations import typing as t -from sqlglot import exp from pathlib import Path +from sqlglot import exp + from sqlmesh.core.engine_adapter.mixins import ( - GetCurrentCatalogFromFunctionMixin, - LogicalMergeMixin, - RowDiffMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, - SourceQuery, - set_catalog, -) + GetCurrentCatalogFromFunctionMixin, LogicalMergeMixin, RowDiffMixin) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + SourceQuery, set_catalog) if t.TYPE_CHECKING: from sqlmesh.core._typing import SchemaName, TableName from sqlmesh.core.engine_adapter._typing import DF -@set_catalog(override_mapping={"_get_data_objects": CatalogSupport.REQUIRES_SET_CATALOG}) -class DuckDBEngineAdapter(LogicalMergeMixin, GetCurrentCatalogFromFunctionMixin, RowDiffMixin): +@set_catalog( + override_mapping={"_get_data_objects": CatalogSupport.REQUIRES_SET_CATALOG} +) +class DuckDBEngineAdapter( + LogicalMergeMixin, GetCurrentCatalogFromFunctionMixin, RowDiffMixin +): DIALECT = "duckdb" SUPPORTS_TRANSACTIONS = False SCHEMA_DIFFER_KWARGS = { @@ -51,12 +49,15 @@ def _create_catalog(self, catalog_name: exp.Identifier) -> None: db_filename = f"{catalog_name.output_name}.db" self.execute( exp.Attach( - this=exp.alias_(exp.Literal.string(db_filename), catalog_name), exists=True + this=exp.alias_(exp.Literal.string(db_filename), catalog_name), + exists=True, ) ) else: self.execute( - exp.Create(this=exp.Table(this=catalog_name), kind="DATABASE", exists=True) + exp.Create( + this=exp.Table(this=catalog_name), kind="DATABASE", exists=True + ) ) def _drop_catalog(self, catalog_name: exp.Identifier) -> None: @@ -68,7 +69,10 @@ def _drop_catalog(self, catalog_name: exp.Identifier) -> None: else: self.execute( exp.Drop( - this=exp.Table(this=catalog_name), kind="DATABASE", cascade=True, exists=True + this=exp.Table(this=catalog_name), + kind="DATABASE", + cascade=True, + exists=True, ) ) @@ -89,7 +93,9 @@ def _df_to_source_queries( self.cursor.sql(f"CREATE TABLE {temp_table} AS {temp_table_sql}") return [ SourceQuery( - query_factory=lambda: self._select_columns(target_columns_to_types).from_( + query_factory=lambda: self._select_columns( + target_columns_to_types + ).from_( temp_table ), # type: ignore cleanup_func=lambda: self.drop_table(temp_table), @@ -129,7 +135,8 @@ def _get_data_objects( ) .from_(exp.to_table("system.information_schema.tables")) .where( - exp.column("table_catalog").eq(catalog), exp.column("table_schema").eq(schema_name) + exp.column("table_catalog").eq(catalog), + exp.column("table_schema").eq(schema_name), ) ) if object_names: @@ -213,7 +220,9 @@ def _create_table( partitioned_by_str = ", ".join( expr.sql(dialect=self.dialect) for expr in partitioned_by_exps ) - self.execute(f"ALTER TABLE {table_name_str} SET PARTITIONED BY ({partitioned_by_str});") + self.execute( + f"ALTER TABLE {table_name_str} SET PARTITIONED BY ({partitioned_by_str});" + ) @property def _is_motherduck(self) -> bool: diff --git a/sqlmesh/core/engine_adapter/fabric.py b/sqlmesh/core/engine_adapter/fabric.py index 7b2f1acd73..a59995ce7b 100644 --- a/sqlmesh/core/engine_adapter/fabric.py +++ b/sqlmesh/core/engine_adapter/fabric.py @@ -1,23 +1,23 @@ from __future__ import annotations -import typing as t import logging -import requests import time +import typing as t from functools import cached_property + +import requests from sqlglot import exp -from tenacity import retry, stop_after_attempt, wait_exponential, retry_if_result +from tenacity import (retry, retry_if_result, stop_after_attempt, + wait_exponential) + from sqlmesh.core.engine_adapter.mssql import MSSQLEngineAdapter -from sqlmesh.core.engine_adapter.shared import ( - CommentCreationTable, - CommentCreationView, - InsertOverwriteStrategy, -) -from sqlmesh.utils.errors import SQLMeshError -from sqlmesh.utils.connection_pool import ConnectionPool +from sqlmesh.core.engine_adapter.shared import (CommentCreationTable, + CommentCreationView, + InsertOverwriteStrategy) from sqlmesh.core.schema_diff import TableAlterOperation from sqlmesh.utils import random_id - +from sqlmesh.utils.connection_pool import ConnectionPool +from sqlmesh.utils.errors import SQLMeshError logger = logging.getLogger(__name__) @@ -38,14 +38,19 @@ class FabricEngineAdapter(MSSQLEngineAdapter): COMMENT_CREATION_VIEW = CommentCreationView.UNSUPPORTED def __init__( - self, connection_factory_or_pool: t.Union[t.Callable, t.Any], *args: t.Any, **kwargs: t.Any + self, + connection_factory_or_pool: t.Union[t.Callable, t.Any], + *args: t.Any, + **kwargs: t.Any, ) -> None: # Wrap connection factory to support changing the catalog dynamically at runtime if not isinstance(connection_factory_or_pool, ConnectionPool): original_connection_factory = connection_factory_or_pool - connection_factory_or_pool = lambda *args, **kwargs: original_connection_factory( - target_catalog=self._target_catalog, *args, **kwargs + connection_factory_or_pool = ( + lambda *args, **kwargs: original_connection_factory( + target_catalog=self._target_catalog, *args, **kwargs + ) ) super().__init__(connection_factory_or_pool, *args, **kwargs) @@ -179,7 +184,9 @@ def set_current_catalog(self, catalog_name: t.Optional[str]) -> None: if self.get_current_catalog() == target_catalog and ( not explicit_default_catalog or connected_catalog is None ): - logger.debug("Already using requested Fabric catalog state, no action needed") + logger.debug( + "Already using requested Fabric catalog state, no action needed" + ) return # Decide whether the open connection needs to be replaced. @@ -237,7 +244,11 @@ def alter_table( # Get the target table from the first expression to determine the correct catalog. first_op = alter_expressions[0] - expression = first_op.expression if isinstance(first_op, TableAlterOperation) else first_op + expression = ( + first_op.expression + if isinstance(first_op, TableAlterOperation) + else first_op + ) if not isinstance(expression, exp.Alter) or not expression.this.catalog: # Fallback for unexpected scenarios logger.warning( @@ -251,7 +262,9 @@ def alter_table( with self.transaction(): for op in alter_expressions: - expression = op.expression if isinstance(op, TableAlterOperation) else op + expression = ( + op.expression if isinstance(op, TableAlterOperation) else op + ) if not isinstance(expression, exp.Alter): self.execute(expression) @@ -263,14 +276,16 @@ def alter_table( table_name_without_catalog = table_name.copy() table_name_without_catalog.set("catalog", None) - is_type_change = isinstance(action, exp.AlterColumn) and action.args.get( - "dtype" - ) + is_type_change = isinstance( + action, exp.AlterColumn + ) and action.args.get("dtype") if is_type_change: column_to_alter = action.this new_type = action.args["dtype"] - temp_column_name_str = f"{column_to_alter.name}__{random_id(short=True)}" + temp_column_name_str = ( + f"{column_to_alter.name}__{random_id(short=True)}" + ) temp_column_name = exp.to_identifier(temp_column_name_str) logger.info( @@ -285,7 +300,9 @@ def alter_table( this=table_name_without_catalog.copy(), kind="TABLE", actions=[ - exp.ColumnDef(this=temp_column_name.copy(), kind=new_type.copy()) + exp.ColumnDef( + this=temp_column_name.copy(), kind=new_type.copy() + ) ], ) add_sql = self._to_sql(add_column_expr) @@ -299,7 +316,8 @@ def alter_table( exp.EQ( this=temp_column_name.copy(), expression=exp.Cast( - this=column_to_alter.copy(), to=new_type.copy() + this=column_to_alter.copy(), + to=new_type.copy(), ), ) ], @@ -312,7 +330,9 @@ def alter_table( exp.Alter( this=table_name_without_catalog.copy(), kind="TABLE", - actions=[exp.Drop(this=column_to_alter.copy(), kind="COLUMN")], + actions=[ + exp.Drop(this=column_to_alter.copy(), kind="COLUMN") + ], ) ) self.execute(drop_sql) @@ -327,13 +347,17 @@ def alter_table( else: # For other alterations, execute directly. direct_alter_expr = exp.Alter( - this=table_name_without_catalog.copy(), kind="TABLE", actions=[action] + this=table_name_without_catalog.copy(), + kind="TABLE", + actions=[action], ) self.execute(direct_alter_expr) class FabricHttpClient: - def __init__(self, tenant_id: str, workspace_id: str, client_id: str, client_secret: str): + def __init__( + self, tenant_id: str, workspace_id: str, client_id: str, client_secret: str + ): self.tenant_id = tenant_id self.client_id = client_id self.client_secret = client_secret @@ -357,7 +381,9 @@ def create_warehouse( "description": f"Warehouse created by SQLMesh: {warehouse_name}", } - response = self.session.post(self._endpoint_url("warehouses"), json=request_data) + response = self.session.post( + self._endpoint_url("warehouses"), json=request_data + ) if ( if_not_exists @@ -368,14 +394,18 @@ def create_warehouse( logger.warning(f"Fabric warehouse {warehouse_name} already exists") return if errorCode == "ItemDisplayNameNotAvailableYet": - logger.warning(f"Fabric warehouse {warehouse_name} is still spinning up; waiting") + logger.warning( + f"Fabric warehouse {warehouse_name} is still spinning up; waiting" + ) # Fabric error message is something like: # - "Requested 'circleci_51d7087e__dev' is not available yet and is expected to become available in the upcoming minutes." # This seems to happen if a catalog is dropped and then a new one with the same name is immediately created. # There appears to be some delayed async process on the Fabric side that actually drops the warehouses and frees up the names to be used again time.sleep(30) return self.create_warehouse( - warehouse_name=warehouse_name, if_not_exists=if_not_exists, attempt=attempt + 1 + warehouse_name=warehouse_name, + if_not_exists=if_not_exists, + attempt=attempt + 1, ) try: @@ -392,12 +422,16 @@ def create_warehouse( logger.info(f"Successfully created Fabric warehouse: {warehouse_name}") return - if response.status_code == 202 and (location_header := response.headers.get("location")): + if response.status_code == 202 and ( + location_header := response.headers.get("location") + ): logger.info(f"Warehouse creation initiated for: {warehouse_name}") self._wait_for_completion(location_header, warehouse_name) logger.info(f"Successfully created Fabric warehouse: {warehouse_name}") else: - logger.error(f"Unexpected response from Fabric API: {response}\n{response.text}") + logger.error( + f"Unexpected response from Fabric API: {response}\n{response.text}" + ) raise SQLMeshError(f"Unable to create warehouse: {response}") def delete_warehouse(self, warehouse_name: str, if_exists: bool = True) -> None: @@ -452,7 +486,9 @@ def _get_access_token(self) -> str: """Get access token using Service Principal authentication.""" # Use Azure AD OAuth2 token endpoint - token_url = f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token" + token_url = ( + f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token" + ) data = { "grant_type": "client_credentials", @@ -489,10 +525,14 @@ def _poll() -> str: elif status in ["InProgress", "Running"]: logger.debug(f"Operation {operation_name} still in progress...") elif status not in ["Succeeded"]: - logger.warning(f"Unknown status '{status}' for operation {operation_name}") + logger.warning( + f"Unknown status '{status}' for operation {operation_name}" + ) return status final_status = _poll() if final_status != "Succeeded": - raise SQLMeshError(f"Operation {operation_name} completed with status: {final_status}") + raise SQLMeshError( + f"Operation {operation_name} completed with status: {final_status}" + ) diff --git a/sqlmesh/core/engine_adapter/mixins.py b/sqlmesh/core/engine_adapter/mixins.py index bf4bb970a2..7e875c2e2e 100644 --- a/sqlmesh/core/engine_adapter/mixins.py +++ b/sqlmesh/core/engine_adapter/mixins.py @@ -9,21 +9,17 @@ from sqlglot.helper import seq_get from sqlglot.optimizer.normalize_identifiers import normalize_identifiers +from sqlmesh.core.dialect import schema_ from sqlmesh.core.engine_adapter.base import EngineAdapter from sqlmesh.core.engine_adapter.shared import DataObjectType from sqlmesh.core.node import IntervalUnit -from sqlmesh.core.dialect import schema_ from sqlmesh.core.schema_diff import TableAlterOperation from sqlmesh.utils.errors import SQLMeshError if t.TYPE_CHECKING: from sqlmesh.core._typing import TableName - from sqlmesh.core.engine_adapter._typing import ( - DCL, - DF, - GrantsConfig, - QueryOrDF, - ) + from sqlmesh.core.engine_adapter._typing import (DCL, DF, GrantsConfig, + QueryOrDF) from sqlmesh.core.engine_adapter.base import QueryOrDF logger = logging.getLogger(__name__) @@ -65,7 +61,11 @@ def _fetch_native_df( from pandas.io.sql import read_sql_query - sql = self._to_sql(query, quote=quote_identifiers) if isinstance(query, exp.Expr) else query + sql = ( + self._to_sql(query, quote=quote_identifiers) + if isinstance(query, exp.Expr) + else query + ) logger.debug(f"Executing SQL:\n{sql}") with catch_warnings(), self.transaction(): filterwarnings( @@ -90,7 +90,8 @@ def _build_partitioned_by_exp( ) -> t.Union[exp.PartitionedByProperty, exp.Property]: if ( self.dialect == "trino" - and self.get_catalog_type(catalog_name or self.get_current_catalog()) == "iceberg" + and self.get_catalog_type(catalog_name or self.get_current_catalog()) + == "iceberg" ): # On the Trino Iceberg catalog, the table property is called "partitioning" - not "partitioned_by" # In addition, partition column transform expressions like `day(col)` or `bucket(col, 5)` are allowed @@ -99,7 +100,10 @@ def _build_partitioned_by_exp( return exp.Property( this=exp.var("PARTITIONING"), value=exp.array( - *(exp.Literal.string(e.sql(dialect=self.dialect)) for e in partitioned_by) + *( + exp.Literal.string(e.sql(dialect=self.dialect)) + for e in partitioned_by + ) ), ) for expr in partitioned_by: @@ -132,7 +136,8 @@ def _build_table_properties_exp( if storage_format: properties.append( exp.Property( - this="write.format.default", value=exp.Literal.string(storage_format) + this="write.format.default", + value=exp.Literal.string(storage_format), ) ) elif storage_format: @@ -150,11 +155,15 @@ def _build_table_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) - properties.extend(self._table_or_view_properties_to_expressions(table_properties)) + properties.extend( + self._table_or_view_properties_to_expressions(table_properties) + ) if properties: return exp.Properties(expressions=properties) @@ -172,11 +181,15 @@ def _build_view_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) - properties.extend(self._table_or_view_properties_to_expressions(view_properties)) + properties.extend( + self._table_or_view_properties_to_expressions(view_properties) + ) if properties: return exp.Properties(expressions=properties) @@ -229,7 +242,9 @@ def _default_precision_to_max( parameter = self.schema_differ.get_type_parameters(col_type) type_default = types_with_max_default_param[col_type.this] if parameter == type_default: - col_type.set("expressions", [exp.DataTypeParam(this=exp.var("max"))]) + col_type.set( + "expressions", [exp.DataTypeParam(this=exp.var("max"))] + ) return columns_to_types @@ -265,7 +280,9 @@ def _build_create_table_exp( # redshift and mssql have a bug where CTAS statements have non determistic types. if a limit # is applied to a ctas statement, VARCHAR types default to 1 in some instances. select_statement = statement.expression.copy() - for select_or_union in select_statement.find_all(exp.Select, exp.SetOperation): + for select_or_union in select_statement.find_all( + exp.Select, exp.SetOperation + ): limit = select_or_union.args.get("limit") if limit is not None and limit.expression.this == "0": limit.pop() @@ -386,7 +403,8 @@ def get_alter_operations( if current_table_info and target_table_info: if target_table_info.is_clustered: if target_table_info.clustering_key and ( - current_table_info.clustering_key != target_table_info.clustering_key + current_table_info.clustering_key + != target_table_info.clustering_key ): operations.append( TableAlterChangeClusterKeyOperation( @@ -396,7 +414,9 @@ def get_alter_operations( ) ) elif current_table_info.is_clustered: - operations.append(TableAlterDropClusterKeyOperation(target_table=current_table)) + operations.append( + TableAlterDropClusterKeyOperation(target_table=current_table) + ) return operations @@ -459,7 +479,10 @@ def concat_columns( exp.func( "COALESCE", self.normalize_value( - exp.to_column(column), type, decimal_precision, timestamp_precision + exp.to_column(column), + type, + decimal_precision, + timestamp_precision, ), exp.Literal.string(""), ) @@ -539,7 +562,10 @@ def _normalize_timestamp_value( expr = exp.TimeToStr(this=expr, format=exp.Literal.string(format)) if digits_to_chop_off > 0: expr = exp.func( - "SUBSTRING", expr, 1, len("2023-01-01 12:13:14.000000") - digits_to_chop_off + "SUBSTRING", + expr, + 1, + len("2023-01-01 12:13:14.000000") - digits_to_chop_off, ) return expr @@ -591,7 +617,9 @@ def _dcl_grants_config_expr( if self.SUPPORTS_MULTIPLE_GRANT_PRINCIPALS: args["principals"] = [ normalize_identifiers( - parse_one(principal, into=exp.GrantPrincipal, dialect=self.dialect), + parse_one( + principal, into=exp.GrantPrincipal, dialect=self.dialect + ), dialect=self.dialect, ) for principal in principals @@ -601,7 +629,9 @@ def _dcl_grants_config_expr( for principal in principals: args["principals"] = [ normalize_identifiers( - parse_one(principal, into=exp.GrantPrincipal, dialect=self.dialect), + parse_one( + principal, into=exp.GrantPrincipal, dialect=self.dialect + ), dialect=self.dialect, ) ] @@ -623,11 +653,14 @@ def _revoke_grants_config_expr( grants_config: GrantsConfig, table_type: DataObjectType = DataObjectType.TABLE, ) -> t.List[exp.Expr]: - return self._dcl_grants_config_expr(exp.Revoke, table, grants_config, table_type) + return self._dcl_grants_config_expr( + exp.Revoke, table, grants_config, table_type + ) def _get_grant_expression(self, table: exp.Table) -> exp.Expr: schema_identifier = table.args.get("db") or normalize_identifiers( - exp.to_identifier(self._get_current_schema(), quoted=True), dialect=self.dialect + exp.to_identifier(self._get_current_schema(), quoted=True), + dialect=self.dialect, ) schema_name = schema_identifier.this table_name = table.args.get("this").this # type: ignore @@ -640,7 +673,9 @@ def _get_grant_expression(self, table: exp.Table) -> exp.Expr: ] info_schema_table = normalize_identifiers( - exp.table_(self.GRANT_INFORMATION_SCHEMA_TABLE_NAME, db="information_schema"), + exp.table_( + self.GRANT_INFORMATION_SCHEMA_TABLE_NAME, db="information_schema" + ), dialect=self.dialect, ) if self.USE_CATALOG_IN_GRANTS: diff --git a/sqlmesh/core/engine_adapter/mssql.py b/sqlmesh/core/engine_adapter/mssql.py index fca6b4cc9f..59db446c0e 100644 --- a/sqlmesh/core/engine_adapter/mssql.py +++ b/sqlmesh/core/engine_adapter/mssql.py @@ -2,36 +2,27 @@ from __future__ import annotations -from textwrap import dedent -import typing as t import logging +import typing as t +from textwrap import dedent from sqlglot import exp -from sqlmesh.core.dialect import to_schema, add_table -from sqlmesh.core.engine_adapter.base import ( - EngineAdapterWithIndexSupport, - EngineAdapter, - InsertOverwriteStrategy, - MERGE_SOURCE_ALIAS, - MERGE_TARGET_ALIAS, - _get_data_object_cache_key, -) +from sqlmesh.core.dialect import add_table, to_schema +from sqlmesh.core.engine_adapter.base import (MERGE_SOURCE_ALIAS, + MERGE_TARGET_ALIAS, + EngineAdapter, + EngineAdapterWithIndexSupport, + InsertOverwriteStrategy, + _get_data_object_cache_key) from sqlmesh.core.engine_adapter.mixins import ( - GetCurrentCatalogFromFunctionMixin, - PandasNativeFetchDFSupportMixin, - VarcharSizeWorkaroundMixin, - RowDiffMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, - SourceQuery, - set_catalog, -) + GetCurrentCatalogFromFunctionMixin, PandasNativeFetchDFSupportMixin, + RowDiffMixin, VarcharSizeWorkaroundMixin) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + SourceQuery, set_catalog) from sqlmesh.utils import get_source_columns_to_types if t.TYPE_CHECKING: @@ -78,7 +69,14 @@ class MSSQLEngineAdapter( exp.DataType.build("NVARCHAR", dialect=DIALECT).this: 2147483647, }, } - VARIABLE_LENGTH_DATA_TYPES = {"binary", "varbinary", "char", "varchar", "nchar", "nvarchar"} + VARIABLE_LENGTH_DATA_TYPES = { + "binary", + "varbinary", + "char", + "varchar", + "nchar", + "nvarchar", + } INSERT_OVERWRITE_STRATEGY = InsertOverwriteStrategy.MERGE @property @@ -135,7 +133,10 @@ def build_var_length_col( ): return (column_name, f"{data_type}(max)") if data_type in ("decimal", "numeric"): - return (column_name, f"{data_type}({numeric_precision}, {numeric_scale})") + return ( + column_name, + f"{data_type}({numeric_precision}, {numeric_scale})", + ) if data_type == "float": return (column_name, f"{data_type}({numeric_precision})") @@ -151,7 +152,9 @@ def build_var_length_col( def table_exists(self, table_name: TableName) -> bool: """MsSql doesn't support describe so we query information_schema.""" table = exp.to_table(table_name) - data_object_cache_key = _get_data_object_cache_key(table.catalog, table.db, table.name) + data_object_cache_key = _get_data_object_cache_key( + table.catalog, table.db, table.name + ) if data_object_cache_key in self._data_object_cache: logger.debug("Table existence cache hit: %s", data_object_cache_key) return self._data_object_cache[data_object_cache_key] is not None @@ -199,7 +202,9 @@ def drop_schema( object_table, exists=ignore_if_not_exists, ) - super().drop_schema(schema_name, ignore_if_not_exists=ignore_if_not_exists, cascade=False) + super().drop_schema( + schema_name, ignore_if_not_exists=ignore_if_not_exists, cascade=False + ) def merge( self, @@ -212,18 +217,24 @@ def merge( source_columns: t.Optional[t.List[str]] = None, **kwargs: t.Any, ) -> None: - mssql_merge_exists = kwargs.get("physical_properties", {}).get("mssql_merge_exists") + mssql_merge_exists = kwargs.get("physical_properties", {}).get( + "mssql_merge_exists" + ) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - source_table, - target_columns_to_types, - target_table=target_table, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + source_table, + target_columns_to_types, + target_table=target_table, + source_columns=source_columns, + ) ) target_columns_to_types = target_columns_to_types or self.columns(target_table) on = exp.and_( *( - add_table(part, MERGE_TARGET_ALIAS).eq(add_table(part, MERGE_SOURCE_ALIAS)) + add_table(part, MERGE_TARGET_ALIAS).eq( + add_table(part, MERGE_SOURCE_ALIAS) + ) for part in unique_key ) ) @@ -283,7 +294,8 @@ def merge( ), expression=exp.Tuple( expressions=[ - exp.column(col, MERGE_SOURCE_ALIAS) for col in target_columns_to_types + exp.column(col, MERGE_SOURCE_ALIAS) + for col in target_columns_to_types ] ), ), @@ -298,7 +310,9 @@ def merge( whens=exp.Whens(expressions=match_expressions), ) - def _convert_df_datetime(self, df: DF, columns_to_types: t.Dict[str, exp.DataType]) -> None: + def _convert_df_datetime( + self, df: DF, columns_to_types: t.Dict[str, exp.DataType] + ) -> None: import pandas as pd from pandas.api.types import is_datetime64_any_dtype # type: ignore @@ -327,8 +341,8 @@ def _df_to_source_queries( target_table: TableName, source_columns: t.Optional[t.List[str]] = None, ) -> t.List[SourceQuery]: - import pandas as pd import numpy as np + import pandas as pd assert isinstance(df, pd.DataFrame) temp_table = self._get_temp_table(target_table or "pandas") @@ -336,7 +350,11 @@ def _df_to_source_queries( # Return the superclass implementation if the connection pool doesn't support bulk_copy if not hasattr(self._connection_pool.get(), "bulk_copy"): return super()._df_to_source_queries( - df, target_columns_to_types, batch_size, target_table, source_columns=source_columns + df, + target_columns_to_types, + batch_size, + target_table, + source_columns=source_columns, ) def query_factory() -> Query: @@ -358,8 +376,12 @@ def query_factory() -> Query: conn = self._connection_pool.get() conn.bulk_copy(temp_table.sql(dialect=self.dialect), rows) return exp.select( - *self._casted_columns(target_columns_to_types, source_columns=source_columns) - ).from_(temp_table) # type: ignore + *self._casted_columns( + target_columns_to_types, source_columns=source_columns + ) + ).from_( + temp_table + ) # type: ignore return [ SourceQuery( @@ -382,7 +404,10 @@ def _get_data_objects( exp.column("TABLE_NAME").as_("name"), exp.column("TABLE_SCHEMA").as_("schema_name"), exp.case() - .when(exp.column("TABLE_TYPE").eq("BASE TABLE"), exp.Literal.string("TABLE")) + .when( + exp.column("TABLE_TYPE").eq("BASE TABLE"), + exp.Literal.string("TABLE"), + ) .else_(exp.column("TABLE_TYPE")) .as_("type"), ) @@ -413,7 +438,9 @@ def _rename_table( ) -> None: # The function that renames tables in MSSQL takes string literals as arguments instead of identifiers, # so we shouldn't quote the identifiers. - self.execute(exp.rename_table(old_table_name, new_table_name), quote_identifiers=False) + self.execute( + exp.rename_table(old_table_name, new_table_name), quote_identifiers=False + ) def _insert_overwrite_by_condition( self, @@ -425,7 +452,9 @@ def _insert_overwrite_by_condition( **kwargs: t.Any, ) -> None: # note that this is passed as table_properties here rather than physical_properties - use_merge_strategy = kwargs.get("table_properties", {}).get("mssql_merge_exists") + use_merge_strategy = kwargs.get("table_properties", {}).get( + "mssql_merge_exists" + ) if (not where or where == exp.true()) and not use_merge_strategy: # this is a full table replacement, call the base strategy to do DELETE+INSERT # which will result in TRUNCATE+INSERT due to how we have overridden self.delete_from() @@ -454,7 +483,9 @@ def delete_from(self, table_name: TableName, where: t.Union[str, exp.Expr]) -> N # "A TRUNCATE TABLE operation can be rolled back within a transaction." # ref: https://learn.microsoft.com/en-us/sql/t-sql/statements/truncate-table-transact-sql?view=sql-server-ver15#remarks return self.execute( - exp.TruncateTable(expressions=[exp.to_table(table_name, dialect=self.dialect)]) + exp.TruncateTable( + expressions=[exp.to_table(table_name, dialect=self.dialect)] + ) ) return super().delete_from(table_name, where) @@ -492,13 +523,19 @@ def _build_create_comment_table_exp( schema_name=exp.Literal.string(table.db or "dbo").sql( dialect=self.dialect, identify=False ), - object_name=exp.Literal.string(table.name).sql(dialect=self.dialect, identify=False), + object_name=exp.Literal.string(table.name).sql( + dialect=self.dialect, identify=False + ), object_kind=table_kind, ) return tsql_text def _build_create_comment_column_exp( - self, table: exp.Table, column_name: str, column_comment: str, table_kind: str = "TABLE" + self, + table: exp.Table, + column_name: str, + column_comment: str, + table_kind: str = "TABLE", ) -> exp.Comment | str: template = dedent(""" DECLARE @comment sql_variant = {comment}; @@ -532,9 +569,13 @@ def _build_create_comment_column_exp( schema_name=exp.Literal.string(table.db or "dbo").sql( dialect=self.dialect, identify=False ), - object_name=exp.Literal.string(table.name).sql(dialect=self.dialect, identify=False), + object_name=exp.Literal.string(table.name).sql( + dialect=self.dialect, identify=False + ), object_kind=table_kind, - column_name=exp.Literal.string(column_name).sql(dialect=self.dialect, identify=False), + column_name=exp.Literal.string(column_name).sql( + dialect=self.dialect, identify=False + ), ) return tsql_text diff --git a/sqlmesh/core/engine_adapter/mysql.py b/sqlmesh/core/engine_adapter/mysql.py index 6918cdec49..b2616152bf 100644 --- a/sqlmesh/core/engine_adapter/mysql.py +++ b/sqlmesh/core/engine_adapter/mysql.py @@ -2,22 +2,17 @@ import logging import typing as t + from sqlglot import exp, parse_one from sqlmesh.core.dialect import to_schema from sqlmesh.core.engine_adapter.mixins import ( - LogicalMergeMixin, - NonTransactionalTruncateMixin, - PandasNativeFetchDFSupportMixin, - RowDiffMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, - set_catalog, -) + LogicalMergeMixin, NonTransactionalTruncateMixin, + PandasNativeFetchDFSupportMixin, RowDiffMixin) +from sqlmesh.core.engine_adapter.shared import (CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + set_catalog) if t.TYPE_CHECKING: from sqlmesh.core._typing import SchemaName, TableName @@ -28,7 +23,10 @@ @set_catalog() class MySQLEngineAdapter( - LogicalMergeMixin, PandasNativeFetchDFSupportMixin, NonTransactionalTruncateMixin, RowDiffMixin + LogicalMergeMixin, + PandasNativeFetchDFSupportMixin, + NonTransactionalTruncateMixin, + RowDiffMixin, ): DEFAULT_BATCH_SIZE = 200 DIALECT = "mysql" @@ -76,7 +74,9 @@ def drop_schema( **drop_args: t.Dict[str, exp.Expr], ) -> None: # MySQL doesn't support CASCADE clause and drops schemas unconditionally. - super().drop_schema(schema_name, ignore_if_not_exists=ignore_if_not_exists, cascade=False) + super().drop_schema( + schema_name, ignore_if_not_exists=ignore_if_not_exists, cascade=False + ) def _get_data_objects( self, schema_name: SchemaName, object_names: t.Optional[t.Set[str]] = None @@ -235,8 +235,12 @@ def _replace_by_key( ] ) - target_table_aliased = exp.to_table(target_table).as_(target_alias, quoted=True) - temp_table_aliased = exp.to_table(temp_table).as_(temp_alias, quoted=True) + target_table_aliased = exp.to_table(target_table).as_( + target_alias, quoted=True + ) + temp_table_aliased = exp.to_table(temp_table).as_( + temp_alias, quoted=True + ) join = exp.Join(this=temp_table_aliased, kind="INNER", on=on_condition) target_table_aliased.append("joins", join) @@ -247,7 +251,9 @@ def _replace_by_key( ) self.execute(delete_stmt) - insert_query = self._select_columns(target_columns_to_types).from_(temp_table) + insert_query = self._select_columns(target_columns_to_types).from_( + temp_table + ) if is_unique_key: insert_query = insert_query.distinct(*key) diff --git a/sqlmesh/core/engine_adapter/postgres.py b/sqlmesh/core/engine_adapter/postgres.py index 6794169322..84f19b59fe 100644 --- a/sqlmesh/core/engine_adapter/postgres.py +++ b/sqlmesh/core/engine_adapter/postgres.py @@ -4,16 +4,13 @@ import re import typing as t from functools import cached_property, partial + from sqlglot import exp from sqlmesh.core.engine_adapter.base_postgres import BasePostgresEngineAdapter from sqlmesh.core.engine_adapter.mixins import ( - GetCurrentCatalogFromFunctionMixin, - PandasNativeFetchDFSupportMixin, - RowDiffMixin, - logical_merge, - GrantsFromInfoSchemaMixin, -) + GetCurrentCatalogFromFunctionMixin, GrantsFromInfoSchemaMixin, + PandasNativeFetchDFSupportMixin, RowDiffMixin, logical_merge) from sqlmesh.core.engine_adapter.shared import set_catalog if t.TYPE_CHECKING: @@ -45,7 +42,10 @@ class PostgresEngineAdapter( SCHEMA_DIFFER_KWARGS = { "parameterized_type_defaults": { # DECIMAL without precision is "up to 131072 digits before the decimal point; up to 16383 digits after the decimal point" - exp.DataType.build("DECIMAL", dialect=DIALECT).this: [(131072 + 16383, 16383), (0,)], + exp.DataType.build("DECIMAL", dialect=DIALECT).this: [ + (131072 + 16383, 16383), + (0,), + ], exp.DataType.build("CHAR", dialect=DIALECT).this: [(1,)], exp.DataType.build("TIME", dialect=DIALECT).this: [(6,)], exp.DataType.build("TIMESTAMP", dialect=DIALECT).this: [(6,)], @@ -99,7 +99,11 @@ def _create_table_like( expressions=[ exp.LikeProperty( this=exp.to_table(source_table_name), - expressions=[exp.Property(this="INCLUDING", value=exp.Var(this="ALL"))], + expressions=[ + exp.Property( + this="INCLUDING", value=exp.Var(this="ALL") + ) + ], ) ], ), diff --git a/sqlmesh/core/engine_adapter/redshift.py b/sqlmesh/core/engine_adapter/redshift.py index 39453f0cd2..334fff2889 100644 --- a/sqlmesh/core/engine_adapter/redshift.py +++ b/sqlmesh/core/engine_adapter/redshift.py @@ -7,30 +7,23 @@ from sqlglot.helper import ensure_list from sqlmesh.core.dialect import to_schema -from sqlmesh.core.engine_adapter.base import MERGE_SOURCE_ALIAS, MERGE_TARGET_ALIAS +from sqlmesh.core.engine_adapter.base import (MERGE_SOURCE_ALIAS, + MERGE_TARGET_ALIAS) from sqlmesh.core.engine_adapter.base_postgres import BasePostgresEngineAdapter from sqlmesh.core.engine_adapter.mixins import ( - GetCurrentCatalogFromFunctionMixin, - NonTransactionalTruncateMixin, - VarcharSizeWorkaroundMixin, - RowDiffMixin, - logical_merge, - GrantsFromInfoSchemaMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CommentCreationView, - DataObject, - DataObjectType, - SourceQuery, - set_catalog, -) + GetCurrentCatalogFromFunctionMixin, GrantsFromInfoSchemaMixin, + NonTransactionalTruncateMixin, RowDiffMixin, VarcharSizeWorkaroundMixin, + logical_merge) +from sqlmesh.core.engine_adapter.shared import (CommentCreationView, + DataObject, DataObjectType, + SourceQuery, set_catalog) from sqlmesh.utils.errors import SQLMeshError if t.TYPE_CHECKING: import pandas as pd from sqlmesh.core._typing import SchemaName, TableName - from sqlmesh.core.engine_adapter.base import QueryOrDF, Query + from sqlmesh.core.engine_adapter.base import Query, QueryOrDF from sqlmesh.core.node import IntervalUnit logger = logging.getLogger(__name__) @@ -66,7 +59,9 @@ class RedshiftEngineAdapter( exp.DataType.build("CHAR", dialect=DIALECT).this: 4096, exp.DataType.build("VARCHAR", dialect=DIALECT).this: 65535, }, - "precision_increase_allowed_types": {exp.DataType.build("VARCHAR", dialect=DIALECT).this}, + "precision_increase_allowed_types": { + exp.DataType.build("VARCHAR", dialect=DIALECT).this + }, "drop_cascade": True, } VARIABLE_LENGTH_DATA_TYPES = { @@ -118,7 +113,10 @@ def build_var_length_col( ): return (column_name, f"{data_type}({character_maximum_length})") if data_type in ("decimal", "numeric"): - return (column_name, f"{data_type}({numeric_precision}, {numeric_scale})") + return ( + column_name, + f"{data_type}({numeric_precision}, {numeric_scale})", + ) return (column_name, data_type) @@ -270,7 +268,9 @@ def _build_table_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) @@ -287,15 +287,21 @@ def _to_identifier_if_string(expression: exp.Expr) -> exp.Expr: diststyle = table_properties.get("DISTSTYLE") if diststyle: - properties.append(exp.DistStyleProperty(this=exp.var(diststyle.name.upper()))) + properties.append( + exp.DistStyleProperty(this=exp.var(diststyle.name.upper())) + ) distkey = table_properties.get("DISTKEY") if distkey: - properties.append(exp.DistKeyProperty(this=_to_identifier_if_string(distkey))) + properties.append( + exp.DistKeyProperty(this=_to_identifier_if_string(distkey)) + ) sortkey = table_properties.get("SORTKEY") if sortkey: - sortkey_expressions = sortkey.expressions if sortkey.expressions else [sortkey] + sortkey_expressions = ( + sortkey.expressions if sortkey.expressions else [sortkey] + ) properties.append( exp.SortKeyProperty( this=[ @@ -331,7 +337,9 @@ def replace_query( target_data_object = self.get_data_object(table_name) table_exists = target_data_object is not None - if self.drop_data_object_on_type_mismatch(target_data_object, DataObjectType.TABLE): + if self.drop_data_object_on_type_mismatch( + target_data_object, DataObjectType.TABLE + ): table_exists = False if not isinstance(query_or_df, pd.DataFrame) or not table_exists: @@ -344,11 +352,13 @@ def replace_query( source_columns=source_columns, **kwargs, ) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - target_table=table_name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + target_table=table_name, + source_columns=source_columns, + ) ) target_columns_to_types = target_columns_to_types or self.columns(table_name) target_table = exp.to_table(table_name) @@ -363,7 +373,9 @@ def replace_query( column_descriptions=column_descriptions, **kwargs, ) - self._insert_append_source_queries(temp_table, source_queries, target_columns_to_types) + self._insert_append_source_queries( + temp_table, source_queries, target_columns_to_types + ) self.rename_table(target_table, old_table) self.rename_table(temp_table, target_table) self.drop_table(old_table) @@ -476,7 +488,8 @@ def resolve_target_table(expression: exp.Expr) -> exp.Expr: # Since Redshift does not support multiple "WHEN MATCHED" clauses. if ( len(whens.expressions) != 2 - or whens.expressions[0].args["matched"] == whens.expressions[1].args["matched"] + or whens.expressions[0].args["matched"] + == whens.expressions[1].args["matched"] ): raise SQLMeshError( "Redshift only supports a single WHEN MATCHED and WHEN NOT MATCHED clause" diff --git a/sqlmesh/core/engine_adapter/risingwave.py b/sqlmesh/core/engine_adapter/risingwave.py index 61b44f5bbb..8d0c2fd187 100644 --- a/sqlmesh/core/engine_adapter/risingwave.py +++ b/sqlmesh/core/engine_adapter/risingwave.py @@ -3,17 +3,13 @@ import logging import typing as t - from sqlglot import exp from sqlmesh.core.engine_adapter.postgres import PostgresEngineAdapter -from sqlmesh.core.engine_adapter.shared import ( - set_catalog, - CatalogSupport, - CommentCreationView, - CommentCreationTable, -) - +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + set_catalog) from sqlmesh.utils.errors import SQLMeshError if t.TYPE_CHECKING: @@ -41,9 +37,13 @@ def columns( table = exp.to_table(table_name) sql = ( - exp.select("rw_columns.name AS column_name", "rw_columns.data_type AS data_type") + exp.select( + "rw_columns.name AS column_name", "rw_columns.data_type AS data_type" + ) .from_("rw_catalog.rw_columns") - .join("rw_catalog.rw_relations", on="rw_relations.id=rw_columns.relation_id") + .join( + "rw_catalog.rw_relations", on="rw_relations.id=rw_columns.relation_id" + ) .join("rw_catalog.rw_schemas", on="rw_schemas.id=rw_relations.schema_id") .where( exp.and_( @@ -60,7 +60,9 @@ def columns( self.execute(sql) resp = self.cursor.fetchall() if not resp: - raise SQLMeshError(f"Could not get columns for table {table_name}. Table not found.") + raise SQLMeshError( + f"Could not get columns for table {table_name}. Table not found." + ) return { column_name: exp.DataType.build(data_type, dialect=self.dialect, udt=True) for column_name, data_type in resp diff --git a/sqlmesh/core/engine_adapter/shared.py b/sqlmesh/core/engine_adapter/shared.py index ba0e1fa619..55d788f7d7 100644 --- a/sqlmesh/core/engine_adapter/shared.py +++ b/sqlmesh/core/engine_adapter/shared.py @@ -11,7 +11,7 @@ from sqlglot import exp from sqlmesh.core.dialect import to_schema -from sqlmesh.utils.errors import UnsupportedCatalogOperationError, SQLMeshError +from sqlmesh.utils.errors import SQLMeshError, UnsupportedCatalogOperationError from sqlmesh.utils.pydantic import PydanticModel if t.TYPE_CHECKING: @@ -172,7 +172,9 @@ def is_clustered(self) -> bool: return bool(self.clustering_key) def to_table(self) -> exp.Table: - return exp.table_(self.name, db=self.schema_name, catalog=self.catalog, quoted=True) + return exp.table_( + self.name, db=self.schema_name, catalog=self.catalog, quoted=True + ) class CatalogSupport(Enum): @@ -299,7 +301,9 @@ def __exit__( return None -def set_catalog(override_mapping: t.Optional[t.Dict[str, CatalogSupport]] = None) -> t.Callable: +def set_catalog( + override_mapping: t.Optional[t.Dict[str, CatalogSupport]] = None, +) -> t.Callable: def set_catalog_decorator( func: t.Callable, target_name: str, @@ -318,7 +322,9 @@ def internal_wrapper(*args: t.Any, **kwargs: t.Any) -> t.Any: return func(*list_args, **kwargs) obj, container, key = t.cast( - t.Tuple[t.Union[str, exp.Table], t.Union[t.Dict, t.List], t.Union[int, str]], + t.Tuple[ + t.Union[str, exp.Table], t.Union[t.Dict, t.List], t.Union[int, str] + ], ( (kwargs.get(target_name), kwargs, target_name) if kwargs.get(target_name) @@ -329,7 +335,9 @@ def internal_wrapper(*args: t.Any, **kwargs: t.Any) -> t.Any: t.Callable[[t.Union[str, exp.Table]], exp.Table], exp.to_table if target_type == "TableName" else to_schema, ) - expression = to_expression_func(obj.copy() if isinstance(obj, exp.Table) else obj) + expression = to_expression_func( + obj.copy() if isinstance(obj, exp.Table) else obj + ) catalog_name = expression.catalog if not catalog_name: return func(*list_args, **kwargs) @@ -373,7 +381,9 @@ def internal_wrapper(*args: t.Any, **kwargs: t.Any) -> t.Any: def wrapper(cls: t.Type[EngineAdapter]) -> t.Callable: for name in dir(cls): - if name in exclusion_list or (name.startswith("_") and name not in inclusion_list): + if name in exclusion_list or ( + name.startswith("_") and name not in inclusion_list + ): continue m = getattr(cls, name) if inspect.isfunction(m): diff --git a/sqlmesh/core/engine_adapter/snowflake.py b/sqlmesh/core/engine_adapter/snowflake.py index d589b5d15b..b996ca09ef 100644 --- a/sqlmesh/core/engine_adapter/snowflake.py +++ b/sqlmesh/core/engine_adapter/snowflake.py @@ -12,19 +12,12 @@ import sqlmesh.core.constants as c from sqlmesh.core.dialect import to_schema from sqlmesh.core.engine_adapter.mixins import ( - GetCurrentCatalogFromFunctionMixin, - ClusteredByMixin, - RowDiffMixin, - GrantsFromInfoSchemaMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - DataObject, - DataObjectType, - SourceQuery, - set_catalog, -) -from sqlmesh.utils import optional_import, get_source_columns_to_types + ClusteredByMixin, GetCurrentCatalogFromFunctionMixin, + GrantsFromInfoSchemaMixin, RowDiffMixin) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, DataObject, + DataObjectType, SourceQuery, + set_catalog) +from sqlmesh.utils import get_source_columns_to_types, optional_import from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.pandas import columns_to_types_from_dtypes @@ -35,12 +28,8 @@ import pandas as pd from sqlmesh.core._typing import SchemaName, SessionProperties, TableName - from sqlmesh.core.engine_adapter._typing import ( - DF, - Query, - QueryOrDF, - SnowparkSession, - ) + from sqlmesh.core.engine_adapter._typing import (DF, Query, QueryOrDF, + SnowparkSession) from sqlmesh.core.node import IntervalUnit @@ -53,7 +42,10 @@ } ) class SnowflakeEngineAdapter( - GetCurrentCatalogFromFunctionMixin, ClusteredByMixin, RowDiffMixin, GrantsFromInfoSchemaMixin + GetCurrentCatalogFromFunctionMixin, + ClusteredByMixin, + RowDiffMixin, + GrantsFromInfoSchemaMixin, ): DIALECT = "snowflake" SUPPORTS_MATERIALIZED_VIEWS = True @@ -118,7 +110,9 @@ def session(self, properties: SessionProperties) -> t.Iterator[None]: def _current_warehouse(self) -> exp.Identifier: current_warehouse_str = self.fetchone("SELECT CURRENT_WAREHOUSE()")[0] # type: ignore # The warehouse value returned by Snowflake is already normalized, so only quoting is needed. - return quote_identifiers(exp.to_identifier(current_warehouse_str), dialect=self.dialect) + return quote_identifiers( + exp.to_identifier(current_warehouse_str), dialect=self.dialect + ) @property def snowpark(self) -> t.Optional[SnowparkSession]: @@ -158,11 +152,16 @@ def _get_current_schema(self) -> str: def _create_catalog(self, catalog_name: exp.Identifier) -> None: props = exp.Properties( - expressions=[exp.SchemaCommentProperty(this=exp.Literal.string(c.SQLMESH_MANAGED))] + expressions=[ + exp.SchemaCommentProperty(this=exp.Literal.string(c.SQLMESH_MANAGED)) + ] ) self.execute( exp.Create( - this=exp.Table(this=catalog_name), kind="DATABASE", exists=True, properties=props + this=exp.Table(this=catalog_name), + kind="DATABASE", + exists=True, + properties=props, ) ) @@ -180,7 +179,11 @@ def _drop_catalog(self, catalog_name: exp.Identifier) -> None: ) normalize_identifiers(exists_check, dialect=self.dialect) if self.fetchone(exists_check, quote_identifiers=True) is not None: - self.execute(exp.Drop(this=exp.Table(this=catalog_name), kind="DATABASE", exists=True)) + self.execute( + exp.Drop( + this=exp.Table(this=catalog_name), kind="DATABASE", exists=True + ) + ) else: logger.warning( f"Not dropping database {catalog_name.sql(dialect=self.dialect)} because there is no indication it is '{c.SQLMESH_MANAGED}'" @@ -250,8 +253,13 @@ def create_managed_table( "`target_lag` must be specified in the model physical_properties for a Snowflake Dynamic Table" ) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query, target_columns_to_types, target_table=target_table, source_columns=source_columns + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query, + target_columns_to_types, + target_table=target_table, + source_columns=source_columns, + ) ) self._create_table_from_source_queries( @@ -328,13 +336,16 @@ def _build_table_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) if ( clustered_by - and (clustered_by_prop := self._build_clustered_by_exp(clustered_by)) is not None + and (clustered_by_prop := self._build_clustered_by_exp(clustered_by)) + is not None ): properties.append(clustered_by_prop) @@ -349,7 +360,9 @@ def _build_table_properties_exp( table_type = self._pop_creatable_type_from_properties(table_properties) properties.extend(ensure_list(table_type)) - properties.extend(self._table_or_view_properties_to_expressions(table_properties)) + properties.extend( + self._table_or_view_properties_to_expressions(table_properties) + ) return exp.Properties(expressions=properties) if properties else None @@ -372,7 +385,9 @@ def _df_to_source_queries( target_table or "pandas", quoted=False ) # write_pandas() re-quotes everything without checking if its already quoted - is_snowpark_dataframe = snowpark and isinstance(df, snowpark.dataframe.DataFrame) + is_snowpark_dataframe = snowpark and isinstance( + df, snowpark.dataframe.DataFrame + ) def query_factory() -> Query: # The catalog needs to be normalized before being passed to Snowflake's library functions because they @@ -397,7 +412,9 @@ def query_factory() -> Query: if not columns_already_quoted: local_df = df.rename( { - col: exp.to_identifier(col).sql(dialect=self.dialect, identify=True) + col: exp.to_identifier(col).sql( + dialect=self.dialect, identify=True + ) for col in source_columns_to_types } ) # type: ignore @@ -423,18 +440,24 @@ def query_factory() -> Query: if kind.is_type("date"): # type: ignore ordered_df[column] = pd.to_datetime(ordered_df[column]).dt.date # type: ignore elif getattr(ordered_df.dtypes[column], "tz", None) is not None: # type: ignore - ordered_df[column] = pd.to_datetime(ordered_df[column]).dt.strftime( + ordered_df[column] = pd.to_datetime( + ordered_df[column] + ).dt.strftime( "%Y-%m-%d %H:%M:%S.%f%z" ) # type: ignore # https://github.com/snowflakedb/snowflake-connector-python/issues/1677 else: # type: ignore - ordered_df[column] = pd.to_datetime(ordered_df[column]).dt.strftime( + ordered_df[column] = pd.to_datetime( + ordered_df[column] + ).dt.strftime( "%Y-%m-%d %H:%M:%S.%f" ) # type: ignore # create the table first using our usual method ensure the column datatypes match what we parsed with sqlglot # otherwise we would be trusting `write_pandas()` from the snowflake lib to do this correctly - self.create_table(temp_table, source_columns_to_types, table_kind="TEMPORARY TABLE") + self.create_table( + temp_table, source_columns_to_types, table_kind="TEMPORARY TABLE" + ) write_pandas( self._connection_pool.get(), @@ -452,7 +475,9 @@ def query_factory() -> Query: ) return exp.select( - *self._casted_columns(target_columns_to_types, source_columns=source_columns) + *self._casted_columns( + target_columns_to_types, source_columns=source_columns + ) ).from_(temp_table) def cleanup() -> None: @@ -520,11 +545,26 @@ def _get_data_objects( ), exp.Literal.string("MANAGED_TABLE"), ) - .when(exp.column("TABLE_TYPE").eq("BASE TABLE"), exp.Literal.string("TABLE")) - .when(exp.column("TABLE_TYPE").eq("TEMPORARY TABLE"), exp.Literal.string("TABLE")) - .when(exp.column("TABLE_TYPE").eq("LOCAL TEMPORARY"), exp.Literal.string("TABLE")) - .when(exp.column("TABLE_TYPE").eq("EXTERNAL TABLE"), exp.Literal.string("TABLE")) - .when(exp.column("TABLE_TYPE").eq("EVENT TABLE"), exp.Literal.string("TABLE")) + .when( + exp.column("TABLE_TYPE").eq("BASE TABLE"), + exp.Literal.string("TABLE"), + ) + .when( + exp.column("TABLE_TYPE").eq("TEMPORARY TABLE"), + exp.Literal.string("TABLE"), + ) + .when( + exp.column("TABLE_TYPE").eq("LOCAL TEMPORARY"), + exp.Literal.string("TABLE"), + ) + .when( + exp.column("TABLE_TYPE").eq("EXTERNAL TABLE"), + exp.Literal.string("TABLE"), + ) + .when( + exp.column("TABLE_TYPE").eq("EVENT TABLE"), + exp.Literal.string("TABLE"), + ) .when(exp.column("TABLE_TYPE").eq("VIEW"), exp.Literal.string("VIEW")) .when( exp.column("TABLE_TYPE").eq("MATERIALIZED VIEW"), @@ -544,7 +584,9 @@ def _get_data_objects( # exclude SNOWPARK_TEMP_TABLE tables that are managed by the Snowpark library and are an implementation # detail of dealing with DataFrame's - query = query.where(exp.column("TABLE_NAME").like("SNOWPARK_TEMP_TABLE%").not_()) + query = query.where( + exp.column("TABLE_NAME").like("SNOWPARK_TEMP_TABLE%").not_() + ) df = self.fetchdf(query, quote_identifiers=True) if df.empty: @@ -558,7 +600,9 @@ def _get_data_objects( clustering_key=row.clustering_key, # type: ignore ) # lowercase the column names for cases where Snowflake might return uppercase column names for certain catalogs - for row in df.rename(columns={col: col.lower() for col in df.columns}).itertuples() + for row in df.rename( + columns={col: col.lower() for col in df.columns} + ).itertuples() ] def _get_grant_expression(self, table: exp.Table) -> exp.Expr: @@ -569,14 +613,22 @@ def _get_grant_expression(self, table: exp.Table) -> exp.Expr: for col_exp in expression.find_all(exp.Column): if col_exp.this.name == "table_catalog": and_exp = col_exp.parent - assert and_exp is not None, "Expected column expression to have a parent" - assert and_exp.expression, "Expected AND expression to have an expression" + assert ( + and_exp is not None + ), "Expected column expression to have a parent" + assert ( + and_exp.expression + ), "Expected AND expression to have an expression" normalized_catalog = self._normalize_catalog( - exp.table_("placeholder", db="placeholder", catalog=and_exp.expression.this) + exp.table_( + "placeholder", db="placeholder", catalog=and_exp.expression.this + ) ) and_exp.set( "expression", - exp.Literal.string(normalized_catalog.args["catalog"].alias_or_name), + exp.Literal.string( + normalized_catalog.args["catalog"].alias_or_name + ), ) return expression @@ -610,8 +662,13 @@ def catalog_rewriter(node: exp.Expr) -> exp.Expr: # only replace the catalog on the model with the target catalog if the two are functionally equivalent if unquote_and_lower(node.catalog) == default_catalog_unquoted: node.set("catalog", default_catalog_normalized) - elif isinstance(node, exp.Use) and isinstance(node.this, exp.Identifier): - if unquote_and_lower(node.this.output_name) == default_catalog_unquoted: + elif isinstance(node, exp.Use) and isinstance( + node.this, exp.Identifier + ): + if ( + unquote_and_lower(node.this.output_name) + == default_catalog_unquoted + ): node.set("this", default_catalog_normalized) return node @@ -644,14 +701,20 @@ def _create_column_comments( list_comment_sql = [] for column_name, column_comment in column_comments.items(): - column_sql = exp.column(column_name).sql(dialect=self.dialect, identify=True) + column_sql = exp.column(column_name).sql( + dialect=self.dialect, identify=True + ) truncated_comment = self._truncate_column_comment(column_comment) - comment_sql = exp.Literal.string(truncated_comment).sql(dialect=self.dialect) + comment_sql = exp.Literal.string(truncated_comment).sql( + dialect=self.dialect + ) list_comment_sql.append(f"COLUMN {column_sql} COMMENT {comment_sql}") - combined_sql = f"ALTER {table_kind} {table_sql} ALTER {', '.join(list_comment_sql)}" + combined_sql = ( + f"ALTER {table_kind} {table_sql} ALTER {', '.join(list_comment_sql)}" + ) try: self.execute(combined_sql) except Exception: @@ -705,11 +768,17 @@ def _columns_to_types( target_columns_to_types: t.Optional[t.Dict[str, exp.DataType]] = None, source_columns: t.Optional[t.List[str]] = None, ) -> t.Tuple[t.Optional[t.Dict[str, exp.DataType]], t.Optional[t.List[str]]]: - if not target_columns_to_types and snowpark and isinstance(query_or_df, snowpark.DataFrame): + if ( + not target_columns_to_types + and snowpark + and isinstance(query_or_df, snowpark.DataFrame) + ): target_columns_to_types = columns_to_types_from_dtypes( query_or_df.sample(n=1).to_pandas().dtypes.items() ) - return target_columns_to_types, list(source_columns or target_columns_to_types) + return target_columns_to_types, list( + source_columns or target_columns_to_types + ) return super()._columns_to_types( query_or_df, target_columns_to_types, source_columns=source_columns diff --git a/sqlmesh/core/engine_adapter/spark.py b/sqlmesh/core/engine_adapter/spark.py index 9199aa3bcd..7533e2fe52 100644 --- a/sqlmesh/core/engine_adapter/spark.py +++ b/sqlmesh/core/engine_adapter/spark.py @@ -8,20 +8,14 @@ from sqlmesh.core.dialect import to_schema from sqlmesh.core.engine_adapter.mixins import ( - GetCurrentCatalogFromFunctionMixin, - HiveMetastoreTablePropertiesMixin, - RowDiffMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, - InsertOverwriteStrategy, - SourceQuery, - set_catalog, -) + GetCurrentCatalogFromFunctionMixin, HiveMetastoreTablePropertiesMixin, + RowDiffMixin) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + InsertOverwriteStrategy, + SourceQuery, set_catalog) from sqlmesh.utils import classproperty, get_source_columns_to_types from sqlmesh.utils.errors import SQLMeshError @@ -30,14 +24,11 @@ from pyspark.sql import types as spark_types from sqlmesh.core._typing import SchemaName, TableName - from sqlmesh.core.engine_adapter._typing import ( - DF, - PySparkDataFrame, - PySparkSession, - Query, - ) + from sqlmesh.core.engine_adapter._typing import (DF, PySparkDataFrame, + PySparkSession, Query) from sqlmesh.core.engine_adapter.base import QueryOrDF - from sqlmesh.engines.spark.db_api.spark_session import SparkSessionConnection + from sqlmesh.engines.spark.db_api.spark_session import \ + SparkSessionConnection logger = logging.getLogger(__name__) @@ -131,7 +122,9 @@ def _spark_to_sqlglot_complex_mapping(self) -> t.Dict[t.Any, t.Any]: return {v: k for k, v in self._sqlglot_to_spark_complex_mapping.items()} @classmethod - def spark_to_sqlglot_types(cls, input: spark_types.StructType) -> t.Dict[str, exp.DataType]: + def spark_to_sqlglot_types( + cls, input: spark_types.StructType + ) -> t.Dict[str, exp.DataType]: from pyspark.sql import types as spark_types def spark_complex_to_sqlglot_complex( @@ -168,9 +161,13 @@ def get_fields( if isinstance(sqlglot_data_type, exp.DataType) else exp.DataType(this=sqlglot_data_type) ) - expressions.append(exp.ColumnDef(this=exp.to_identifier(field.name), kind=kind)) + expressions.append( + exp.ColumnDef(this=exp.to_identifier(field.name), kind=kind) + ) else: - kind = exp.DataType(this=cls._spark_to_sqlglot_primitive_mapping[type(field)]) + kind = exp.DataType( + this=cls._spark_to_sqlglot_primitive_mapping[type(field)] + ) expressions.append(kind) dtype = cls._spark_to_sqlglot_complex_mapping[type(complex_type)] return exp.DataType( @@ -180,19 +177,28 @@ def get_fields( ) resp = spark_complex_to_sqlglot_complex(input) - return {column_def.this.name: column_def.args["kind"] for column_def in resp.expressions} + return { + column_def.this.name: column_def.args["kind"] + for column_def in resp.expressions + } @classmethod - def sqlglot_to_spark_types(cls, input: t.Dict[str, exp.DataType]) -> spark_types.StructType: + def sqlglot_to_spark_types( + cls, input: t.Dict[str, exp.DataType] + ) -> spark_types.StructType: from pyspark.sql import types as spark_types - def sqlglot_complex_to_spark_complex(complex_type: exp.DataType) -> spark_types.DataType: + def sqlglot_complex_to_spark_complex( + complex_type: exp.DataType, + ) -> spark_types.DataType: is_struct = complex_type.is_type(exp.DataType.Type.STRUCT) expressions = [] for column_def in complex_type.expressions: col_name = column_def.this.name if is_struct else None data_type = column_def.args["kind"] if is_struct else column_def - primitive_func = cls._sqlglot_to_spark_primitive_mapping.get(data_type.this) + primitive_func = cls._sqlglot_to_spark_primitive_mapping.get( + data_type.this + ) type_func = ( primitive_func if primitive_func @@ -261,14 +267,18 @@ def _columns_to_types( source_columns: t.Optional[t.List[str]] = None, ) -> t.Tuple[t.Optional[t.Dict[str, exp.DataType]], t.Optional[t.List[str]]]: if target_columns_to_types: - return target_columns_to_types, list(source_columns or target_columns_to_types) + return target_columns_to_types, list( + source_columns or target_columns_to_types + ) if self.is_pyspark_df(query_or_df): from pyspark.sql import DataFrame target_columns_to_types = self.spark_to_sqlglot_types( t.cast(DataFrame, query_or_df).schema ) - return target_columns_to_types, list(source_columns or target_columns_to_types) + return target_columns_to_types, list( + source_columns or target_columns_to_types + ) return super()._columns_to_types( query_or_df, target_columns_to_types, source_columns=source_columns ) @@ -281,13 +291,17 @@ def _df_to_source_queries( target_table: TableName, source_columns: t.Optional[t.List[str]] = None, ) -> t.List[SourceQuery]: - df = self._ensure_pyspark_df(df, target_columns_to_types, source_columns=source_columns) + df = self._ensure_pyspark_df( + df, target_columns_to_types, source_columns=source_columns + ) def query_factory() -> Query: temp_table = self._get_temp_table(target_table or "spark", table_only=True) df.createOrReplaceGlobalTempView(temp_table.sql(dialect=self.dialect)) # type: ignore temp_table.set("db", "global_temp") - return exp.select(*self._select_columns(target_columns_to_types)).from_(temp_table) + return exp.select(*self._select_columns(target_columns_to_types)).from_( + temp_table + ) return [SourceQuery(query_factory=query_factory)] @@ -342,7 +356,9 @@ def _get_temp_table( def fetchdf( self, query: t.Union[exp.Expr, str], quote_identifiers: bool = False ) -> pd.DataFrame: - return self.fetch_pyspark_df(query, quote_identifiers=quote_identifiers).toPandas() + return self.fetch_pyspark_df( + query, quote_identifiers=quote_identifiers + ).toPandas() def fetch_pyspark_df( self, query: t.Union[exp.Expr, str], quote_identifiers: bool = False @@ -406,7 +422,9 @@ def get_data_object( self, target_name: TableName, safe_to_cache: bool = False ) -> t.Optional[DataObject]: target_table = exp.to_table(target_name) - if isinstance(target_table.this, exp.Dot) and target_table.this.expression.name.startswith( + if isinstance( + target_table.this, exp.Dot + ) and target_table.this.expression.name.startswith( f"{self.BRANCH_PREFIX}{self.WAP_PREFIX}" ): # Exclude the branch name @@ -422,7 +440,9 @@ def create_state_table( self.create_table( table_name, target_columns_to_types, - partitioned_by=[exp.column(x) for x in primary_key] if primary_key else None, + partitioned_by=( + [exp.column(x) for x in primary_key] if primary_key else None + ), ) def _native_df_to_pandas_df( @@ -465,7 +485,9 @@ def _create_table( kwargs.get("storage_format") or "" ).lower() == "iceberg" or self.wap_supported(table_name) do_dummy_insert = ( - False if not wap_supported or not exists else not self.table_exists(table_name) + False + if not wap_supported or not exists + else not self.table_exists(table_name) ) super()._create_table( table_name_or_schema, @@ -499,14 +521,16 @@ def wap_supported(self, table_name: TableName) -> bool: def wap_table_name(self, table_name: TableName, wap_id: str) -> str: branch_name = self._wap_branch_name(wap_id) fqn = self._ensure_fqn(table_name) - return exp.Dot.build([fqn, exp.to_identifier(f"{self.BRANCH_PREFIX}{branch_name}")]).sql( - dialect=self.dialect - ) + return exp.Dot.build( + [fqn, exp.to_identifier(f"{self.BRANCH_PREFIX}{branch_name}")] + ).sql(dialect=self.dialect) def wap_prepare(self, table_name: TableName, wap_id: str) -> str: branch_name = self._wap_branch_name(wap_id) fqn = self._ensure_fqn(table_name) - self.execute(f"ALTER TABLE {fqn.sql(dialect=self.dialect)} CREATE BRANCH {branch_name}") + self.execute( + f"ALTER TABLE {fqn.sql(dialect=self.dialect)} CREATE BRANCH {branch_name}" + ) return self.wap_table_name(table_name, wap_id) def wap_publish(self, table_name: TableName, wap_id: str) -> None: @@ -524,13 +548,17 @@ def wap_publish(self, table_name: TableName, wap_id: str) -> None: iceberg_snapshot_id = iceberg_snapshot_ids[0][0] logger.info( - "Cherry-picking Iceberg snapshot %s into table '%s'...", iceberg_snapshot_id, fqn + "Cherry-picking Iceberg snapshot %s into table '%s'...", + iceberg_snapshot_id, + fqn, ) self.execute( f"CALL {fqn.catalog}.system.cherrypick_snapshot('{fqn.db}.{fqn.name}', {iceberg_snapshot_id})" ) - self.execute(f"ALTER TABLE {fqn.sql(dialect=self.dialect)} DROP BRANCH {branch_name}") + self.execute( + f"ALTER TABLE {fqn.sql(dialect=self.dialect)} DROP BRANCH {branch_name}" + ) def _ensure_fqn(self, table_name: TableName) -> exp.Table: if isinstance(table_name, exp.Table): @@ -543,7 +571,11 @@ def _ensure_fqn(self, table_name: TableName) -> exp.Table: return table def _build_create_comment_column_exp( - self, table: exp.Table, column_name: str, column_comment: str, table_kind: str = "TABLE" + self, + table: exp.Table, + column_name: str, + column_comment: str, + table_kind: str = "TABLE", ) -> exp.Comment | str: table_sql = table.sql(dialect=self.dialect, identify=True) column_sql = exp.column(column_name).sql(dialect=self.dialect, identify=True) @@ -551,7 +583,9 @@ def _build_create_comment_column_exp( truncated_comment = self._truncate_column_comment(column_comment) comment_sql = exp.Literal.string(truncated_comment).sql(dialect=self.dialect) - return f"ALTER TABLE {table_sql} ALTER COLUMN {column_sql} COMMENT {comment_sql}" + return ( + f"ALTER TABLE {table_sql} ALTER COLUMN {column_sql} COMMENT {comment_sql}" + ) @classmethod def _wap_branch_name(cls, wap_id: str) -> str: diff --git a/sqlmesh/core/engine_adapter/starrocks.py b/sqlmesh/core/engine_adapter/starrocks.py index 05120db0e3..5240ff806e 100644 --- a/sqlmesh/core/engine_adapter/starrocks.py +++ b/sqlmesh/core/engine_adapter/starrocks.py @@ -2,27 +2,19 @@ import logging import re +import typing as t + import sqlglot from sqlglot import exp -import typing as t -from sqlmesh.core.engine_adapter.base import ( - InsertOverwriteStrategy, - get_source_columns_to_types, -) +from sqlmesh.core.engine_adapter.base import (InsertOverwriteStrategy, + get_source_columns_to_types) from sqlmesh.core.engine_adapter.mixins import ( - ClusteredByMixin, - LogicalMergeMixin, - PandasNativeFetchDFSupportMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, - set_catalog, - to_schema, -) + ClusteredByMixin, LogicalMergeMixin, PandasNativeFetchDFSupportMixin) +from sqlmesh.core.engine_adapter.shared import (CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + set_catalog, to_schema) from sqlmesh.core.node import IntervalUnit from sqlmesh.utils.errors import SQLMeshError @@ -113,7 +105,9 @@ def validate(self, value: t.Any) -> t.Optional[Validated]: """Check if value conforms to this type. Return validated value or None. String that can be parsed as literal """ - raise NotImplementedError(f"{self.__class__.__name__}.validate() must be implemented") + raise NotImplementedError( + f"{self.__class__.__name__}.validate() must be implemented" + ) def normalize(self, validated: Validated) -> Normalized: """Convert validated intermediate value to final output format.""" @@ -124,7 +118,9 @@ def __call__(self, value: t.Any) -> Normalized: """Validate and normalize in one step.""" validated = self.validate(value) if validated is None: - raise ValueError(f"Value {value!r} does not conform to type {self.__class__.__name__}") + raise ValueError( + f"Value {value!r} does not conform to type {self.__class__.__name__}" + ) return self.normalize(validated) @@ -342,7 +338,9 @@ def validate(self, value: t.Any) -> t.Optional[t.Tuple[str, t.Any]]: key_name = None if isinstance(left, exp.Column): - key_name = left.this.name if hasattr(left.this, "name") else str(left.this) + key_name = ( + left.this.name if hasattr(left.this, "name") else str(left.this) + ) elif isinstance(left, exp.Identifier): key_name = left.this elif isinstance(left, str): @@ -399,13 +397,18 @@ def __init__( self.case_sensitive = bool(case_sensitive) self.normalized_type = normalized_type - if self.normalized_type is not None and self.normalized_type not in PROPERTY_OUTPUT_TYPES: + if ( + self.normalized_type is not None + and self.normalized_type not in PROPERTY_OUTPUT_TYPES + ): raise ValueError( f"normalized_type must be one of {PROPERTY_OUTPUT_TYPES}, got {self.normalized_type!r}" ) # Pre-compute normalized values for efficient lookup - self._values_normalized = [v if case_sensitive else v.upper() for v in self.valid_values] + self._values_normalized = [ + v if case_sensitive else v.upper() for v in self.valid_values + ] def _extract_text(self, value: t.Any) -> t.Optional[str]: """Extract text from various value types.""" @@ -534,7 +537,9 @@ def __init__(self, *types: DeclarativeType): # Validate all types are DeclarativeType instances for type_ in types: if not isinstance(type_, DeclarativeType): - raise TypeError(f"AnyOf expects DeclarativeType instances, got {type_!r}") + raise TypeError( + f"AnyOf expects DeclarativeType instances, got {type_!r}" + ) self.types: t.List[DeclarativeType] = list(types) @@ -611,7 +616,9 @@ def __init__( self.allow_single = allow_single self.output_as = output_as - def validate(self, value: t.Any) -> t.Optional[t.List[t.Tuple[DeclarativeType, Validated]]]: + def validate( + self, value: t.Any + ) -> t.Optional[t.List[t.Tuple[DeclarativeType, Validated]]]: """Validate each element in the sequence. Returns list of (matched_type, validated_value) tuples or None.""" # Extract elements from various container types elems = self._extract_elements(value) @@ -640,7 +647,9 @@ def normalize( self, validated: t.List[t.Tuple[DeclarativeType, Validated]] ) -> t.Union[t.List[Normalized], t.Tuple[Normalized, ...]]: """Normalize each validated element using its matched type's normalize method.""" - normalized_items = [elem_type.normalize(value) for elem_type, value in validated] + normalized_items = [ + elem_type.normalize(value) for elem_type, value in validated + ] # Convert to desired output format if self.output_as == "tuple": @@ -662,7 +671,9 @@ def _extract_elements(self, value: t.Any) -> t.Optional[t.List[t.Any]]: value = parse_fragment(value) except Exception: # If parsing fails and we accept single strings, promote to list - if self.allow_single and any(isinstance(t, StringType) for t in self.elem_types): + if self.allow_single and any( + isinstance(t, StringType) for t in self.elem_types + ): return [value] return None @@ -768,7 +779,9 @@ class DistributionTupleInputType(StructuredTupleType): FIELDS: t.Dict[str, Field] = {} # Subclasses override this - def __init__(self, error_on_unknown_field: bool = True, error_on_invalid_field: bool = True): + def __init__( + self, error_on_unknown_field: bool = True, error_on_invalid_field: bool = True + ): self.error_on_unknown_field = error_on_unknown_field self.error_on_invalid_field = error_on_invalid_field @@ -1026,7 +1039,9 @@ def validate(self, value: t.Any) -> t.Optional[t.Dict[str, t.Any]]: # ============================================================ @staticmethod - def from_enum(enum_value: str, buckets: t.Optional[int] = None) -> t.Dict[str, t.Any]: + def from_enum( + enum_value: str, buckets: t.Optional[int] = None + ) -> t.Dict[str, t.Any]: """ Create distribution dict from EnumType normalized value. @@ -1062,11 +1077,15 @@ def from_func( >> DistributionTupleOutputType.from_func(func) {"kind": "HASH", "columns": [exp.Column("id"), exp.Column("dt")], "buckets": None} """ - func_name = func.name.upper() if hasattr(func, "name") else str(func.this).upper() + func_name = ( + func.name.upper() if hasattr(func, "name") else str(func.this).upper() + ) if func_name == "HASH": # Extract columns from HASH(col1, col2, ...) - columns: list[exp.Column] = [func.this] if isinstance(func.this, exp.Column) else [] + columns: list[exp.Column] = ( + [func.this] if isinstance(func.this, exp.Column) else [] + ) columns.extend(func.expressions) return {"kind": "HASH", "columns": columns, "buckets": buckets} elif func_name == "RANDOM": # noqa: RET505 @@ -1193,7 +1212,9 @@ class PropertySpecs: RefreshSchemeInputSpec = AnyOf( EnumType(["ASYNC", "MANUAL"], normalized_type="var"), ColumnType(normalized_type="str"), # Columns → will be converted to string - IdentifierType(normalized_type="str"), # Identifiers → will be converted to string + IdentifierType( + normalized_type="str" + ), # Identifiers → will be converted to string LiteralType(normalized_type="str"), # Numbers and string → to string StringType(), # Plain strings ) @@ -1204,8 +1225,12 @@ class PropertySpecs: # So we normalize everything to string for consistent SQL generation GenericPropertyInputSpec = AnyOf( StringType(), # Plain strings - LiteralType(normalized_type="str"), # Numbers and string → will be converted to string - IdentifierType(normalized_type="str"), # Identifiers → will be converted to string + LiteralType( + normalized_type="str" + ), # Numbers and string → will be converted to string + IdentifierType( + normalized_type="str" + ), # Identifiers → will be converted to string ColumnType(normalized_type="str"), # Columns → will be converted to string ) @@ -1298,7 +1323,9 @@ class PropertySpecs: - order_by: List[exp.Expr] - columns - generic properties: str - normalized string values """ - GeneralColumnListOutputSpec: DeclarativeType = SequenceOf(ColumnType(), allow_single=False) + GeneralColumnListOutputSpec: DeclarativeType = SequenceOf( + ColumnType(), allow_single=False + ) PROPERTY_OUTPUT_SPECS: t.Dict[str, DeclarativeType] = { "primary_key": GeneralColumnListOutputSpec, @@ -1430,7 +1457,11 @@ def ensure_parenthesized(value: t.Any) -> t.Any: if isinstance(value, exp.Literal) and value.is_string: value = value.this # Extract string content from Column (quoted) - elif isinstance(value, exp.Column) and hasattr(value.this, "quoted") and value.this.quoted: + elif ( + isinstance(value, exp.Column) + and hasattr(value.this, "quoted") + and value.this.quoted + ): value = value.name # Column.name returns the string elif not isinstance(value, str): return value @@ -1490,7 +1521,9 @@ def validate_and_normalize_property( # Step 3: Validate validated = input_spec.validate(value) if validated is None: - raise SQLMeshError(f"Invalid value type for property '{property_name}': {value!r}.") + raise SQLMeshError( + f"Invalid value type for property '{property_name}': {value!r}." + ) # Step 4: Normalize normalized = input_spec.normalize(validated) @@ -1589,14 +1622,17 @@ def check_at_most_one( Only one is allowed. """ if not exclusive_property_names: - exclusive_property_names = PropertyValidator.EXCLUSIVE_PROPERTY_NAME_MAP.get( - property_name, set() - ) | {property_name} + exclusive_property_names = ( + PropertyValidator.EXCLUSIVE_PROPERTY_NAME_MAP.get(property_name, set()) + | {property_name} + ) # logger.debug("Checking at most one property for '%s': %s", property_name, exclusive_property_names) # Check parameter first (highest priority) if parameter_value is not None: # Check if any conflicting properties exist in table_properties - conflicts = [name for name in exclusive_property_names if name in table_properties] + conflicts = [ + name for name in exclusive_property_names if name in table_properties + ] if conflicts: param_display = f"{property_name} (parameter)" raise SQLMeshError( @@ -1607,7 +1643,9 @@ def check_at_most_one( return None # Check table_properties for multiple definitions - present = [name for name in exclusive_property_names if name in table_properties] + present = [ + name for name in exclusive_property_names if name in table_properties + ] # logger.debug("Get table key names for %s from table_properties: %s", property_name, present) if len(present) > 1: @@ -1799,7 +1837,9 @@ def _get_data_objects( # StarRocks may treat information_schema table_name comparisons as case-sensitive. # Use LOWER(table_name) to match case-insensitively. lowered_names = [name.lower() for name in object_names] - query = query.where(exp.func("LOWER", exp.column("table_name")).isin(*lowered_names)) + query = query.where( + exp.func("LOWER", exp.column("table_name")).isin(*lowered_names) + ) df = self.fetchdf(query) objects = [ @@ -1937,12 +1977,16 @@ def delete_from( # If no where clause or WHERE TRUE, use TRUNCATE TABLE (for all table types) if not where_expr or where_expr == exp.true(): - table_expr = exp.to_table(table_name) if isinstance(table_name, str) else table_name + table_expr = ( + exp.to_table(table_name) if isinstance(table_name, str) else table_name + ) logger.info( f"Converting DELETE FROM {table_name} WHERE TRUE to TRUNCATE TABLE " "(StarRocks does not support WHERE TRUE in DELETE)" ) - self.execute(f"TRUNCATE TABLE {table_expr.sql(dialect=self.dialect, identify=True)}") + self.execute( + f"TRUNCATE TABLE {table_expr.sql(dialect=self.dialect, identify=True)}" + ) return # For non-PRIMARY KEY tables, apply WHERE clause restrictions @@ -1990,10 +2034,14 @@ def transform(node: exp.Expr) -> exp.Expr: # Handle standalone TRUE/FALSE at the top level if node == exp.true(): # Convert TRUE to 1=1 - return exp.EQ(this=exp.Literal.number(1), expression=exp.Literal.number(1)) + return exp.EQ( + this=exp.Literal.number(1), expression=exp.Literal.number(1) + ) elif node == exp.false(): # noqa: RET505 # Convert FALSE to 1=0 - return exp.EQ(this=exp.Literal.number(1), expression=exp.Literal.number(0)) + return exp.EQ( + this=exp.Literal.number(1), expression=exp.Literal.number(0) + ) # Handle AND expressions elif isinstance(node, exp.And): @@ -2022,7 +2070,9 @@ def transform(node: exp.Expr) -> exp.Expr: # Transform the expression tree return expression.transform(transform, copy=True) - def _where_clause_convert_between_to_comparison(self, expression: exp.Expr) -> exp.Expr: + def _where_clause_convert_between_to_comparison( + self, expression: exp.Expr + ) -> exp.Expr: """ Convert BETWEEN expressions to >= AND <= comparisons. @@ -2152,7 +2202,9 @@ def adjust_physical_properties_for_incremental( # statements remain supported. if unique_key: physical_properties["primary_key"] = ( - unique_key[0] if len(unique_key) == 1 else exp.Tuple(expressions=unique_key) + unique_key[0] + if len(unique_key) == 1 + else exp.Tuple(expressions=unique_key) ) logger.info( "Model '%s' promoted to PRIMARY KEY table on StarRocks to support rich DELETE operations.", @@ -2254,7 +2306,9 @@ def _create_table_from_columns( # it's passed as a model parameter rather than in physical_properties if primary_key: table_properties["primary_key"] = primary_key - logger.debug("_create_table_from_columns: unified primary_key into table_properties") + logger.debug( + "_create_table_from_columns: unified primary_key into table_properties" + ) elif key_type: # logger.debug( # "table key type '%s' may be handled in _build_table_key_property", key_type @@ -2392,21 +2446,27 @@ def _create_materialized_view( batch_end=len(values), ) - source_queries, target_columns_to_types = self._get_source_queries_and_columns_to_types( - query_or_df, - target_columns_to_types, - batch_size=0, - target_table=view_name, - source_columns=source_columns, + source_queries, target_columns_to_types = ( + self._get_source_queries_and_columns_to_types( + query_or_df, + target_columns_to_types, + batch_size=0, + target_table=view_name, + source_columns=source_columns, + ) ) if len(source_queries) != 1: - raise SQLMeshError("Only one source query is supported for creating materialized views") + raise SQLMeshError( + "Only one source query is supported for creating materialized views" + ) target_table = exp.to_table(view_name) - schema: t.Union[exp.Table, exp.Schema] = self._build_materialized_view_schema_exp( - target_table, - target_columns_to_types=target_columns_to_types, - column_descriptions=column_descriptions, + schema: t.Union[exp.Table, exp.Schema] = ( + self._build_materialized_view_schema_exp( + target_table, + target_columns_to_types=target_columns_to_types, + column_descriptions=column_descriptions, + ) ) # Pass model materialized properties through the existing properties builder @@ -2416,7 +2476,9 @@ def _create_materialized_view( if materialized_properties: partitioned_by = materialized_properties.get("partitioned_by") clustered_by = materialized_properties.get("clustered_by") - partition_interval_unit = materialized_properties.get("partition_interval_unit") + partition_interval_unit = materialized_properties.get( + "partition_interval_unit" + ) # logger.debug( # f"Get info from materialized_properties: {materialized_properties}, " # f"partitioned_by: {partitioned_by}, " @@ -2488,7 +2550,9 @@ def _build_materialized_view_schema_exp( constraints.append( exp.ColumnConstraint( kind=exp.CommentColumnConstraint( - this=exp.Literal.string(self._truncate_column_comment(comment)) + this=exp.Literal.string( + self._truncate_column_comment(comment) + ) ) ) ) @@ -2594,7 +2658,9 @@ def _build_table_properties_exp( key_columns = tuple(col.name for col in normalized) # 1. Handle key constraints (ALL types including PRIMARY KEY) - key_prop = self._build_table_key_property(table_properties_copy, active_key_type) + key_prop = self._build_table_key_property( + table_properties_copy, active_key_type + ) if key_prop: properties.append(key_prop) @@ -2602,7 +2668,9 @@ def _build_table_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) @@ -2620,7 +2688,9 @@ def _build_table_properties_exp( properties.append(partition_prop) # 4. Handle distributed_by (DISTRIBUTED BY HASH/RANDOM) - distributed_prop = self._build_distributed_by_property(table_properties_copy, key_columns) + distributed_prop = self._build_distributed_by_property( + table_properties_copy, key_columns + ) if distributed_prop: properties.append(distributed_prop) @@ -2639,7 +2709,9 @@ def _build_table_properties_exp( properties.append(refresh_prop) # 6. Handle order_by/clustered_by (ORDER BY ...) - order_prop = self._build_order_by_property(table_properties_copy, clustered_by or None) + order_prop = self._build_order_by_property( + table_properties_copy, clustered_by or None + ) if order_prop: properties.append(order_prop) @@ -2666,7 +2738,9 @@ def _build_view_properties_exp( if table_description: properties.append( exp.SchemaCommentProperty( - this=exp.Literal.string(self._truncate_table_comment(table_description)) + this=exp.Literal.string( + self._truncate_table_comment(table_description) + ) ) ) @@ -2678,9 +2752,13 @@ def _build_view_properties_exp( "security", security ) # exp.SqlSecurityProperty renders as `SECURITY ` (no '=') - properties.append(exp.SqlSecurityProperty(this=exp.Var(this=security_text))) + properties.append( + exp.SqlSecurityProperty(this=exp.Var(this=security_text)) + ) - properties.extend(self._table_or_view_properties_to_expressions(view_properties_copy)) + properties.extend( + self._table_or_view_properties_to_expressions(view_properties_copy) + ) if properties: return exp.Properties(expressions=properties) @@ -2800,7 +2878,9 @@ def _build_partition_property( return None # Parse partition expressions to extract columns and kind (RANGE/LIST) - partition_kind, partition_cols = self._parse_partition_expressions(partitioned_by) + partition_kind, partition_cols = self._parse_partition_expressions( + partitioned_by + ) logger.debug( "_build_partition_property: partition_kind=%s, partition_cols=%s", partition_kind, @@ -2817,7 +2897,9 @@ def extract_column_name(expr: exp.Expr) -> t.Optional[str]: # Validate partition columns are in key columns (StarRocks requirement) if key_columns: - partition_col_names = set(extract_column_name(expr) for expr in partition_cols) - {None} + partition_col_names = set( + extract_column_name(expr) for expr in partition_cols + ) - {None} key_cols_set = set(key_columns) not_in_key = partition_col_names - key_cols_set if not_in_key: @@ -2967,7 +3049,9 @@ def _build_partitioned_by_exp( create_expressions=create_expressions, ) elif partition_kind is None: - return exp.PartitionedByProperty(this=exp.Schema(expressions=partitioned_by)) + return exp.PartitionedByProperty( + this=exp.Schema(expressions=partitioned_by) + ) return None @@ -3063,7 +3147,9 @@ def _validate_deferred_refresh_for_audits( """ refresh_moment = (view_properties or {}).get("refresh_moment") normalized = ( - PropertyValidator.validate_and_normalize_property("refresh_moment", refresh_moment) + PropertyValidator.validate_and_normalize_property( + "refresh_moment", refresh_moment + ) if refresh_moment is not None else None ) @@ -3120,8 +3206,8 @@ def _build_refresh_property( if isinstance(scheme_text, exp.Var): kind_expr = scheme_text else: - kind_expr, starts_expr, every_expr, unit_expr = self._parse_refresh_scheme( - scheme_text + kind_expr, starts_expr, every_expr, unit_expr = ( + self._parse_refresh_scheme(scheme_text) ) return exp.RefreshTriggerProperty( @@ -3132,9 +3218,7 @@ def _build_refresh_property( unit=unit_expr, ) - def _parse_refresh_scheme( - self, refresh_scheme: str - ) -> t.Tuple[ + def _parse_refresh_scheme(self, refresh_scheme: str) -> t.Tuple[ t.Optional[exp.Expr], t.Optional[exp.Expr], t.Optional[exp.Expr], @@ -3164,10 +3248,14 @@ def _parse_refresh_scheme( every_expr: t.Optional[exp.Expr] = None unit_expr: t.Optional[exp.Expr] = None m_start = re.search( - r"\bSTART\s*\(\s*(?:'([^']*)'|\"([^\"]*)\"|([^)]*))\s*\)", text, flags=re.IGNORECASE + r"\bSTART\s*\(\s*(?:'([^']*)'|\"([^\"]*)\"|([^)]*))\s*\)", + text, + flags=re.IGNORECASE, ) if m_start: - start_inner = (m_start.group(1) or m_start.group(2) or m_start.group(3) or "").strip() + start_inner = ( + m_start.group(1) or m_start.group(2) or m_start.group(3) or "" + ).strip() starts_expr = exp.Literal.string(start_inner) m_every = re.search( r"\bEVERY\s*\(\s*INTERVAL\s+(\d+)\s+(\w+)\s*\)", text, flags=re.IGNORECASE @@ -3211,7 +3299,9 @@ def _parse_distribution_with_buckets( return None # Split on BUCKETS (case-insensitive) - match = re.match(r"^(.+?)\s+BUCKETS\s+(\d+)\s*$", text.strip(), flags=re.IGNORECASE) + match = re.match( + r"^(.+?)\s+BUCKETS\s+(\d+)\s*$", text.strip(), flags=re.IGNORECASE + ) if not match: return None @@ -3219,7 +3309,9 @@ def _parse_distribution_with_buckets( buckets_str = match.group(2) # Parse the HASH/RANDOM part via SPEC - normalized = PropertyValidator.validate_and_normalize_property("distributed_by", hash_part) + normalized = PropertyValidator.validate_and_normalize_property( + "distributed_by", hash_part + ) return DistributionTupleOutputType.to_unified_dict(normalized, int(buckets_str)) @@ -3269,7 +3361,9 @@ def _build_order_by_property( else: # noqa: RET505 return None - def _build_other_properties(self, table_properties: t.Dict[str, t.Any]) -> t.List[exp.Property]: + def _build_other_properties( + self, table_properties: t.Dict[str, t.Any] + ) -> t.List[exp.Property]: """ Build other literal properties (replication_num, storage_medium, etc.). @@ -3287,7 +3381,9 @@ def _build_other_properties(self, table_properties: t.Dict[str, t.Any]) -> t.Lis for key, value in list(table_properties.items()): # Skip special keys handled elsewhere if key in PropertyValidator.IMPORTANT_PROPERTY_NAMES: - logger.warning(f"[StarRocks] {key!r} should have been processed already, skipping") + logger.warning( + f"[StarRocks] {key!r} should have been processed already, skipping" + ) continue # Remove from properties @@ -3296,7 +3392,9 @@ def _build_other_properties(self, table_properties: t.Dict[str, t.Any]) -> t.Lis # Validate and normalize to string # All other properties are treated as generic string properties try: - normalized = PropertyValidator.validate_and_normalize_property(key, value) + normalized = PropertyValidator.validate_and_normalize_property( + key, value + ) other_props.append( exp.Property( this=exp.to_identifier(key), @@ -3304,7 +3402,9 @@ def _build_other_properties(self, table_properties: t.Dict[str, t.Any]) -> t.Lis ) ) except SQLMeshError as e: - logger.warning("[StarRocks] skipping property %s due to error: %s", key, e) + logger.warning( + "[StarRocks] skipping property %s due to error: %s", key, e + ) return other_props @@ -3449,9 +3549,9 @@ def _build_create_comment_table_exp( SQL string for ALTER TABLE COMMENT """ table_sql = table.sql(dialect=self.dialect, identify=True) - comment_sql = exp.Literal.string(self._truncate_table_comment(table_comment)).sql( - dialect=self.dialect - ) + comment_sql = exp.Literal.string( + self._truncate_table_comment(table_comment) + ).sql(dialect=self.dialect) return f"ALTER TABLE {table_sql} COMMENT = {comment_sql}" def _build_create_comment_column_exp( @@ -3482,13 +3582,17 @@ def _build_create_comment_column_exp( SQL string for ALTER TABLE MODIFY COLUMN with COMMENT """ table_sql = table.sql(dialect=self.dialect, identify=True) - column_sql = exp.to_identifier(column_name).sql(dialect=self.dialect, identify=True) - - comment_sql = exp.Literal.string(self._truncate_column_comment(column_comment)).sql( - dialect=self.dialect + column_sql = exp.to_identifier(column_name).sql( + dialect=self.dialect, identify=True ) - return f"ALTER TABLE {table_sql} MODIFY COLUMN {column_sql} COMMENT {comment_sql}" + comment_sql = exp.Literal.string( + self._truncate_column_comment(column_comment) + ).sql(dialect=self.dialect) + + return ( + f"ALTER TABLE {table_sql} MODIFY COLUMN {column_sql} COMMENT {comment_sql}" + ) # ==================== Methods NOT Needing Override (Base Class Works) ==================== # The following methods work correctly with base class implementation: diff --git a/sqlmesh/core/engine_adapter/trino.py b/sqlmesh/core/engine_adapter/trino.py index 00acddb26c..d2b0ff289b 100644 --- a/sqlmesh/core/engine_adapter/trino.py +++ b/sqlmesh/core/engine_adapter/trino.py @@ -7,28 +7,21 @@ from sqlglot import exp from sqlglot.helper import seq_get -from tenacity import retry, wait_fixed, stop_after_attempt, retry_if_result +from tenacity import retry, retry_if_result, stop_after_attempt, wait_fixed from sqlmesh.core.dialect import schema_, to_schema from sqlmesh.core.engine_adapter.mixins import ( - GetCurrentCatalogFromFunctionMixin, - HiveMetastoreTablePropertiesMixin, - PandasNativeFetchDFSupportMixin, - RowDiffMixin, -) -from sqlmesh.core.engine_adapter.shared import ( - CatalogSupport, - CommentCreationTable, - CommentCreationView, - DataObject, - DataObjectType, - InsertOverwriteStrategy, - SourceQuery, - set_catalog, -) + GetCurrentCatalogFromFunctionMixin, HiveMetastoreTablePropertiesMixin, + PandasNativeFetchDFSupportMixin, RowDiffMixin) +from sqlmesh.core.engine_adapter.shared import (CatalogSupport, + CommentCreationTable, + CommentCreationView, + DataObject, DataObjectType, + InsertOverwriteStrategy, + SourceQuery, set_catalog) from sqlmesh.utils import get_source_columns_to_types -from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.date import TimeLike +from sqlmesh.utils.errors import SQLMeshError if t.TYPE_CHECKING: from sqlmesh.core._typing import SchemaName, SessionProperties, TableName @@ -191,11 +184,15 @@ def _insert_overwrite_by_condition( # These session properties are only valid for the Trino Hive connector # Attempting to set them on an Iceberg catalog will throw an error: # "Session property 'catalog.insert_existing_partitions_behavior' does not exist" - self.execute(f"SET SESSION {catalog}.insert_existing_partitions_behavior='OVERWRITE'") + self.execute( + f"SET SESSION {catalog}.insert_existing_partitions_behavior='OVERWRITE'" + ) super()._insert_overwrite_by_condition( table_name, source_queries, target_columns_to_types, where ) - self.execute(f"SET SESSION {catalog}.insert_existing_partitions_behavior='APPEND'") + self.execute( + f"SET SESSION {catalog}.insert_existing_partitions_behavior='APPEND'" + ) else: super()._insert_overwrite_by_condition( table_name, @@ -243,8 +240,12 @@ def _get_data_objects( exp.column("catalog_name", table="mv").eq( exp.column("table_catalog", table="t") ), - exp.column("schema_name", table="mv").eq(exp.column("table_schema", table="t")), - exp.column("name", table="mv").eq(exp.column("table_name", table="t")), + exp.column("schema_name", table="mv").eq( + exp.column("table_schema", table="t") + ), + exp.column("name", table="mv").eq( + exp.column("table_name", table="t") + ), ), join_type="left", ) @@ -296,11 +297,18 @@ def _df_to_source_queries( # timestamp in Trino. for column, kind in source_columns_to_types.items(): dtype = df.dtypes[column] - if is_datetime64_any_dtype(dtype) and getattr(dtype, "tz", None) is not None: + if ( + is_datetime64_any_dtype(dtype) + and getattr(dtype, "tz", None) is not None + ): df[column] = pd.to_datetime(df[column]).map(lambda x: x.isoformat(" ")) return super()._df_to_source_queries( - df, target_columns_to_types, batch_size, target_table, source_columns=source_columns + df, + target_columns_to_types, + batch_size, + target_table, + source_columns=source_columns, ) def _build_schema_exp( @@ -316,7 +324,9 @@ def _build_schema_exp( target_columns_to_types ) if "delta_lake" in self.get_catalog_type_from_table(table): - target_columns_to_types = self._to_delta_ts(target_columns_to_types, mapped_columns) + target_columns_to_types = self._to_delta_ts( + target_columns_to_types, mapped_columns + ) return super()._build_schema_exp( table, target_columns_to_types, column_descriptions, expressions, is_view @@ -347,10 +357,13 @@ def _scd_type_2( target_columns_to_types, mapped_columns = self._apply_timestamp_mapping( target_columns_to_types ) - if target_columns_to_types and "delta_lake" in self.get_catalog_type_from_table( - target_table + if ( + target_columns_to_types + and "delta_lake" in self.get_catalog_type_from_table(target_table) ): - target_columns_to_types = self._to_delta_ts(target_columns_to_types, mapped_columns) + target_columns_to_types = self._to_delta_ts( + target_columns_to_types, mapped_columns + ) return super()._scd_type_2( target_table, @@ -394,13 +407,21 @@ def _to_delta_ts( } delta_columns_to_types = { - k: ts3_tz if k not in skip and v.is_type(exp.DataType.Type.TIMESTAMPTZ) else v + k: ( + ts3_tz + if k not in skip and v.is_type(exp.DataType.Type.TIMESTAMPTZ) + else v + ) for k, v in delta_columns_to_types.items() } return delta_columns_to_types - @retry(wait=wait_fixed(1), stop=stop_after_attempt(10), retry=retry_if_result(lambda v: not v)) + @retry( + wait=wait_fixed(1), + stop=stop_after_attempt(10), + retry=retry_if_result(lambda v: not v), + ) def _block_until_table_exists(self, table_name: TableName) -> bool: return self.table_exists(table_name) @@ -413,7 +434,9 @@ def _create_schema( kind: str, ) -> None: if mapped_location := self._schema_location(schema_name): - properties.append(exp.LocationProperty(this=exp.Literal.string(mapped_location))) + properties.append( + exp.LocationProperty(this=exp.Literal.string(mapped_location)) + ) return super()._create_schema( schema_name=schema_name, diff --git a/sqlmesh/core/environment.py b/sqlmesh/core/environment.py index 4594dc120d..3106762fa4 100644 --- a/sqlmesh/core/environment.py +++ b/sqlmesh/core/environment.py @@ -11,13 +11,14 @@ from sqlmesh.core.engine_adapter.base import EngineAdapter from sqlmesh.core.macros import RuntimeStage from sqlmesh.core.renderer import render_statements -from sqlmesh.core.snapshot import SnapshotId, SnapshotTableInfo, Snapshot +from sqlmesh.core.snapshot import Snapshot, SnapshotId, SnapshotTableInfo from sqlmesh.utils import word_characters_only from sqlmesh.utils.date import TimeLike, now_timestamp from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.jinja import JinjaMacroRegistry from sqlmesh.utils.metaprogramming import Executable -from sqlmesh.utils.pydantic import PydanticModel, field_validator, ValidationInfo +from sqlmesh.utils.pydantic import (PydanticModel, ValidationInfo, + field_validator) T = t.TypeVar("T", bound="EnvironmentNamingInfo") PydanticType = t.TypeVar("PydanticType", bound="PydanticModel") @@ -38,7 +39,9 @@ class EnvironmentNamingInfo(PydanticModel): """ name: str = c.PROD - suffix_target: EnvironmentSuffixTarget = Field(default=EnvironmentSuffixTarget.SCHEMA) + suffix_target: EnvironmentSuffixTarget = Field( + default=EnvironmentSuffixTarget.SCHEMA + ) catalog_name_override: t.Optional[str] = None normalize_name: bool = True gateway_managed: bool = False @@ -175,14 +178,19 @@ def _load_requirements(cls, v: t.Any) -> t.Any: @property def snapshots(self) -> t.List[SnapshotTableInfo]: - return self._convert_list_to_models_and_store("snapshots_", SnapshotTableInfo) or [] + return ( + self._convert_list_to_models_and_store("snapshots_", SnapshotTableInfo) + or [] + ) def snapshot_dicts(self) -> t.List[dict]: return self._convert_list_to_dicts(self.snapshots_) @property def promoted_snapshot_ids(self) -> t.Optional[t.List[SnapshotId]]: - return self._convert_list_to_models_and_store("promoted_snapshot_ids_", SnapshotId) + return self._convert_list_to_models_and_store( + "promoted_snapshot_ids_", SnapshotId + ) def promoted_snapshot_id_dicts(self) -> t.List[dict]: return self._convert_list_to_dicts(self.promoted_snapshot_ids_) @@ -275,7 +283,9 @@ def render_before_all( default_catalog: t.Optional[str] = None, **render_kwargs: t.Any, ) -> t.List[str]: - return self.render(RuntimeStage.BEFORE_ALL, dialect, default_catalog, **render_kwargs) + return self.render( + RuntimeStage.BEFORE_ALL, dialect, default_catalog, **render_kwargs + ) def render_after_all( self, @@ -283,7 +293,9 @@ def render_after_all( default_catalog: t.Optional[str] = None, **render_kwargs: t.Any, ) -> t.List[str]: - return self.render(RuntimeStage.AFTER_ALL, dialect, default_catalog, **render_kwargs) + return self.render( + RuntimeStage.AFTER_ALL, dialect, default_catalog, **render_kwargs + ) def render( self, diff --git a/sqlmesh/core/janitor.py b/sqlmesh/core/janitor.py index 92d889e276..df7643ec56 100644 --- a/sqlmesh/core/janitor.py +++ b/sqlmesh/core/janitor.py @@ -4,18 +4,15 @@ from sqlglot import exp -from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.console import Console from sqlmesh.core.dialect import schema_ +from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.environment import Environment from sqlmesh.core.snapshot import SnapshotEvaluator from sqlmesh.core.state_sync import StateSync -from sqlmesh.core.state_sync.common import ( - logger, - iter_expired_snapshot_batches, - RowBoundary, - ExpiredBatchRange, -) +from sqlmesh.core.state_sync.common import (ExpiredBatchRange, RowBoundary, + iter_expired_snapshot_batches, + logger) def cleanup_expired_views( @@ -32,11 +29,15 @@ def cleanup_expired_views( if environment.suffix_target.is_schema or environment.suffix_target.is_catalog ] expired_table_environments = [ - environment for environment in environments if environment.suffix_target.is_table + 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: + def get_adapter( + gateway_managed: bool, gateway: t.Optional[str] = None + ) -> EngineAdapter: if gateway_managed and gateway: return engine_adapters.get(gateway, default_adapter) return default_adapter @@ -47,7 +48,11 @@ def get_adapter(gateway_managed: bool, gateway: t.Optional[str] = None) -> Engin # Collect schemas and catalogs to drop for engine_adapter, expired_catalog, expired_schema, suffix_target in { ( - (engine_adapter := get_adapter(environment.gateway_managed, snapshot.model_gateway)), + ( + engine_adapter := get_adapter( + environment.gateway_managed, snapshot.model_gateway + ) + ), snapshot.qualified_view_name.catalog_for_environment( environment.naming_info, dialect=engine_adapter.dialect ), @@ -70,7 +75,11 @@ def get_adapter(gateway_managed: bool, gateway: t.Optional[str] = None) -> Engin # Drop the views for the expired environments for engine_adapter, expired_view in { ( - (engine_adapter := get_adapter(environment.gateway_managed, snapshot.model_gateway)), + ( + engine_adapter := get_adapter( + environment.gateway_managed, snapshot.model_gateway + ) + ), snapshot.qualified_view_name.for_environment( environment.naming_info, dialect=engine_adapter.dialect ), @@ -84,7 +93,9 @@ def get_adapter(gateway_managed: bool, gateway: t.Optional[str] = None) -> Engin if console: console.update_cleanup_progress(expired_view) except Exception as e: - message = f"Failed to drop the expired environment view '{expired_view}': {e}" + message = ( + f"Failed to drop the expired environment view '{expired_view}': {e}" + ) logger.warning(message) failures.append(message) @@ -97,7 +108,9 @@ def get_adapter(gateway_managed: bool, gateway: t.Optional[str] = None) -> Engin cascade=True, ) if console: - console.update_cleanup_progress(schema.sql(dialect=engine_adapter.dialect)) + console.update_cleanup_progress( + schema.sql(dialect=engine_adapter.dialect) + ) except Exception as e: message = f"Failed to drop the expired environment schema '{schema}': {e}" logger.warning(message) @@ -112,7 +125,9 @@ def get_adapter(gateway_managed: bool, gateway: t.Optional[str] = None) -> Engin if console: console.update_cleanup_progress(catalog) except Exception as e: - message = f"Failed to drop the expired environment catalog '{catalog}': {e}" + message = ( + f"Failed to drop the expired environment catalog '{catalog}': {e}" + ) logger.warning(message) failures.append(message) diff --git a/sqlmesh/core/lineage.py b/sqlmesh/core/lineage.py index 8363979034..a5de98efc5 100644 --- a/sqlmesh/core/lineage.py +++ b/sqlmesh/core/lineage.py @@ -51,7 +51,9 @@ def lineage( query = qualify.qualify( query, dialect=model.dialect, - schema=normalize_mapping_schema(model.mapping_schema, dialect=model.dialect), + schema=normalize_mapping_schema( + model.mapping_schema, dialect=model.dialect + ), **{"validate_qualify_columns": False, "infer_schema": True, **kwargs}, ) @@ -101,7 +103,9 @@ def column_description( if column in model.column_descriptions: return model.column_descriptions[column] - dependencies = column_dependencies(context, model_name, exp.column(column, quoted=quote_column)) + dependencies = column_dependencies( + context, model_name, exp.column(column, quoted=quote_column) + ) if len(dependencies) != 1: return None diff --git a/sqlmesh/core/linter/definition.py b/sqlmesh/core/linter/definition.py index 7dc64bbf95..1e8f5386cb 100644 --- a/sqlmesh/core/linter/definition.py +++ b/sqlmesh/core/linter/definition.py @@ -2,12 +2,12 @@ import operator as op import typing as t -from collections.abc import Iterator, Iterable, Set, Mapping, Callable +from collections.abc import Callable, Iterable, Iterator, Mapping, Set from functools import reduce from sqlmesh.core.config.linter import LinterConfig from sqlmesh.core.console import LinterConsole, get_console -from sqlmesh.core.linter.rule import Rule, RuleViolation, Range, Fix +from sqlmesh.core.linter.rule import Fix, Range, Rule, RuleViolation from sqlmesh.core.model import Model from sqlmesh.utils.errors import raise_config_error @@ -55,7 +55,10 @@ def from_rules(cls, all_rules: RuleSet, config: LinterConfig) -> Linter: return Linter(config.enabled, all_rules, rules, warn_rules) def lint_model( - self, model: Model, context: GenericContext, console: LinterConsole = get_console() + self, + model: Model, + context: GenericContext, + console: LinterConsole = get_console(), ) -> t.Tuple[bool, t.List[AnnotatedRuleViolation]]: if not self.enabled: return False, [] @@ -103,7 +106,9 @@ class RuleSet(Mapping[str, type[Rule]]): def __init__(self, rules: Iterable[type[Rule]] = ()) -> None: self._underlying = {rule.name: rule for rule in rules} - def check_model(self, model: Model, context: GenericContext) -> t.List[RuleViolation]: + def check_model( + self, model: Model, context: GenericContext + ) -> t.List[RuleViolation]: violations = [] for rule in self._underlying.values(): diff --git a/sqlmesh/core/linter/helpers.py b/sqlmesh/core/linter/helpers.py index 3c79f83a43..57eeade200 100644 --- a/sqlmesh/core/linter/helpers.py +++ b/sqlmesh/core/linter/helpers.py @@ -1,9 +1,10 @@ +import typing as t from pathlib import Path -from sqlmesh.core.linter.rule import Range, Position +from sqlglot import Token, TokenType, tokenize + +from sqlmesh.core.linter.rule import Position, Range from sqlmesh.utils.pydantic import PydanticModel -from sqlglot import tokenize, TokenType, Token -import typing as t class TokenPositionDetails(PydanticModel): @@ -53,7 +54,9 @@ def to_range(self, read_file: t.Optional[t.List[str]]) -> Range: ) if read_file is None: - raise ValueError("read_file must be provided when start and end positions differ.") + raise ValueError( + "read_file must be provided when start and end positions differ." + ) # Convert from 1-indexed to 0-indexed for line only end_line_0 = self.line - 1 @@ -178,7 +181,7 @@ def get_range_of_model_block( block = get_start_and_end_of_model_block(tokens) if not block: return None - (start_idx, end_idx) = block + start_idx, end_idx = block start = tokens[start_idx - 1] end = tokens[end_idx + 1] start_position = TokenPositionDetails( @@ -217,7 +220,7 @@ def get_range_of_a_key_in_model_block( block = get_start_and_end_of_model_block(tokens) if not block: return None - (lparen_idx, rparen_idx) = block + lparen_idx, rparen_idx = block # 4) Scan within the MODEL property list for the key at top-level (depth == 1) # Initialize depth to 1 since we're inside the first parentheses @@ -275,7 +278,8 @@ def is_close(t: TokenType) -> bool: # End of value: at top-level (nested == 0) encountering a comma or the end paren if nested == 0 and ( - ttype is TokenType.COMMA or (ttype is TokenType.R_PAREN and depth == 1) + ttype is TokenType.COMMA + or (ttype is TokenType.R_PAREN and depth == 1) ): # For comma, don't include it in the value range # For closing paren, include it only if it's part of the value structure diff --git a/sqlmesh/core/linter/rule.py b/sqlmesh/core/linter/rule.py index 8dd1a2ebbd..24cf62ad58 100644 --- a/sqlmesh/core/linter/rule.py +++ b/sqlmesh/core/linter/rule.py @@ -1,18 +1,14 @@ from __future__ import annotations import abc +import typing as t from dataclasses import dataclass, field from pathlib import Path - -from sqlmesh.core.model import Model - from typing import Type -import typing as t - +from sqlmesh.core.model import Model from sqlmesh.utils.pydantic import PydanticModel - if t.TYPE_CHECKING: from sqlmesh.core.context import GenericContext diff --git a/sqlmesh/core/linter/rules/builtin.py b/sqlmesh/core/linter/rules/builtin.py index 8dc4172f9f..87d13944fb 100644 --- a/sqlmesh/core/linter/rules/builtin.py +++ b/sqlmesh/core/linter/rules/builtin.py @@ -9,23 +9,15 @@ from sqlmesh.core.constants import EXTERNAL_MODELS_YAML from sqlmesh.core.dialect import normalize_model_name -from sqlmesh.core.linter.helpers import ( - TokenPositionDetails, - get_range_of_model_block, - read_range_from_string, -) -from sqlmesh.core.linter.rule import ( - Rule, - RuleViolation, - Range, - Fix, - TextEdit, - Position, - CreateFile, -) from sqlmesh.core.linter.definition import RuleSet -from sqlmesh.core.model import Model, SqlModel, ExternalModel -from sqlmesh.utils.lineage import extract_references_from_query, ExternalModelReference +from sqlmesh.core.linter.helpers import (TokenPositionDetails, + get_range_of_model_block, + read_range_from_string) +from sqlmesh.core.linter.rule import (CreateFile, Fix, Position, Range, Rule, + RuleViolation, TextEdit) +from sqlmesh.core.model import ExternalModel, Model, SqlModel +from sqlmesh.utils.lineage import (ExternalModelReference, + extract_references_from_query) class NoSelectStar(Rule): @@ -44,10 +36,12 @@ def check_model(self, model: Model) -> t.Optional[RuleViolation]: def _get_range(self, model: SqlModel) -> t.Optional[Range]: """Get the range of the violation if available.""" try: - if len(model.query.expressions) == 1 and isinstance(model.query.expressions[0], Star): - return TokenPositionDetails.from_meta(model.query.expressions[0].meta).to_range( - None - ) + if len(model.query.expressions) == 1 and isinstance( + model.query.expressions[0], Star + ): + return TokenPositionDetails.from_meta( + model.query.expressions[0].meta + ).to_range(None) except Exception: pass @@ -101,9 +95,7 @@ def check_model(self, model: Model) -> t.Optional[RuleViolation]: if not sqlglot_err: return None - violation_msg = ( - f"{sqlglot_err} for model '{model.fqn}', the column may not exist or is ambiguous." - ) + violation_msg = f"{sqlglot_err} for model '{model.fqn}', the column may not exist or is ambiguous." return self.violation(violation_msg) @@ -168,7 +160,11 @@ def check_model( # If the model is anything other than a sql model that and has a path # that ends with .sql, we cannot extract the references from the query. path = model._path - if not isinstance(model, SqlModel) or not path or not str(path).endswith(".sql"): + if ( + not isinstance(model, SqlModel) + or not path + or not str(path).endswith(".sql") + ): return self._standard_error_message( model_name=model.fqn, external_models=not_registered_external_models, @@ -274,7 +270,8 @@ def create_fix(self, model_name: str) -> t.Optional[Fix]: else: new_text = f"\n- name: '{model_name}'\n" position = Position( - line=len(split_lines) - 1, character=len(split_lines[-1]) if split_lines else 0 + line=len(split_lines) - 1, + character=len(split_lines[-1]) if split_lines else 0, ) return Fix( diff --git a/sqlmesh/core/loader.py b/sqlmesh/core/loader.py index cb951b4f9e..c4c9fd9999 100644 --- a/sqlmesh/core/loader.py +++ b/sqlmesh/core/loader.py @@ -1,6 +1,7 @@ from __future__ import annotations import abc +import concurrent.futures import glob import itertools import linecache @@ -10,28 +11,25 @@ from collections import Counter, defaultdict from dataclasses import dataclass from pathlib import Path -from pydantic import ValidationError -import concurrent.futures -from sqlglot.errors import SqlglotError +from pydantic import ValidationError from sqlglot import exp +from sqlglot.errors import SqlglotError from sqlglot.helper import subclasses from sqlmesh.core import constants as c -from sqlmesh.core.audit import Audit, ModelAudit, StandaloneAudit, load_multiple_audits +from sqlmesh.core.audit import (Audit, ModelAudit, StandaloneAudit, + load_multiple_audits) from sqlmesh.core.console import Console from sqlmesh.core.dialect import parse from sqlmesh.core.environment import EnvironmentStatements -from sqlmesh.core.linter.rule import Rule from sqlmesh.core.linter.definition import RuleSet +from sqlmesh.core.linter.rule import Rule from sqlmesh.core.macros import MacroRegistry, macro -from sqlmesh.core.metric import Metric, MetricMeta, expand_metrics, load_metric_ddl -from sqlmesh.core.model import ( - Model, - ModelCache, - create_external_model, - load_sql_based_models, -) +from sqlmesh.core.metric import (Metric, MetricMeta, expand_metrics, + load_metric_ddl) +from sqlmesh.core.model import (Model, ModelCache, create_external_model, + load_sql_based_models) from sqlmesh.core.model import model as model_registry from sqlmesh.core.model.common import make_python_env from sqlmesh.core.signal import signal @@ -40,10 +38,10 @@ from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.jinja import JinjaMacroRegistry, MacroExtractor from sqlmesh.utils.metaprogramming import import_python_file -from sqlmesh.utils.pydantic import validation_error_message from sqlmesh.utils.process import create_process_pool_executor -from sqlmesh.utils.yaml import YAML, load as yaml_load - +from sqlmesh.utils.pydantic import validation_error_message +from sqlmesh.utils.yaml import YAML +from sqlmesh.utils.yaml import load as yaml_load if t.TYPE_CHECKING: from sqlmesh.core.context import GenericContext @@ -214,7 +212,9 @@ def load(self) -> LoadedProject: self._track_file(config_file) config_mtimes[c.SQLMESH_PATH].append(self._path_mtimes[config_file]) - self._config_mtimes = {path: max(mtimes) for path, mtimes in config_mtimes.items()} + self._config_mtimes = { + path: max(mtimes) for path, mtimes in config_mtimes.items() + } macros, jinja_macros = self._load_scripts() audits: UniqueKeyDict[str, ModelAudit] = UniqueKeyDict("audits") @@ -222,7 +222,9 @@ def load(self) -> LoadedProject: "standalone_audits" ) - for name, audit in self._load_audits(macros=macros, jinja_macros=jinja_macros).items(): + for name, audit in self._load_audits( + macros=macros, jinja_macros=jinja_macros + ).items(): if isinstance(audit, ModelAudit): audits[name] = audit else: @@ -295,7 +297,9 @@ def _load_audits( ) -> UniqueKeyDict[str, Audit]: """Loads all audits.""" - def _load_environment_statements(self, macros: MacroRegistry) -> t.List[EnvironmentStatements]: + def _load_environment_statements( + self, macros: MacroRegistry + ) -> t.List[EnvironmentStatements]: """Loads environment statements.""" return [] @@ -329,7 +333,9 @@ def _load_external_models( paths_to_load.append(deprecated_yaml) if external_models_path.exists() and external_models_path.is_dir(): - paths_to_load.extend(self._glob_paths(external_models_path, extension=".yaml")) + paths_to_load.extend( + self._glob_paths(external_models_path, extension=".yaml") + ) def _load(path: Path) -> t.List[Model]: try: @@ -380,7 +386,8 @@ def _load(path: Path) -> t.List[Model]: if model.fqn in models and models[model.fqn].gateway == gateway: raise ConfigError( self._failed_to_load_model_error( - path, f"Duplicate external model name: '{model.name}'." + path, + f"Duplicate external model name: '{model.name}'.", ), path, ) @@ -457,10 +464,15 @@ def _glob_paths( ignored_filepaths = set(ignore_patterns) | { ignored_path for ignore_pattern in ignore_patterns - for ignored_path in glob.glob(str(self.config_path / ignore_pattern), recursive=True) + for ignored_path in glob.glob( + str(self.config_path / ignore_pattern), recursive=True + ) } for filepath in path.glob(f"**/*{extension}"): - if any(filepath.match(ignored_filepath) for ignored_filepath in ignored_filepaths): + if any( + filepath.match(ignored_filepath) + for ignored_filepath in ignored_filepaths + ): continue yield filepath @@ -469,7 +481,9 @@ def _track_file(self, path: Path) -> None: """Project file to track for modifications""" self._path_mtimes[path] = path.stat().st_mtime - def _failed_to_load_model_error(self, path: Path, error: t.Union[str, Exception]) -> str: + def _failed_to_load_model_error( + self, path: Path, error: t.Union[str, Exception] + ) -> str: base_message = f"Failed to load model from file '{path}':" if isinstance(error, ValidationError): return validation_error_message(error, base_message) @@ -512,11 +526,15 @@ def _load_scripts(self) -> t.Tuple[MacroRegistry, JinjaMacroRegistry]: self._track_file(path) macro_file_mtime = self._path_mtimes[path] macros_max_mtime = ( - max(macros_max_mtime, macro_file_mtime) if macros_max_mtime else macro_file_mtime + max(macros_max_mtime, macro_file_mtime) + if macros_max_mtime + else macro_file_mtime ) with open(path, "r", encoding="utf-8") as file: jinja_macros.add_macros( - extractor.extract(file.read(), dialect=self.config.model_defaults.dialect) + extractor.extract( + file.read(), dialect=self.config.model_defaults.dialect + ) ) self._macros_max_mtime = macros_max_mtime @@ -540,14 +558,20 @@ def _load_models( """ cache = SqlMeshLoader._Cache(self, self.config_path) - sql_models = self._load_sql_models(macros, jinja_macros, audits, signals, cache, gateway) + sql_models = self._load_sql_models( + macros, jinja_macros, audits, signals, cache, gateway + ) external_models = self._load_external_models(audits, cache, gateway) python_models = self._load_python_models(macros, jinja_macros, audits, signals) all_model_names = list(sql_models) + list(external_models) + list(python_models) - duplicates = [name for name, count in Counter(all_model_names).items() if count > 1] + duplicates = [ + name for name, count in Counter(all_model_names).items() if count > 1 + ] if duplicates: - raise ConfigError(f"Duplicate model name(s) found: {', '.join(duplicates)}.") + raise ConfigError( + f"Duplicate model name(s) found: {', '.join(duplicates)}." + ) return UniqueKeyDict("models", **sql_models, **external_models, **python_models) @@ -614,7 +638,9 @@ def _load_sql_models( ), max_workers=c.MAX_FORK_WORKERS, ) as pool: - futures_to_paths = {pool.submit(load_sql_models, path): path for path in paths} + futures_to_paths = { + pool.submit(load_sql_models, path): path for path in paths + } for future in concurrent.futures.as_completed(futures_to_paths): path = futures_to_paths[future] try: @@ -623,7 +649,8 @@ def _load_sql_models( if model.fqn in models: raise ConfigError( self._failed_to_load_model_error( - path, f"Duplicate SQL model name: '{model.name}'." + path, + f"Duplicate SQL model name: '{model.name}'.", ), path, ) @@ -631,7 +658,9 @@ def _load_sql_models( model._path = path models[model.fqn] = model except Exception as ex: - raise ConfigError(self._failed_to_load_model_error(path, ex), path) + raise ConfigError( + self._failed_to_load_model_error(path, ex), path + ) return models @@ -755,7 +784,9 @@ def _load_audits( if audits_max_mtime else audits_file_mtime ) - expressions = parse(file.read(), default_dialect=self.config.model_defaults.dialect) + expressions = parse( + file.read(), default_dialect=self.config.model_defaults.dialect + ) audits = load_multiple_audits( expressions=expressions, path=path, @@ -800,7 +831,9 @@ def _load_metrics(self) -> UniqueKeyDict[str, MetricMeta]: return metrics - def _load_environment_statements(self, macros: MacroRegistry) -> t.List[EnvironmentStatements]: + def _load_environment_statements( + self, macros: MacroRegistry + ) -> t.List[EnvironmentStatements]: """Loads environment statements.""" if self.config.before_all or self.config.after_all: @@ -824,7 +857,9 @@ def _load_environment_statements(self, macros: MacroRegistry) -> t.List[Environm return [ EnvironmentStatements( - **statements, python_env=python_env, project=self.config.project or None + **statements, + python_env=python_env, + project=self.config.project or None, ) ] return [] @@ -856,7 +891,9 @@ def _load_model_test_file(self, path: Path) -> dict[str, ModelTestMetadata]: # If the user has specified a quoted/escaped gateway (e.g. "gateway: 'ma\tin'"), we need to # parse it as YAML to match the gateway name stored in the config gateway_line = GATEWAY_PATTERN.search(source) - gateway = YAML().load(gateway_line.group(0))["gateway"] if gateway_line else None + gateway = ( + YAML().load(gateway_line.group(0))["gateway"] if gateway_line else None + ) contents = yaml_load(source, variables=get_variables(gateway)) @@ -950,6 +987,7 @@ def _model_cache_entry_id(self, model_path: Path) -> str: # gateway is configurable, and it is retained in a cached # model's python environment if the @gateway macro variable is # used in the model - self._loader.context.gateway or self._loader.config.default_gateway_name, + self._loader.context.gateway + or self._loader.config.default_gateway_name, ] ) diff --git a/sqlmesh/core/macros.py b/sqlmesh/core/macros.py index 2e995003bf..8f4892096c 100644 --- a/sqlmesh/core/macros.py +++ b/sqlmesh/core/macros.py @@ -4,12 +4,12 @@ import sys import types import typing as t +from datetime import date, datetime from enum import Enum from functools import lru_cache, reduce from itertools import chain from pathlib import Path from string import Template -from datetime import datetime, date import sqlglot from sqlglot import Generator, exp, parse_one @@ -21,38 +21,25 @@ from sqlglot.schema import MappingSchema from sqlmesh.core import constants as c -from sqlmesh.core.dialect import ( - SQLMESH_MACRO_PREFIX, - Dialect, - MacroDef, - MacroFunc, - MacroSQL, - MacroStrReplace, - MacroVar, - StagedFilePath, - normalize_model_name, -) -from sqlmesh.utils import ( - DECORATOR_RETURN_TYPE, - UniqueKeyDict, - columns_to_types_all_known, - registry_decorator, -) -from sqlmesh.utils.date import DatetimeRanges, to_datetime, to_date +from sqlmesh.core.dialect import (SQLMESH_MACRO_PREFIX, Dialect, MacroDef, + MacroFunc, MacroSQL, MacroStrReplace, + MacroVar, StagedFilePath, + normalize_model_name) +from sqlmesh.utils import (DECORATOR_RETURN_TYPE, UniqueKeyDict, + columns_to_types_all_known, registry_decorator) +from sqlmesh.utils.date import DatetimeRanges, to_date, to_datetime from sqlmesh.utils.errors import MacroEvalError, SQLMeshError -from sqlmesh.utils.metaprogramming import ( - Executable, - SqlValue, - format_evaluated_code_exception, - prepare_env, -) +from sqlmesh.utils.metaprogramming import (Executable, SqlValue, + format_evaluated_code_exception, + prepare_env) if t.TYPE_CHECKING: from sqlglot.dialects.dialect import DialectType + from sqlmesh.core._typing import TableName from sqlmesh.core.engine_adapter import EngineAdapter - from sqlmesh.core.snapshot import Snapshot from sqlmesh.core.environment import EnvironmentNamingInfo + from sqlmesh.core.snapshot import Snapshot if sys.version_info >= (3, 10): @@ -147,7 +134,9 @@ class Generator(PythonGenerator): exp.Column: lambda self, e: f"exp.to_column('{self.sql(e, 'this')}')", exp.Lambda: lambda self, e: f"lambda {self.expressions(e)}: {self.sql(e, 'this')}", MacroFunc: _macro_func_sql, - MacroSQL: lambda self, e: _macro_sql(self.sql(e, "this"), e.args.get("into")), + MacroSQL: lambda self, e: _macro_sql( + self.sql(e, "this"), e.args.get("into") + ), MacroStrReplace: lambda self, e: _macro_str_replace(self.sql(e, "this")), } @@ -199,7 +188,9 @@ def __init__( "MacroEvaluator": MacroEvaluator, } self.python_env = python_env or {} - self.macros = {normalize_macro_name(k): v.func for k, v in macro.get_registry().items()} + self.macros = { + normalize_macro_name(k): v.func for k, v in macro.get_registry().items() + } self.columns_to_types_called = False self.default_catalog = default_catalog @@ -246,7 +237,11 @@ def send( try: return call_macro( - func, self.dialect, self._path, provided_args=(self, *args), provided_kwargs=kwargs + func, + self.dialect, + self._path, + provided_args=(self, *args), + provided_kwargs=kwargs, ) # type: ignore except Exception as e: raise MacroEvalError( @@ -272,7 +267,9 @@ def evaluate_macros( if var_name not in self.locals and var_name not in variables: if not isinstance(node.parent, StagedFilePath): - raise SQLMeshError(f"Macro variable '{node.name}' is undefined.") + raise SQLMeshError( + f"Macro variable '{node.name}' is undefined." + ) return node @@ -280,10 +277,15 @@ def evaluate_macros( value = self.locals.get(var_name, variables.get(var_name)) if isinstance(value, list): return exp.convert( - tuple(self.transform(v) if isinstance(v, exp.Expr) else v for v in value) + tuple( + self.transform(v) if isinstance(v, exp.Expr) else v + for v in value + ) ) - return exp.convert(self.transform(value) if isinstance(value, exp.Expr) else value) + return exp.convert( + self.transform(value) if isinstance(value, exp.Expr) else value + ) if isinstance(node, exp.Identifier) and "@" in node.this: text = self.template(node.this, {}) if node.this != text: @@ -327,14 +329,18 @@ def template(self, text: t.Any, local_variables: t.Dict[str, t.Any]) -> str: # into strings; in sql we don't convert strings because that would result in adding quotes base_mapping = { k.lower(): convert_sql(v, self.dialect) - for k, v in chain(self.variables.items(), self.locals.items(), local_variables.items()) + for k, v in chain( + self.variables.items(), self.locals.items(), local_variables.items() + ) if k.lower() not in ( "engine_adapter", "snapshot", ) } - return MacroStrTemplate(str(text)).safe_substitute(CaseInsensitiveMapping(base_mapping)) + return MacroStrTemplate(str(text)).safe_substitute( + CaseInsensitiveMapping(base_mapping) + ) def evaluate(self, node: MacroFunc) -> exp.Expr | t.List[exp.Expr] | None: if isinstance(node, MacroDef): @@ -389,18 +395,18 @@ def evaluate(self, node: MacroFunc) -> exp.Expr | t.List[exp.Expr] | None: - and that output is something that _norm_var_arg_lambda() will unpack into varargs > (a list containing a single item of type exp.Tuple/exp.Array) then we will get inconsistent behaviour depending on if this node emits a list with a single item vs multiple items. - + In the first case, emitting a list containing a single array item will cause that array to get unpacked and its *members* passed to the calling macro In the second case, emitting a list containing multiple array items will cause each item to get passed as-is to the calling macro - + To prevent this inconsistency, we wrap this node output in an exp.Array so that _norm_var_arg_lambda() can "unpack" that into the actual argument we want to pass to the parent macro function - + Note we only do this for evaluation results that get passed as an argument to another macro, because when the final result is given to something like SELECT, we still want that to be unpacked into a list of items like: - SELECT ARRAY(1), ARRAY(2) rather than a single item like: - - SELECT ARRAY(ARRAY(1), ARRAY(2)) + - SELECT ARRAY(ARRAY(1), ARRAY(2)) """ result = [exp.Array(expressions=result)] else: @@ -446,13 +452,17 @@ def parse_one( """ return sqlglot.maybe_parse(sql, dialect=self.dialect, into=into, **opts) - def columns_to_types(self, model_name: TableName | exp.Column) -> t.Dict[str, exp.DataType]: + def columns_to_types( + self, model_name: TableName | exp.Column + ) -> t.Dict[str, exp.DataType]: """Returns the columns-to-types mapping corresponding to the specified model.""" # We only return this dummy schema at load time, because if we don't actually know the # target model's schema at creation/evaluation time, returning a dummy schema could lead # to unintelligible errors when the query is executed - if (self._schema is None or self._schema.empty) and self.runtime_stage == "loading": + if ( + self._schema is None or self._schema.empty + ) and self.runtime_stage == "loading": self.columns_to_types_called = True return {"__schema_unavailable_at_load__": exp.DataType.build("unknown")} @@ -464,7 +474,9 @@ def columns_to_types(self, model_name: TableName | exp.Column) -> t.Dict[str, ex model_name = exp.to_table(normalized_model_name) columns_to_types = ( - self._schema.find(model_name, ensure_data_types=True) if self._schema else None + self._schema.find(model_name, ensure_data_types=True) + if self._schema + else None ) if columns_to_types is None: snapshot = self.get_snapshot(model_name) @@ -472,7 +484,9 @@ def columns_to_types(self, model_name: TableName | exp.Column) -> t.Dict[str, ex columns_to_types = snapshot.node.columns_to_types # type: ignore if columns_to_types is None: - raise SQLMeshError(f"Schema for model '{model_name}' can't be statically determined.") + raise SQLMeshError( + f"Schema for model '{model_name}' can't be statically determined." + ) return columns_to_types @@ -545,7 +559,9 @@ def snapshots(self) -> t.Dict[str, Snapshot]: def this_env(self) -> str: """Returns the name of the current environment in before after all.""" if "this_env" not in self.locals: - raise SQLMeshError("Environment name is only available in before_all and after_all") + raise SQLMeshError( + "Environment name is only available in before_all and after_all" + ) return self.locals["this_env"] @property @@ -562,14 +578,18 @@ def views(self) -> t.List[str]: raise SQLMeshError("Views are only available in before_all and after_all") return self.locals["views"] - def var(self, var_name: str, default: t.Optional[t.Any] = None) -> t.Optional[t.Any]: + def var( + self, var_name: str, default: t.Optional[t.Any] = None + ) -> t.Optional[t.Any]: """Returns the value of the specified variable, or the default value if it doesn't exist.""" return { **(self.locals.get(c.SQLMESH_VARS) or {}), **(self.locals.get(c.SQLMESH_VARS_METADATA) or {}), }.get(var_name.lower(), default) - def blueprint_var(self, var_name: str, default: t.Optional[t.Any] = None) -> t.Optional[t.Any]: + def blueprint_var( + self, var_name: str, default: t.Optional[t.Any] = None + ) -> t.Optional[t.Any]: """Returns the value of the specified blueprint variable, or the default value if it doesn't exist.""" return { **(self.locals.get(c.SQLMESH_BLUEPRINT_VARS) or {}), @@ -609,7 +629,9 @@ def add_one(evaluator: MacroEvaluator, column: exp.Literal) -> exp.Add: registry_name = "macros" - def __init__(self, *args: t.Any, metadata_only: bool = False, **kwargs: t.Any) -> None: + def __init__( + self, *args: t.Any, metadata_only: bool = False, **kwargs: t.Any + ) -> None: super().__init__(*args, **kwargs) self.metadata_only = metadata_only @@ -656,7 +678,9 @@ def substitute( return exp.convert(evaluator.locals[name]) if SQLMESH_MACRO_PREFIX in node.name: return node.__class__( - this=evaluator.template(node.name, {k: v.name for k, v in args.items()}) + this=evaluator.template( + node.name, {k: v.name for k, v in args.items()} + ) ) elif isinstance(node, MacroFunc): local_copy = evaluator.locals.copy() @@ -671,9 +695,7 @@ def substitute( expressions = ( item.expressions if isinstance(item, (exp.Array, exp.Tuple)) - else [item.this] - if isinstance(item, exp.Paren) - else item + else [item.this] if isinstance(item, exp.Paren) else item ) else: expressions = items @@ -684,7 +706,8 @@ def substitute( { expression.name.lower(): arg for expression, arg in zip( - func.expressions, args.expressions if isinstance(args, exp.Tuple) else [args] + func.expressions, + args.expressions if isinstance(args, exp.Tuple) else [args], ) }, ) @@ -761,7 +784,9 @@ def reduce_(evaluator: MacroEvaluator, *args: t.Any) -> t.Any: """ *items, func = args items, func = _norm_var_arg_lambda(evaluator, func, *items) # type: ignore - return reduce(lambda a, b: func(exp.Tuple(expressions=[a, b])), ensure_collection(items)) + return reduce( + lambda a, b: func(exp.Tuple(expressions=[a, b])), ensure_collection(items) + ) @macro("FILTER") @@ -789,7 +814,11 @@ def filter_(evaluator: MacroEvaluator, *args: t.Any) -> t.List[t.Any]: """ *items, func = args items, func = _norm_var_arg_lambda(evaluator, func, *items) # type: ignore - return list(filter(lambda arg: evaluator.eval_expression(func(arg)), ensure_collection(items))) + return list( + filter( + lambda arg: evaluator.eval_expression(func(arg)), ensure_collection(items) + ) + ) def _optional_expression( @@ -902,7 +931,9 @@ def star( if suffix and not isinstance(suffix, exp.Literal): raise SQLMeshError(f"Invalid suffix '{suffix}'. Expected a literal.") if not isinstance(quote_identifiers, exp.Boolean): - raise SQLMeshError(f"Invalid quote_identifiers '{quote_identifiers}'. Expected a boolean.") + raise SQLMeshError( + f"Invalid quote_identifiers '{quote_identifiers}'. Expected a boolean." + ) excluded_names = { normalize_identifiers(excluded, dialect=evaluator.dialect).name @@ -914,7 +945,9 @@ def star( ).name columns_to_types = { - k: v for k, v in evaluator.columns_to_types(relation).items() if k not in excluded_names + k: v + for k, v in evaluator.columns_to_types(relation).items() + if k not in excluded_names } if columns_to_types_all_known(columns_to_types): return [ @@ -1007,7 +1040,9 @@ def generate_surrogate_key( return exp.Lower( this=exp.Hex( this=exp.SHA2( - this=exp.Encode(this=func.this, charset=exp.Literal.string("utf-8")), + this=exp.Encode( + this=func.this, charset=exp.Literal.string("utf-8") + ), length=func.args.get("length"), ) ) @@ -1027,7 +1062,9 @@ def generate_surrogate_key( def _is_presto_family(dialect: DialectType) -> bool: """Whether this dialect is Presto, Trino or Athena.""" - return (str(dialect) if dialect else "").split(",")[0].strip().lower() in _PRESTO_FAMILY + return (str(dialect) if dialect else "").split(",")[ + 0 + ].strip().lower() in _PRESTO_FAMILY @lru_cache(maxsize=None) @@ -1128,7 +1165,9 @@ def union( if isinstance(condition, bool): arg_idx += 1 if arg_idx >= len(args): - raise SQLMeshError("Expected more arguments after the condition of the `@UNION` macro.") + raise SQLMeshError( + "Expected more arguments after the condition of the `@UNION` macro." + ) # Check for union type type_ = exp.Literal.string("ALL") @@ -1141,7 +1180,8 @@ def union( # Remaining args should be tables tables = [ - exp.to_table(e.sql(evaluator.dialect), dialect=evaluator.dialect) for e in args[arg_idx:] + exp.to_table(e.sql(evaluator.dialect), dialect=evaluator.dialect) + for e in args[arg_idx:] ] columns = { @@ -1200,10 +1240,14 @@ def haversine_distance( "ASIN", exp.func( "SQRT", - exp.func("POWER", exp.func("SIN", exp.func("RADIANS", (lat2 - lat1) / 2)), 2) + exp.func( + "POWER", exp.func("SIN", exp.func("RADIANS", (lat2 - lat1) / 2)), 2 + ) + exp.func("COS", exp.func("RADIANS", lat1)) * exp.func("COS", exp.func("RADIANS", lat2)) - * exp.func("POWER", exp.func("SIN", exp.func("RADIANS", (lon2 - lon1) / 2)), 2), + * exp.func( + "POWER", exp.func("SIN", exp.func("RADIANS", (lon2 - lon1) / 2)), 2 + ), ), ) * conversion_rate @@ -1260,7 +1304,9 @@ def pivot( @macro("AND") -def and_(evaluator: MacroEvaluator, *expressions: t.Optional[exp.Expr]) -> exp.Condition: +def and_( + evaluator: MacroEvaluator, *expressions: t.Optional[exp.Expr] +) -> exp.Condition: """Returns an AND statement filtering out any NULL expressions.""" conditions = [e for e in expressions if not isinstance(e, exp.Null)] @@ -1287,7 +1333,9 @@ def var( ) -> exp.Expr: """Returns the value of a variable or the default value if the variable is not set.""" if not var_name.is_string: - raise SQLMeshError(f"Invalid variable name '{var_name.sql()}'. Expected a string literal.") + raise SQLMeshError( + f"Invalid variable name '{var_name.sql()}'. Expected a string literal." + ) return exp.convert(evaluator.var(var_name.this, default)) @@ -1340,7 +1388,9 @@ def deduplicate( partition_clause = exp.tuple_(*partition_by) order_expressions = [ - evaluator.transform(parse_one(order_item, into=exp.Ordered, dialect=evaluator.dialect)) + evaluator.transform( + parse_one(order_item, into=exp.Ordered, dialect=evaluator.dialect) + ) for order_item in order_by ] @@ -1422,7 +1472,9 @@ def date_spine( ): date_interval = exp.Interval(this=exp.Literal.number(3), unit=exp.var("month")) else: - date_interval = exp.Interval(this=exp.Literal.number(1), unit=exp.var(datepart_name)) + date_interval = exp.Interval( + this=exp.Literal.number(1), unit=exp.var(datepart_name) + ) generate_date_array = exp.func( "GENERATE_DATE_ARRAY", @@ -1432,7 +1484,9 @@ def date_spine( ) alias_name = f"date_{datepart_name}" - exploded = exp.alias_(exp.func("unnest", generate_date_array), "_exploded", table=[alias_name]) + exploded = exp.alias_( + exp.func("unnest", generate_date_array), "_exploded", table=[alias_name] + ) return exp.select(alias_name).from_(exploded) @@ -1555,9 +1609,13 @@ def call_macro( # https://docs.python.org/3/library/inspect.html#inspect.BoundArguments.arguments param = sig.parameters[arg] if param.kind is inspect.Parameter.VAR_POSITIONAL: - bound.arguments[arg] = tuple(_coerce(v, typ, dialect, path) for v in value) + bound.arguments[arg] = tuple( + _coerce(v, typ, dialect, path) for v in value + ) elif param.kind is inspect.Parameter.VAR_KEYWORD: - bound.arguments[arg] = {k: _coerce(v, typ, dialect, path) for k, v in value.items()} + bound.arguments[arg] = { + k: _coerce(v, typ, dialect, path) for k, v in value.items() + } else: bound.arguments[arg] = _coerce(value, typ, dialect, path) @@ -1599,7 +1657,11 @@ def _coerce( for literal_type_arg in literal_type_args: expr_is_bool = isinstance(expr.this, bool) literal_is_bool = isinstance(literal_type_arg, bool) - if (expr_is_bool and literal_is_bool and literal_type_arg == expr.this) or ( + if ( + expr_is_bool + and literal_is_bool + and literal_type_arg == expr.this + ) or ( not expr_is_bool and not literal_is_bool and str(literal_type_arg) == str(expr.this) @@ -1623,7 +1685,8 @@ def _coerce( ) else: coerced = parse_one( - expr.this if isinstance(expr, exp.Literal) else expr.sql(), into=into + expr.this if isinstance(expr, exp.Literal) else expr.sql(), + into=into, ) if isinstance(coerced, base): return coerced @@ -1644,7 +1707,10 @@ def _coerce( if not generic: return tuple(expr.expressions) if generic[-1] is ...: - return tuple(_coerce(expr, generic[0], dialect, path) for expr in expr.expressions) + return tuple( + _coerce(expr, generic[0], dialect, path) + for expr in expr.expressions + ) if len(generic) == len(expr.expressions): return tuple( _coerce(expr, generic[i], dialect, path) @@ -1655,7 +1721,9 @@ def _coerce( generic = t.get_args(typ) if not generic: return expr.expressions - return [_coerce(expr, generic[0], dialect, path) for expr in expr.expressions] + return [ + _coerce(expr, generic[0], dialect, path) for expr in expr.expressions + ] raise SQLMeshError(base_err_msg) except Exception: if strict: diff --git a/sqlmesh/core/metric/__init__.py b/sqlmesh/core/metric/__init__.py index e3dea9d8ca..f0a51f7e40 100644 --- a/sqlmesh/core/metric/__init__.py +++ b/sqlmesh/core/metric/__init__.py @@ -1,7 +1,5 @@ -from sqlmesh.core.metric.definition import ( - Metric as Metric, - MetricMeta as MetricMeta, - expand_metrics as expand_metrics, - load_metric_ddl as load_metric_ddl, -) +from sqlmesh.core.metric.definition import Metric as Metric +from sqlmesh.core.metric.definition import MetricMeta as MetricMeta +from sqlmesh.core.metric.definition import expand_metrics as expand_metrics +from sqlmesh.core.metric.definition import load_metric_ddl as load_metric_ddl from sqlmesh.core.metric.rewriter import rewrite as rewrite diff --git a/sqlmesh/core/metric/definition.py b/sqlmesh/core/metric/definition.py index 6119a883ed..92e9388a24 100644 --- a/sqlmesh/core/metric/definition.py +++ b/sqlmesh/core/metric/definition.py @@ -10,7 +10,8 @@ from sqlmesh.core.node import str_or_exp_to_str from sqlmesh.utils import UniqueKeyDict from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.pydantic import PydanticModel, ValidationInfo, field_validator, validation_data +from sqlmesh.utils.pydantic import (PydanticModel, ValidationInfo, + field_validator, validation_data) MeasureAndDimTables = t.Tuple[str, t.Tuple[str, ...]] @@ -21,7 +22,8 @@ def load_metric_ddl( """Returns a MetricMeta from raw Metric DDL.""" if not isinstance(expression, d.Metric): _raise_metric_config_error( - f"Only METRIC(...) statements are allowed. Found {expression.sql(pretty=True)}", path + f"Only METRIC(...) statements are allowed. Found {expression.sql(pretty=True)}", + path, ) metric = MetricMeta( @@ -32,7 +34,10 @@ def load_metric_ddl( if expression.comments else None ), - **{prop.name.lower(): prop.args.get("value") for prop in expression.expressions}, + **{ + prop.name.lower(): prop.args.get("value") + for prop in expression.expressions + }, **kwargs, } ) @@ -105,7 +110,8 @@ def to_metric( for node in self.expression.walk(): if isinstance(node, exp.Alias): _raise_metric_config_error( - f"Alias found for metric '{self.name}' which is not allowed", self._path + f"Alias found for metric '{self.name}' which is not allowed", + self._path, ) elif isinstance(node, exp.AggFunc): agg_or_ref = True @@ -170,7 +176,11 @@ def formula(self) -> exp.Expr: """ return exp.alias_( self.expanded.transform( - lambda node: exp.column(node.args["alias"]) if isinstance(node, exp.Alias) else node + lambda node: ( + exp.column(node.args["alias"]) + if isinstance(node, exp.Alias) + else node + ) ), self.name, copy=False, diff --git a/sqlmesh/core/metric/rewriter.py b/sqlmesh/core/metric/rewriter.py index 6c9ec429a8..44918c6edc 100644 --- a/sqlmesh/core/metric/rewriter.py +++ b/sqlmesh/core/metric/rewriter.py @@ -15,7 +15,9 @@ from sqlmesh.core.reference import ReferenceGraph -SourceAggsAndJoins = t.Dict[str, t.Tuple[t.Set[exp.AggFunc], t.Dict[str, t.Optional[exp.Join]]]] +SourceAggsAndJoins = t.Dict[ + str, t.Tuple[t.Set[exp.AggFunc], t.Dict[str, t.Optional[exp.Join]]] +] class Rewriter: @@ -75,7 +77,9 @@ def _expand(self, select: exp.Select) -> None: if name != base_alias } - explicit_joins = {exp.table_name(join.this): join for join in select.args.pop("joins", [])} + explicit_joins = { + exp.table_name(join.this): join for join in select.args.pop("joins", []) + } for i, (name, (aggs, joins)) in enumerate(sources.items()): source: exp.Expr = exp.to_table(name) @@ -98,7 +102,9 @@ def _expand(self, select: exp.Select) -> None: where = select.args.pop("where", None) if where: - query.where(_replace_table(where.this, table_name, base_alias), copy=False) + query.where( + _replace_table(where.this, table_name, base_alias), copy=False + ) select.from_(query.subquery(base_alias, copy=False), copy=False) else: @@ -143,7 +149,9 @@ def _add_joins( (model for model in models if remove_namespace(model) == t), models[0], ) - node.args["table"] = exp.to_identifier(t or remove_namespace(model)) + node.args["table"] = exp.to_identifier( + t or remove_namespace(model) + ) if model not in joins: joins[model] = None diff --git a/sqlmesh/core/model/__init__.py b/sqlmesh/core/model/__init__.py index c2ab47d9e7..16ffc7e068 100644 --- a/sqlmesh/core/model/__init__.py +++ b/sqlmesh/core/model/__init__.py @@ -1,42 +1,48 @@ -from sqlmesh.core.model.cache import ( - ModelCache as ModelCache, - OptimizedQueryCache as OptimizedQueryCache, -) +from sqlmesh.core.model.cache import ModelCache as ModelCache +from sqlmesh.core.model.cache import OptimizedQueryCache as OptimizedQueryCache from sqlmesh.core.model.decorator import model as model -from sqlmesh.core.model.definition import ( - AuditResult as AuditResult, - ExternalModel as ExternalModel, - Model as Model, - PythonModel as PythonModel, - SeedModel as SeedModel, - SqlModel as SqlModel, - create_external_model as create_external_model, - create_python_model as create_python_model, - create_seed_model as create_seed_model, - create_sql_model as create_sql_model, - load_sql_based_model as load_sql_based_model, - load_sql_based_models as load_sql_based_models, -) -from sqlmesh.core.model.kind import ( - CustomKind as CustomKind, - EmbeddedKind as EmbeddedKind, - ExternalKind as ExternalKind, - FullKind as FullKind, - IncrementalByTimeRangeKind as IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind as IncrementalByUniqueKeyKind, - IncrementalUnmanagedKind as IncrementalUnmanagedKind, - IncrementalByPartitionKind as IncrementalByPartitionKind, - ModelKind as ModelKind, - ModelKindMixin as ModelKindMixin, - ModelKindName as ModelKindName, - SCDType2ByColumnKind as SCDType2ByColumnKind, - SCDType2ByTimeKind as SCDType2ByTimeKind, - SeedKind as SeedKind, - TimeColumn as TimeColumn, - ViewKind as ViewKind, - ManagedKind as ManagedKind, - model_kind_validator as model_kind_validator, -) +from sqlmesh.core.model.definition import AuditResult as AuditResult +from sqlmesh.core.model.definition import ExternalModel as ExternalModel +from sqlmesh.core.model.definition import Model as Model +from sqlmesh.core.model.definition import PythonModel as PythonModel +from sqlmesh.core.model.definition import SeedModel as SeedModel +from sqlmesh.core.model.definition import SqlModel as SqlModel +from sqlmesh.core.model.definition import \ + create_external_model as create_external_model +from sqlmesh.core.model.definition import \ + create_python_model as create_python_model +from sqlmesh.core.model.definition import \ + create_seed_model as create_seed_model +from sqlmesh.core.model.definition import create_sql_model as create_sql_model +from sqlmesh.core.model.definition import \ + load_sql_based_model as load_sql_based_model +from sqlmesh.core.model.definition import \ + load_sql_based_models as load_sql_based_models +from sqlmesh.core.model.kind import CustomKind as CustomKind +from sqlmesh.core.model.kind import EmbeddedKind as EmbeddedKind +from sqlmesh.core.model.kind import ExternalKind as ExternalKind +from sqlmesh.core.model.kind import FullKind as FullKind +from sqlmesh.core.model.kind import \ + IncrementalByPartitionKind as IncrementalByPartitionKind +from sqlmesh.core.model.kind import \ + IncrementalByTimeRangeKind as IncrementalByTimeRangeKind +from sqlmesh.core.model.kind import \ + IncrementalByUniqueKeyKind as IncrementalByUniqueKeyKind +from sqlmesh.core.model.kind import \ + IncrementalUnmanagedKind as IncrementalUnmanagedKind +from sqlmesh.core.model.kind import ManagedKind as ManagedKind +from sqlmesh.core.model.kind import ModelKind as ModelKind +from sqlmesh.core.model.kind import ModelKindMixin as ModelKindMixin +from sqlmesh.core.model.kind import ModelKindName as ModelKindName +from sqlmesh.core.model.kind import \ + SCDType2ByColumnKind as SCDType2ByColumnKind +from sqlmesh.core.model.kind import SCDType2ByTimeKind as SCDType2ByTimeKind +from sqlmesh.core.model.kind import SeedKind as SeedKind +from sqlmesh.core.model.kind import TimeColumn as TimeColumn +from sqlmesh.core.model.kind import ViewKind as ViewKind +from sqlmesh.core.model.kind import \ + model_kind_validator as model_kind_validator from sqlmesh.core.model.meta import ModelMeta as ModelMeta -from sqlmesh.core.model.schema import update_model_schemas as update_model_schemas +from sqlmesh.core.model.schema import \ + update_model_schemas as update_model_schemas from sqlmesh.core.model.seed import Seed as Seed diff --git a/sqlmesh/core/model/cache.py b/sqlmesh/core/model/cache.py index 1f038c5d79..ddb5c27697 100644 --- a/sqlmesh/core/model/cache.py +++ b/sqlmesh/core/model/cache.py @@ -2,6 +2,7 @@ import logging import typing as t +from dataclasses import dataclass from pathlib import Path from sqlglot import exp @@ -10,18 +11,17 @@ from sqlglot.schema import MappingSchema from sqlmesh.core import constants as c -from sqlmesh.core.model.definition import ExternalModel, Model, SqlModel, _Model +from sqlmesh.core.model.definition import (ExternalModel, Model, SqlModel, + _Model) from sqlmesh.utils.cache import FileCache from sqlmesh.utils.hashing import crc32 from sqlmesh.utils.process import PoolExecutor, create_process_pool_executor -from dataclasses import dataclass - logger = logging.getLogger(__name__) if t.TYPE_CHECKING: - from sqlmesh.core.snapshot import SnapshotId from sqlmesh.core.linter.rule import Rule + from sqlmesh.core.snapshot import SnapshotId T = t.TypeVar("T") @@ -52,11 +52,15 @@ def get_or_load( The model definition. """ cache_entry = self._file_cache.get(name, entry_id) - if isinstance(cache_entry, list) and isinstance(seq_get(cache_entry, 0), _Model): + if isinstance(cache_entry, list) and isinstance( + seq_get(cache_entry, 0), _Model + ): return cache_entry models = loader() - if isinstance(models, list) and isinstance(seq_get(models, 0), (SqlModel, ExternalModel)): + if isinstance(models, list) and isinstance( + seq_get(models, 0), (SqlModel, ExternalModel) + ): # make sure we preload full_depends_on for model in models: model.full_depends_on @@ -159,7 +163,9 @@ def _entry_name(model: SqlModel) -> str: return f"{model.name}_{crc32(hash_data)}" -def optimized_query_cache_pool(optimized_query_cache: OptimizedQueryCache) -> PoolExecutor: +def optimized_query_cache_pool( + optimized_query_cache: OptimizedQueryCache, +) -> PoolExecutor: return create_process_pool_executor( initializer=_init_optimized_query_cache, initargs=(optimized_query_cache,), @@ -190,7 +196,9 @@ def load_optimized_query( # this can happen if there is a query rendering error. # for example, the model query references some python library or function that was available # at the time the model was created but has since been removed locally - logger.exception(f"Failed to cache optimized query for model '{model.name}'") + logger.exception( + f"Failed to cache optimized query for model '{model.name}'" + ) return snapshot_id, entry_name @@ -221,7 +229,9 @@ def load_optimized_query_and_mapping( def _mapping_schema_hash_data(schema: t.Dict[str, t.Any]) -> t.List[str]: - keys = sorted(schema) if all(isinstance(v, dict) for v in schema.values()) else schema + keys = ( + sorted(schema) if all(isinstance(v, dict) for v in schema.values()) else schema + ) data = [] for k in keys: diff --git a/sqlmesh/core/model/common.py b/sqlmesh/core/model/common.py index f03cf49753..c00e6b4213 100644 --- a/sqlmesh/core/model/common.py +++ b/sqlmesh/core/model/common.py @@ -2,9 +2,9 @@ import ast import typing as t +from difflib import get_close_matches from pathlib import Path -from difflib import get_close_matches from sqlglot import exp from sqlglot.helper import ensure_list @@ -13,23 +13,15 @@ from sqlmesh.core.macros import MacroRegistry, MacroStrTemplate from sqlmesh.utils import str_to_bool from sqlmesh.utils.errors import ConfigError, SQLMeshError, raise_config_error -from sqlmesh.utils.metaprogramming import ( - Executable, - SqlValue, - build_env, - prepare_env, - serialize_env, -) -from sqlmesh.utils.pydantic import ( - PydanticModel, - ValidationInfo, - field_validator, - get_dialect, - validation_data, -) +from sqlmesh.utils.metaprogramming import (Executable, SqlValue, build_env, + prepare_env, serialize_env) +from sqlmesh.utils.pydantic import (PydanticModel, ValidationInfo, + field_validator, get_dialect, + validation_data) if t.TYPE_CHECKING: from sqlglot.dialects.dialect import DialectType + from sqlmesh.utils import registry_decorator from sqlmesh.utils.jinja import MacroReference @@ -83,7 +75,9 @@ def _is_metadata_var( # We've concluded this variable is definitely not metadata-only return False - appears_under_metadata_macro_func = expr_under_metadata_macro_func.get(id(expression)) + appears_under_metadata_macro_func = expr_under_metadata_macro_func.get( + id(expression) + ) if is_metadata_so_far and ( appears_in_metadata_expression or appears_under_metadata_macro_func ): @@ -115,13 +109,18 @@ def _is_metadata_macro(name: str, appears_in_metadata_expression: bool) -> bool: if isinstance(expression, d.Jinja): continue - for macro_func_or_var in expression.find_all(d.MacroFunc, d.MacroVar, exp.Identifier): + for macro_func_or_var in expression.find_all( + d.MacroFunc, d.MacroVar, exp.Identifier + ): if macro_func_or_var.__class__ is d.MacroFunc: name = macro_func_or_var.this.name.lower() if name not in macros: continue - used_macros[name] = (macros[name], _is_metadata_macro(name, is_metadata)) + used_macros[name] = ( + macros[name], + _is_metadata_macro(name, is_metadata), + ) if name in (c.VAR, c.BLUEPRINT_VAR): args = macro_func_or_var.this.expressions @@ -150,11 +149,17 @@ def _is_metadata_macro(name: str, appears_in_metadata_expression: bool) -> bool: # metadata expression, then we can avoid traversing nested macro function calls. var_refs, _expr_under_metadata_macro_func, _visited_macro_funcs = ( - _extract_macro_func_variable_references(macro_func_or_var, is_metadata) + _extract_macro_func_variable_references( + macro_func_or_var, is_metadata + ) + ) + expr_under_metadata_macro_func.update( + _expr_under_metadata_macro_func ) - expr_under_metadata_macro_func.update(_expr_under_metadata_macro_func) visited_macro_funcs.update(_visited_macro_funcs) - outermost_macro_func_ancestor_by_var |= {var_ref: name for var_ref in var_refs} + outermost_macro_func_ancestor_by_var |= { + var_ref: name for var_ref in var_refs + } elif macro_func_or_var.__class__ is d.MacroVar: var_name = macro_func_or_var.name.lower() if var_name in macros: @@ -167,11 +172,16 @@ def _is_metadata_macro(name: str, appears_in_metadata_expression: bool) -> bool: var_name, macro_func_or_var, is_metadata ) elif ( - isinstance(macro_func_or_var, (exp.Identifier, d.MacroStrReplace, d.MacroSQL)) + isinstance( + macro_func_or_var, (exp.Identifier, d.MacroStrReplace, d.MacroSQL) + ) ) and "@" in macro_func_or_var.name: - for _, identifier, braced_identifier, _ in MacroStrTemplate.pattern.findall( - macro_func_or_var.name - ): + for ( + _, + identifier, + braced_identifier, + _, + ) in MacroStrTemplate.pattern.findall(macro_func_or_var.name): var_name = braced_identifier or identifier if var_name in variables or var_name in blueprint_variables: used_variables[var_name] = _is_metadata_var( @@ -221,16 +231,25 @@ def _extract_macro_func_variable_references( this = n.this args = this.expressions - if this.name.lower() in (c.VAR, c.BLUEPRINT_VAR) and args and args[0].is_string: + if ( + this.name.lower() in (c.VAR, c.BLUEPRINT_VAR) + and args + and args[0].is_string + ): var_references.add(args[0].this.lower()) expr_under_metadata_macro_func[id(n)] = is_metadata elif isinstance(n, d.MacroVar): var_references.add(n.name.lower()) expr_under_metadata_macro_func[id(n)] = is_metadata - elif isinstance(n, (exp.Identifier, d.MacroStrReplace, d.MacroSQL)) and "@" in n.name: + elif ( + isinstance(n, (exp.Identifier, d.MacroStrReplace, d.MacroSQL)) + and "@" in n.name + ): var_references.update( (braced_identifier or identifier).lower() - for _, identifier, braced_identifier, _ in MacroStrTemplate.pattern.findall(n.name) + for _, identifier, braced_identifier, _ in MacroStrTemplate.pattern.findall( + n.name + ) ) expr_under_metadata_macro_func[id(n)] = is_metadata @@ -265,14 +284,19 @@ def _add_variables_to_python_env( metadata_used_variables = { var_name for var_name, is_metadata in used_variables.items() if is_metadata } - for used_var, outermost_macro_func in (outermost_macro_func_ancestor_by_var or {}).items(): + for used_var, outermost_macro_func in ( + outermost_macro_func_ancestor_by_var or {} + ).items(): used_var_is_metadata = used_variables.get(used_var) if used_var_is_metadata is False: continue # At this point we can decide whether a variable reference in a macro call's AST is # metadata-only, because we've annotated the corresponding macro call in the python env. - if outermost_macro_func in python_env and python_env[outermost_macro_func].is_metadata: + if ( + outermost_macro_func in python_env + and python_env[outermost_macro_func].is_metadata + ): metadata_used_variables.add(used_var) non_metadata_used_variables = set(used_variables) - metadata_used_variables @@ -286,7 +310,9 @@ def _add_variables_to_python_env( metadata_variables = { k: v for k, v in (variables or {}).items() if k in metadata_used_variables } - variables = {k: v for k, v in (variables or {}).items() if k in non_metadata_used_variables} + variables = { + k: v for k, v in (variables or {}).items() if k in non_metadata_used_variables + } if variables: python_env[c.SQLMESH_VARS] = Executable.value(variables, sort_root_dict=True) @@ -302,7 +328,9 @@ def _add_variables_to_python_env( if k in metadata_used_variables } blueprint_variables = { - k.lower(): SqlValue(sql=v.sql(dialect=dialect)) if isinstance(v, exp.Expr) else v + k.lower(): ( + SqlValue(sql=v.sql(dialect=dialect)) if isinstance(v, exp.Expr) else v + ) for k, v in blueprint_variables.items() if k in non_metadata_used_variables } @@ -350,7 +378,9 @@ def var(var_name: str, default: t.Optional[t.Any] = None) -> t.Optional[t.Any]: return (variables or {}).get(var_name.lower(), default) @staticmethod - def blueprint_var(var_name: str, default: t.Optional[t.Any] = None) -> t.Optional[t.Any]: + def blueprint_var( + var_name: str, default: t.Optional[t.Any] = None + ) -> t.Optional[t.Any]: return (blueprint_variables or {}).get(var_name.lower(), default) env = prepare_env(python_env) @@ -369,7 +399,9 @@ def blueprint_var(var_name: str, default: t.Optional[t.Any] = None) -> t.Optiona if isinstance(node, ast.Call): func = node.func - if not isinstance(func, ast.Attribute) or not isinstance(func.value, ast.Name): + if not isinstance(func, ast.Attribute) or not isinstance( + func.value, ast.Name + ): continue def get_first_arg(keyword_arg_name: str) -> t.Any: @@ -395,7 +427,10 @@ def get_first_arg(keyword_arg_name: str) -> t.Any: f"Argument '{expression.strip()}' must be resolvable at parse time." ) - if func.value.id == "context" and func.attr in ("table", "resolve_table"): + if func.value.id == "context" and func.attr in ( + "table", + "resolve_table", + ): depends_on.add(get_first_arg("model_name")) elif func.value.id in ("context", "evaluator") and func.attr in ( c.VAR, @@ -420,7 +455,9 @@ def get_first_arg(keyword_arg_name: str) -> t.Any: ) for var_name in next_variables: - used_variables[var_name] = used_variables.get(var_name, True) and bool(is_metadata) + used_variables[var_name] = used_variables.get(var_name, True) and bool( + is_metadata + ) return depends_on, used_variables @@ -451,10 +488,13 @@ def validate_extra_and_required_fields( close_matches[field] = matches[0] if len(close_matches) == 1: - similar_msg = ". Did you mean " + "'" + "', '".join(close_matches.values()) + "'?" + similar_msg = ( + ". Did you mean " + "'" + "', '".join(close_matches.values()) + "'?" + ) else: similar = [ - f"- {field}: Did you mean '{match}'?" for field, match in close_matches.items() + f"- {field}: Did you mean '{match}'?" + for field, match in close_matches.items() ] similar_msg = "\n\n " + "\n ".join(similar) if similar else "" @@ -541,7 +581,9 @@ def parse_properties( ) properties = ( - exp.Tuple(expressions=eq_expressions) if isinstance(v, (exp.Paren, exp.Array)) else v + exp.Tuple(expressions=eq_expressions) + if isinstance(v, (exp.Paren, exp.Array)) + else v ) elif isinstance(v, dict): properties = exp.Tuple( @@ -579,17 +621,23 @@ def depends_on(cls: t.Type, v: t.Any, info: ValidationInfo) -> t.Optional[t.Set[ for table in v.expressions } if isinstance(v, (exp.Table, exp.Column)): - return {d.normalize_model_name(v, default_catalog=default_catalog, dialect=dialect)} + return { + d.normalize_model_name(v, default_catalog=default_catalog, dialect=dialect) + } if hasattr(v, "__iter__") and not isinstance(v, str): return { - d.normalize_model_name(name, default_catalog=default_catalog, dialect=dialect) + d.normalize_model_name( + name, default_catalog=default_catalog, dialect=dialect + ) for name in v } return v -def sort_python_env(python_env: t.Dict[str, Executable]) -> t.List[t.Tuple[str, Executable]]: +def sort_python_env( + python_env: t.Dict[str, Executable], +) -> t.List[t.Tuple[str, Executable]]: """Returns the python env sorted.""" return sorted(python_env.items(), key=lambda x: (x[1].kind, x[0])) @@ -710,11 +758,17 @@ def _validate_parsable_sql( if isinstance(v, list): dialect = get_dialect(info.data) return [ - ParsableSql(sql=s) - if isinstance(s, str) - else ParsableSql.from_parsed_expression(s, dialect, use_meta_sql=False) - if isinstance(s, exp.Expr) - else ParsableSql.parse_obj(s) + ( + ParsableSql(sql=s) + if isinstance(s, str) + else ( + ParsableSql.from_parsed_expression( + s, dialect, use_meta_sql=False + ) + if isinstance(s, exp.Expr) + else ParsableSql.parse_obj(s) + ) + ) for s in v ] return ParsableSql.parse_obj(v) diff --git a/sqlmesh/core/model/decorator.py b/sqlmesh/core/model/decorator.py index 304c07276c..bbe2d8a0f5 100644 --- a/sqlmesh/core/model/decorator.py +++ b/sqlmesh/core/model/decorator.py @@ -1,35 +1,31 @@ from __future__ import annotations -import typing as t -from pathlib import Path import inspect import re +import typing as t +from pathlib import Path from sqlglot import exp from sqlglot.dialects.dialect import DialectType -from sqlmesh.core.config.common import VirtualEnvironmentMode -from sqlmesh.core.macros import MacroRegistry -from sqlmesh.core.signal import SignalRegistry -from sqlmesh.utils.jinja import JinjaMacroRegistry from sqlmesh.core import constants as c +from sqlmesh.core.config.common import VirtualEnvironmentMode from sqlmesh.core.dialect import MacroFunc, parse_one -from sqlmesh.core.model.definition import ( - Model, - create_python_model, - create_sql_model, - create_models_from_blueprints, - get_model_name, - parse_defaults_properties, - render_meta_fields, - render_model_defaults, -) +from sqlmesh.core.macros import MacroRegistry +from sqlmesh.core.model.definition import (Model, + create_models_from_blueprints, + create_python_model, + create_sql_model, get_model_name, + parse_defaults_properties, + render_meta_fields, + render_model_defaults) from sqlmesh.core.model.kind import ModelKindName, _ModelKind -from sqlmesh.utils import registry_decorator, DECORATOR_RETURN_TYPE +from sqlmesh.core.signal import SignalRegistry +from sqlmesh.utils import DECORATOR_RETURN_TYPE, registry_decorator from sqlmesh.utils.errors import ConfigError, raise_config_error +from sqlmesh.utils.jinja import JinjaMacroRegistry from sqlmesh.utils.metaprogramming import build_env, serialize_env - if t.TYPE_CHECKING: from sqlmesh.core.audit import ModelAudit @@ -40,7 +36,9 @@ class model(registry_decorator): registry_name = "python_models" _dialect: DialectType = None - def __init__(self, name: t.Optional[str] = None, is_sql: bool = False, **kwargs: t.Any) -> None: + def __init__( + self, name: t.Optional[str] = None, is_sql: bool = False, **kwargs: t.Any + ) -> None: if not is_sql and "columns" not in kwargs: raise ConfigError("Python model must define column schema.") @@ -60,7 +58,9 @@ def __init__(self, name: t.Optional[str] = None, is_sql: bool = False, **kwargs: call[0], { arg_key: exp.convert( - tuple(arg_value) if isinstance(arg_value, list) else arg_value + tuple(arg_value) + if isinstance(arg_value, list) + else arg_value ) for arg_key, arg_value in call[1].items() }, @@ -177,7 +177,9 @@ def model( f"""Python model "{self.name}"'s `kind` argument was passed a SQLMesh `{type(kind).__name__}` object. This may result in unexpected behavior - provide a dictionary instead.""" ) elif isinstance(kind, dict): - if "name" not in kind or not isinstance(kind.get("name"), ModelKindName): + if "name" not in kind or not isinstance( + kind.get("name"), ModelKindName + ): raise ConfigError( f"""Python model "{self.name}"'s `kind` dictionary must contain a `name` key with a valid ModelKindName enum value.""" ) @@ -217,7 +219,9 @@ def model( else {} ) - rendered_defaults = parse_defaults_properties(rendered_defaults, dialect=dialect) + rendered_defaults = parse_defaults_properties( + rendered_defaults, dialect=dialect + ) common_kwargs = { "defaults": rendered_defaults, @@ -244,7 +248,11 @@ def model( statements = common_kwargs.get(key) if statements: common_kwargs[key] = [ - parse_one(s, dialect=common_kwargs.get("dialect")) if isinstance(s, str) else s + ( + parse_one(s, dialect=common_kwargs.get("dialect")) + if isinstance(s, str) + else s + ) for s in statements ] diff --git a/sqlmesh/core/model/definition.py b/sqlmesh/core/model/definition.py index e8e122dece..08effbb887 100644 --- a/sqlmesh/core/model/definition.py +++ b/sqlmesh/core/model/definition.py @@ -2,8 +2,8 @@ import json import logging -import types import re +import types import typing as t from functools import cached_property, partial from pathlib import Path @@ -11,64 +11,55 @@ from pydantic import Field from sqlglot import exp from sqlglot.helper import seq_get +from sqlglot.optimizer.normalize_identifiers import normalize_identifiers from sqlglot.optimizer.qualify_columns import quote_identifiers from sqlglot.optimizer.simplify import gen -from sqlglot.optimizer.normalize_identifiers import normalize_identifiers from sqlglot.schema import MappingSchema, nested_set from sqlglot.time import format_time from sqlmesh.core import constants as c from sqlmesh.core import dialect as d from sqlmesh.core.audit import Audit, ModelAudit -from sqlmesh.core.node import IntervalUnit from sqlmesh.core.macros import MacroRegistry, macro -from sqlmesh.core.model.common import ( - ParsableSql, - make_python_env, - parse_dependencies, - parse_strings_with_macro_refs, - single_value_or_tuple, - sorted_python_env_payloads, - validate_extra_and_required_fields, -) +from sqlmesh.core.model.common import (ParsableSql, make_python_env, + parse_dependencies, + parse_strings_with_macro_refs, + single_value_or_tuple, + sorted_python_env_payloads, + validate_extra_and_required_fields) +from sqlmesh.core.model.kind import (CustomKind, ExternalKind, FullKind, + ModelKind, ModelKindName, SeedKind, + create_model_kind) from sqlmesh.core.model.meta import ModelMeta -from sqlmesh.core.model.kind import ( - ExternalKind, - ModelKindName, - SeedKind, - ModelKind, - FullKind, - create_model_kind, - CustomKind, -) from sqlmesh.core.model.seed import CsvSeedReader, Seed, create_seed +from sqlmesh.core.node import IntervalUnit from sqlmesh.core.renderer import ExpressionRenderer, QueryRenderer from sqlmesh.core.signal import SignalRegistry -from sqlmesh.utils import columns_to_types_all_known, str_to_bool, UniqueKeyDict +from sqlmesh.utils import (UniqueKeyDict, columns_to_types_all_known, + str_to_bool) from sqlmesh.utils.cron import CroniterCache -from sqlmesh.utils.date import TimeLike, make_inclusive, to_datetime, to_time_column -from sqlmesh.utils.errors import ConfigError, SQLMeshError, raise_config_error, PythonModelEvalError +from sqlmesh.utils.date import (TimeLike, make_inclusive, to_datetime, + to_time_column) +from sqlmesh.utils.errors import (ConfigError, PythonModelEvalError, + SQLMeshError, raise_config_error) from sqlmesh.utils.hashing import hash_data -from sqlmesh.utils.jinja import JinjaMacroRegistry, extract_macro_references_and_variables -from sqlmesh.utils.pydantic import PydanticModel, PRIVATE_FIELDS -from sqlmesh.utils.metaprogramming import ( - Executable, - SqlValue, - build_env, - prepare_env, - serialize_env, - format_evaluated_code_exception, -) +from sqlmesh.utils.jinja import (JinjaMacroRegistry, + extract_macro_references_and_variables) +from sqlmesh.utils.metaprogramming import (Executable, SqlValue, build_env, + format_evaluated_code_exception, + prepare_env, serialize_env) +from sqlmesh.utils.pydantic import PRIVATE_FIELDS, PydanticModel if t.TYPE_CHECKING: from sqlglot.dialects.dialect import DialectType - from sqlmesh.core.node import _Node - from sqlmesh.core._typing import Self, TableName, SessionProperties + + from sqlmesh.core._typing import Self, SessionProperties, TableName from sqlmesh.core.context import ExecutionContext from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.engine_adapter._typing import QueryOrDF from sqlmesh.core.engine_adapter.shared import DataObjectType from sqlmesh.core.linter.rule import Rule + from sqlmesh.core.node import _Node from sqlmesh.core.snapshot import DeployabilityIndex, Node, Snapshot from sqlmesh.utils.jinja import MacroReference @@ -152,8 +143,12 @@ class _Model(ModelMeta, frozen=True): audit_definitions: t.Dict[str, ModelAudit] = {} mapping_schema: t.Dict[str, t.Any] = {} extract_dependencies_from_query: bool = True - pre_statements_: t.Optional[t.List[ParsableSql]] = Field(default=None, alias="pre_statements") - post_statements_: t.Optional[t.List[ParsableSql]] = Field(default=None, alias="post_statements") + pre_statements_: t.Optional[t.List[ParsableSql]] = Field( + default=None, alias="pre_statements" + ) + post_statements_: t.Optional[t.List[ParsableSql]] = Field( + default=None, alias="post_statements" + ) on_virtual_update_: t.Optional[t.List[ParsableSql]] = Field( default=None, alias="on_virtual_update" ) @@ -246,9 +241,9 @@ def render_definition( expressions.append( exp.Property( this=field_info.alias or field_name, - value=META_FIELD_CONVERTER.get(field_name, exp.to_identifier)( - field_value - ), + value=META_FIELD_CONVERTER.get( + field_name, exp.to_identifier + )(field_value), ) ) @@ -258,7 +253,9 @@ def render_definition( jinja_expressions = [] python_expressions = [] if include_python: - python_env = d.PythonCode(expressions=sorted_python_env_payloads(self.python_env)) + python_env = d.PythonCode( + expressions=sorted_python_env_payloads(self.python_env) + ) if python_env.expressions: python_expressions.append(python_env) @@ -302,7 +299,9 @@ def render_query( """ return exp.select( *( - exp.cast(exp.Null(), column_type, copy=False).as_(name, copy=False, quoted=True) + exp.cast(exp.Null(), column_type, copy=False).as_( + name, copy=False, quoted=True + ) for name, column_type in (self.columns_to_types or {}).items() ), copy=False, @@ -536,9 +535,11 @@ def render_audit_query( deployability_index=deployability_index, **{ **audit.defaults, - "this_model": exp.select("*").from_(quoted_model_name).where(where).subquery() - if where is not None - else quoted_model_name, + "this_model": ( + exp.select("*").from_(quoted_model_name).where(where).subquery() + if where is not None + else quoted_model_name + ), **kwargs, }, # type: ignore ) @@ -632,11 +633,15 @@ def render_signals( def _render(e: exp.Expr) -> str | int | float | bool: rendered_exprs = ( - self._create_renderer(e).render(start=start, end=end, execution_time=execution_time) + self._create_renderer(e).render( + start=start, end=end, execution_time=execution_time + ) or [] ) if len(rendered_exprs) != 1: - raise SQLMeshError(f"Expected one expression but got {len(rendered_exprs)}") + raise SQLMeshError( + f"Expected one expression but got {len(rendered_exprs)}" + ) rendered = rendered_exprs[0] if rendered.is_int: @@ -649,7 +654,9 @@ def _render(e: exp.Expr) -> str | int | float | bool: # airflow only return [ - {k: _render(v) for k, v in signal.items()} for name, signal in self.signals if not name + {k: _render(v) for k, v in signal.items()} + for name, signal in self.signals + if not name ] def render_signal_calls(self) -> EvaluatableSignals: @@ -657,7 +664,8 @@ def render_signal_calls(self) -> EvaluatableSignals: env = prepare_env(python_env) signals_to_kwargs = { name: { - k: seq_get(self._create_renderer(v).render() or [], 0) for k, v in kwargs.items() + k: seq_get(self._create_renderer(v).render() or [], 0) + for k, v in kwargs.items() } for name, kwargs in self.signals if name @@ -686,18 +694,27 @@ def render_merge_filter( ) if len(rendered_exprs) != 1: raise SQLMeshError(f"Expected one expression but got {len(rendered_exprs)}") - return rendered_exprs[0].transform(d.replace_merge_table_aliases, dialect=self.dialect) + return rendered_exprs[0].transform( + d.replace_merge_table_aliases, dialect=self.dialect + ) def _render_properties( - self, properties: t.Dict[str, exp.Expr] | SessionProperties, **render_kwargs: t.Any + self, + properties: t.Dict[str, exp.Expr] | SessionProperties, + **render_kwargs: t.Any, ) -> t.Dict[str, t.Any]: def _render(expression: exp.Expr) -> exp.Expr | None: # note: we use the _statement_renderer instead of _create_renderer because it sets model_fqn which # in turn makes @this_model available in the evaluation context - rendered_exprs = self._statement_renderer(expression).render(**render_kwargs) + rendered_exprs = self._statement_renderer(expression).render( + **render_kwargs + ) # Inform instead of raising for cases where a property is conditionally assigned - if not rendered_exprs or rendered_exprs[0].sql().lower() in {"none", "null"}: + if not rendered_exprs or rendered_exprs[0].sql().lower() in { + "none", + "null", + }: logger.info( f"Rendering '{expression.sql(dialect=self.dialect)}' did not return an expression" ) @@ -717,7 +734,9 @@ def _render(expression: exp.Expr) -> exp.Expr | None: } def render_physical_properties(self, **render_kwargs: t.Any) -> t.Dict[str, t.Any]: - rendered = self._render_properties(properties=self.physical_properties, **render_kwargs) + rendered = self._render_properties( + properties=self.physical_properties, **render_kwargs + ) # Some engines (e.g. StarRocks) accept properties whose values reference other models and # need the physical table name rather than the logical view SQLMesh exposes. Resolve those. @@ -743,10 +762,14 @@ def render_physical_properties(self, **render_kwargs: t.Any) -> t.Dict[str, t.An return rendered def render_virtual_properties(self, **render_kwargs: t.Any) -> t.Dict[str, t.Any]: - return self._render_properties(properties=self.virtual_properties, **render_kwargs) + return self._render_properties( + properties=self.virtual_properties, **render_kwargs + ) def render_session_properties(self, **render_kwargs: t.Any) -> t.Dict[str, t.Any]: - return self._render_properties(properties=self.session_properties, **render_kwargs) + return self._render_properties( + properties=self.session_properties, **render_kwargs + ) def _create_renderer(self, expression: exp.Expr) -> ExpressionRenderer: return ExpressionRenderer( @@ -775,7 +798,9 @@ def ctas_query(self, **render_kwarg: t.Any) -> exp.Query: query = self.render_query_or_raise(**render_kwarg).limit(0) for select_or_set_op in query.find_all(exp.Select, exp.SetOperation): - if isinstance(select_or_set_op, exp.Select) and select_or_set_op.args.get("from_"): + if isinstance(select_or_set_op, exp.Select) and select_or_set_op.args.get( + "from_" + ): select_or_set_op.where(exp.false(), copy=False) if self.managed_columns: @@ -822,7 +847,9 @@ def text_diff(self, other: Node, rendered: bool = False) -> str: return text_diff - def set_time_format(self, default_time_format: str = c.DEFAULT_TIME_COLUMN_FORMAT) -> None: + def set_time_format( + self, default_time_format: str = c.DEFAULT_TIME_COLUMN_FORMAT + ) -> None: """Sets the default time format for a model. Args: @@ -843,7 +870,9 @@ def set_time_format(self, default_time_format: str = c.DEFAULT_TIME_COLUMN_FORMA self.time_column.format = default_time_format def convert_to_time_column( - self, time: TimeLike, columns_to_types: t.Optional[t.Dict[str, exp.DataType]] = None + self, + time: TimeLike, + columns_to_types: t.Optional[t.Dict[str, exp.DataType]] = None, ) -> exp.Expr: """Convert a TimeLike object to the same time format and type as the model's time column.""" if self.time_column: @@ -879,7 +908,10 @@ def update_schema(self, schema: MappingSchema) -> None: nested_set( self.mapping_schema, tuple(part.sql(copy=False) for part in table.parts), - {col: dtype.sql(dialect=self.dialect) for col, dtype in mapping_schema.items()}, + { + col: dtype.sql(dialect=self.dialect) + for col, dtype in mapping_schema.items() + }, ) @property @@ -903,7 +935,9 @@ def columns_to_types_or_raise(self) -> t.Dict[str, exp.DataType]: """Returns the mapping of column names to types of this model or raise if not available.""" columns_to_types = self.columns_to_types if columns_to_types is None: - raise SQLMeshError(f"Column information is not available for model '{self.name}'") + raise SQLMeshError( + f"Column information is not available for model '{self.name}'" + ) return columns_to_types @property @@ -912,7 +946,9 @@ def annotated(self) -> bool: if self.columns_to_types is None: return False columns_to_types = { - k: v for k, v in self.columns_to_types.items() if k not in self.managed_columns + k: v + for k, v in self.columns_to_types.items() + if k not in self.managed_columns } if not columns_to_types: return False @@ -975,7 +1011,10 @@ def auto_restatement_croniter(self, value: TimeLike) -> CroniterCache: @property def wap_supported(self) -> bool: - return self.kind.is_materialized and (self.storage_format or "").lower() == "iceberg" + return ( + self.kind.is_materialized + and (self.storage_format or "").lower() == "iceberg" + ) def validate_definition(self) -> None: """Validates the model's definition. @@ -1014,7 +1053,9 @@ def validate_definition(self) -> None: if columns_to_types is not None: missing_keys = unique_keys - set(columns_to_types) if missing_keys: - missing_keys_str = ", ".join(f"'{k}'" for k in sorted(missing_keys)) + missing_keys_str = ", ".join( + f"'{k}'" for k in sorted(missing_keys) + ) raise_config_error( f"{field} keys [{missing_keys_str}] are missing in the model definition", self._path, @@ -1038,7 +1079,10 @@ def validate_definition(self) -> None: if self.kind.is_managed: # TODO: would this sort of logic be better off moved into the Kind? - if self.dialect == "snowflake" and "target_lag" not in self.physical_properties: + if ( + self.dialect == "snowflake" + and "target_lag" not in self.physical_properties + ): raise_config_error( "Snowflake managed tables must specify the 'target_lag' physical property", self._path, @@ -1059,7 +1103,8 @@ def validate_definition(self) -> None: ) if isinstance(self.kind, CustomKind): - from sqlmesh.core.snapshot.evaluator import get_custom_materialization_type_or_raise + from sqlmesh.core.snapshot.evaluator import \ + get_custom_materialization_type_or_raise # Will raise if the custom materialization points to an invalid class get_custom_materialization_type_or_raise(self.kind.materialization) @@ -1108,12 +1153,16 @@ def is_metadata_only_change(self, other: _Node) -> bool: if len(this_statements) != len(other_statements): is_metadata_change = False else: - for this_statement, other_statement in zip(this_statements, other_statements): + for this_statement, other_statement in zip( + this_statements, other_statements + ): this_rendered = ( - self._statement_renderer(this_statement).render() or this_statement + self._statement_renderer(this_statement).render() + or this_statement ) other_rendered = ( - other._statement_renderer(other_statement).render() or other_statement + other._statement_renderer(other_statement).render() + or other_statement ) if this_rendered != other_rendered: is_metadata_change = False @@ -1225,7 +1274,11 @@ def metadata_hash(self) -> str: str(self.end) if self.end else None, str(self.retention) if self.retention else None, str(self.batch_size) if self.batch_size is not None else None, - str(self.batch_concurrency) if self.batch_concurrency is not None else None, + ( + str(self.batch_concurrency) + if self.batch_concurrency is not None + else None + ), json.dumps(self.mapping_schema, sort_keys=True), *sorted(self.tags), *sorted(ref.json(sort_keys=True) for ref in self.all_references), @@ -1271,7 +1324,9 @@ def grants_table_type(self) -> DataObjectType: from sqlmesh.core.engine_adapter.shared import DataObjectType if self.kind.is_view: - if hasattr(self.kind, "materialized") and getattr(self.kind, "materialized", False): + if hasattr(self.kind, "materialized") and getattr( + self.kind, "materialized", False + ): return DataObjectType.MATERIALIZED_VIEW return DataObjectType.VIEW if self.kind.is_managed: @@ -1283,11 +1338,17 @@ def grants_table_type(self) -> DataObjectType: def _additional_metadata(self) -> t.List[str]: additional_metadata = [] - metadata_only_macros = [(k, v) for k, v in self.sorted_python_env if v.is_metadata] + metadata_only_macros = [ + (k, v) for k, v in self.sorted_python_env if v.is_metadata + ] if metadata_only_macros: additional_metadata.append(str(metadata_only_macros)) - for statements in [self.pre_statements_, self.post_statements_, self.on_virtual_update_]: + for statements in [ + self.pre_statements_, + self.post_statements_, + self.on_virtual_update_, + ]: for statement in statements or []: additional_metadata.append(statement.sql) @@ -1465,7 +1526,9 @@ def render_definition( result.append(self.query) result.extend(self.post_statements) if self.on_virtual_update: - result.append(d.VirtualUpdateStatement(expressions=self.on_virtual_update)) + result.append( + d.VirtualUpdateStatement(expressions=self.on_virtual_update) + ) return result @@ -1547,7 +1610,9 @@ def validate_definition(self) -> None: return if not isinstance(query, exp.Query): - raise_config_error("Missing SELECT query in the model definition", self._path) + raise_config_error( + "Missing SELECT query in the model definition", self._path + ) projection_list = query.selects if not projection_list: @@ -1655,7 +1720,9 @@ class SeedModel(_Model): kind: SeedKind seed: Seed - column_hashes_: t.Optional[t.Dict[str, str]] = Field(default=None, alias="column_hashes") + column_hashes_: t.Optional[t.Dict[str, str]] = Field( + default=None, alias="column_hashes" + ) derived_columns_to_types: t.Optional[t.Dict[str, exp.DataType]] = None is_hydrated: bool = True source_type: t.Literal["seed"] = "seed" @@ -1711,7 +1778,9 @@ def render_seed(self) -> t.Iterator[QueryOrDF]: rename_dict = {} for column in columns_to_types: if column not in df: - normalized_name = normalize_identifiers(column, dialect=self.dialect).name + normalized_name = normalize_identifiers( + column, dialect=self.dialect + ).name if normalized_name in df: rename_dict[normalized_name] = column if rename_dict: @@ -1722,7 +1791,8 @@ def render_seed(self) -> t.Iterator[QueryOrDF]: missing_columns = column_names_to_check - set(df.columns) if missing_columns: raise_config_error( - f"Seed model '{self.name}' has missing columns: {missing_columns}", self._path + f"Seed model '{self.name}' has missing columns: {missing_columns}", + self._path, ) # convert all date/time types to native pandas timestamp @@ -1742,7 +1812,9 @@ def render_seed(self) -> t.Iterator[QueryOrDF]: ) for column in bool_columns: - df[column] = df[column].apply(lambda i: None if pd.isna(i) else str_to_bool(str(i))) + df[column] = df[column].apply( + lambda i: None if pd.isna(i) else str_to_bool(str(i)) + ) df.loc[:, string_columns] = df[string_columns].mask( cond=lambda x: x.notna(), # type: ignore @@ -1813,9 +1885,9 @@ def to_dehydrated(self) -> SeedModel: "seed": Seed(content=""), "is_hydrated": False, "column_hashes_": self.column_hashes, - "derived_columns_to_types": self.columns_to_types - if self.columns_to_types_ is None - else None, + "derived_columns_to_types": ( + self.columns_to_types if self.columns_to_types_ is None else None + ), } ) @@ -1914,7 +1986,11 @@ def render( **kwargs.pop("variables", {}), } blueprint_variables = { - k: d.parse_one(v.sql, dialect=self.dialect) if isinstance(v, SqlValue) else v + k: ( + d.parse_one(v.sql, dialect=self.dialect) + if isinstance(v, SqlValue) + else v + ) for k, v in { **env.get(c.SQLMESH_BLUEPRINT_VARS, {}), **env.get(c.SQLMESH_BLUEPRINT_VARS_METADATA, {}), @@ -1930,7 +2006,9 @@ def render( "latest": execution_time, # TODO: Preserved for backward compatibility. Remove in 1.0.0. } df_or_iter = env[self.entrypoint]( - context=context.with_variables(variables, blueprint_variables=blueprint_variables), + context=context.with_variables( + variables, blueprint_variables=blueprint_variables + ), **kwargs, ) @@ -1940,7 +2018,9 @@ def render( for df in df_or_iter: yield df except Exception as e: - raise PythonModelEvalError(format_evaluated_code_exception(e, self.python_env)) + raise PythonModelEvalError( + format_evaluated_code_exception(e, self.python_env) + ) def render_definition( self, @@ -1951,7 +2031,9 @@ def render_definition( # Ignore the provided value for the include_python flag, since the Pyhon model's # definition without Python code is meaningless. return super().render_definition( - include_python=True, include_defaults=include_defaults, render_query=render_query + include_python=True, + include_defaults=include_defaults, + render_query=render_query, ) @property @@ -1977,7 +2059,10 @@ class ExternalModel(_Model): def is_breaking_change(self, previous: Model) -> t.Optional[bool]: if not isinstance(previous, ExternalModel): return None - if not previous.columns_to_types_or_raise.items() - self.columns_to_types_or_raise.items(): + if ( + not previous.columns_to_types_or_raise.items() + - self.columns_to_types_or_raise.items() + ): return False return None @@ -2087,7 +2172,9 @@ def create_models_from_blueprints( default_catalog=loader_kwargs.get("default_catalog"), blueprint_variables=blueprint_variables, ) - gateway_name = rendered_gateway[0].name.lower() if rendered_gateway else None + gateway_name = ( + rendered_gateway[0].name.lower() if rendered_gateway else None + ) elif configured_gateway := (loader_kwargs.get("defaults") or {}).get("gateway"): # Config gateway names are literals, not SQL expressions. In particular, parsing a # gateway such as "secondary-gw" as SQL would interpret it as subtraction. @@ -2297,9 +2384,13 @@ def load_sql_based_model( rendered_defaults = parse_defaults_properties(rendered_defaults, dialect=dialect) # Extract the query and any pre/post statements - query_or_seed_insert, pre_statements, post_statements, on_virtual_update, inline_audits = ( - _split_sql_model_statements(expressions[1:], path, dialect=dialect) - ) + ( + query_or_seed_insert, + pre_statements, + post_statements, + on_virtual_update, + inline_audits, + ) = _split_sql_model_statements(expressions[1:], path, dialect=dialect) meta_fields: t.Dict[str, t.Any] = { "dialect": dialect, @@ -2308,7 +2399,10 @@ def load_sql_based_model( if rendered_meta.comments else None ), - **{prop.name.lower(): prop.args.get("value") for prop in rendered_meta.expressions}, + **{ + prop.name.lower(): prop.args.get("value") + for prop in rendered_meta.expressions + }, **kwargs, } @@ -2495,7 +2589,9 @@ def create_python_model( else: depends_on_rendered = render_expression( expression=exp.Array( - expressions=[exp.maybe_parse(dep, dialect=dialect) for dep in depends_on or []] + expressions=[ + exp.maybe_parse(dep, dialect=dialect) for dep in depends_on or [] + ] ), module_path=module_path, macros=macros, @@ -2510,9 +2606,13 @@ def create_python_model( for dep in t.cast(t.List[exp.Expr], depends_on_rendered)[0].expressions } - used_variables = {k: v for k, v in (variables or {}).items() if k in referenced_variables} + used_variables = { + k: v for k, v in (variables or {}).items() if k in referenced_variables + } if used_variables: - python_env[c.SQLMESH_VARS] = Executable.value(used_variables, sort_root_dict=True) + python_env[c.SQLMESH_VARS] = Executable.value( + used_variables, sort_root_dict=True + ) return _create_model( PythonModel, @@ -2629,7 +2729,8 @@ def _create_model( for statement_field in ["pre_statements", "post_statements", "on_virtual_update"]: if statement_field in defaults: kwargs[statement_field] = [ - exp.maybe_parse(stmt, dialect=dialect) for stmt in defaults[statement_field] + exp.maybe_parse(stmt, dialect=dialect) + for stmt in defaults[statement_field] ] + kwargs.get(statement_field, []) if statement_field in kwargs: # Macros extracted from these statements need to be treated as metadata only @@ -2640,10 +2741,12 @@ def _create_model( statements.append((expr, is_metadata)) kwargs[statement_field] = [ # this to retain the transaction information - stmt - if isinstance(stmt, ParsableSql) - else ParsableSql.from_parsed_expression( - stmt, dialect, use_meta_sql=use_original_sql + ( + stmt + if isinstance(stmt, ParsableSql) + else ParsableSql.from_parsed_expression( + stmt, dialect, use_meta_sql=use_original_sql + ) ) for stmt in kwargs[statement_field] ] @@ -2659,13 +2762,17 @@ def _create_model( if isinstance(getattr(kwargs.get("kind"), "merge_filter", None), exp.Expr): statements.append(kwargs["kind"].merge_filter) - jinja_macro_references, referenced_variables = extract_macro_references_and_variables( - *(gen(e if isinstance(e, exp.Expr) else e[0]) for e in statements) + jinja_macro_references, referenced_variables = ( + extract_macro_references_and_variables( + *(gen(e if isinstance(e, exp.Expr) else e[0]) for e in statements) + ) ) if jinja_macros: jinja_macros = ( - jinja_macros if jinja_macros.trimmed else jinja_macros.trim(jinja_macro_references) + jinja_macros + if jinja_macros.trimmed + else jinja_macros.trim(jinja_macro_references) ) else: jinja_macros = JinjaMacroRegistry() @@ -2677,7 +2784,9 @@ def _create_model( # Merge model-specific audits with default audits if default_audits := defaults.pop("audits", None): - kwargs["audits"] = default_audits + d.extract_function_calls(kwargs.pop("audits", [])) + kwargs["audits"] = default_audits + d.extract_function_calls( + kwargs.pop("audits", []) + ) model = klass( name=name, @@ -2715,7 +2824,9 @@ def _create_model( available_audits = BUILT_IN_AUDITS.keys() | model.audit_definitions.keys() for referenced_audit, audit_args in model.audits: if referenced_audit not in available_audits: - raise_config_error(f"Audit '{referenced_audit}' is undefined", location=path) + raise_config_error( + f"Audit '{referenced_audit}' is undefined", location=path + ) statements.extend( (audit_arg_expression, True) for audit_arg_expression in audit_args.values() @@ -2725,7 +2836,9 @@ def _create_model( for referenced_signal, kwargs in model.signals: if referenced_signal and referenced_signal not in signal_definitions: - raise_config_error(f"Signal '{referenced_signal}' is undefined", location=path) + raise_config_error( + f"Signal '{referenced_signal}' is undefined", location=path + ) statements.extend((signal_kwarg, True) for signal_kwarg in kwargs.values()) @@ -2826,7 +2939,13 @@ def _split_sql_model_statements( raise_config_error("Only one SELECT query is allowed per model", path) query, pos = query_positions[0] - return query, sql_statements[:pos], sql_statements[pos + 1 :], on_virtual_update, inline_audits + return ( + query, + sql_statements[:pos], + sql_statements[pos + 1 :], + on_virtual_update, + inline_audits, + ) def _resolve_model_refs_to_physical_tables( @@ -2849,7 +2968,8 @@ def resolve(ref: str) -> str: physical = table_mapping.get(exp.table_name(table, identify=True)) # Managed model -> physical table; otherwise keep the reference (just unquoted/normalized). return exp.table_name( - exp.to_table(physical, dialect=dialect) if physical else table, identify=False + exp.to_table(physical, dialect=dialect) if physical else table, + identify=False, ) return exp.Literal.string(",".join(resolve(ref) for ref in refs if ref.strip())) @@ -2898,12 +3018,16 @@ def _list_of_calls_to_exp(value: t.List[t.Tuple[str, t.Dict[str, t.Any]]]) -> ex def _has_ordinal_references(query: exp.Query) -> bool: order = query.args.get("order") if order and any( - isinstance(ob.this, exp.Literal) and ob.this.is_number for ob in order.expressions + isinstance(ob.this, exp.Literal) and ob.this.is_number + for ob in order.expressions ): return True group = query.args.get("group") return bool( - group and any(isinstance(gb, exp.Literal) and gb.is_number for gb in group.expressions) + group + and any( + isinstance(gb, exp.Literal) and gb.is_number for gb in group.expressions + ) ) @@ -2947,7 +3071,9 @@ def _added_projection_preserves_cardinality(projection: exp.Expr) -> bool: ) -def _projections_only_safely_added(previous_query: exp.Select, this_query: exp.Select) -> bool: +def _projections_only_safely_added( + previous_query: exp.Select, this_query: exp.Select +) -> bool: """Return whether a SELECT's projections differ only through safe additions. Every previous projection must occur unchanged and in the same order in the current list. @@ -2967,7 +3093,9 @@ def _projections_only_safely_added(previous_query: exp.Select, this_query: exp.S this_index < len(this_projections) and previous_projection != this_projections[this_index] ): - if not _added_projection_preserves_cardinality(this_projections[this_index]): + if not _added_projection_preserves_cardinality( + this_projections[this_index] + ): return False added_before_existing = True @@ -3004,7 +3132,9 @@ def _is_only_projection_additions( This specialized comparison avoids the candidate matching performed by SQLGlot's general tree diff while remaining conservative for every change other than an added projection. """ - expression_pairs: t.List[t.Tuple[exp.Expr, exp.Expr]] = [(previous_query, this_query)] + expression_pairs: t.List[t.Tuple[exp.Expr, exp.Expr]] = [ + (previous_query, this_query) + ] while expression_pairs: previous_expression, this_expression = expression_pairs.pop() @@ -3030,7 +3160,9 @@ def _is_only_projection_additions( and isinstance(this_expression, exp.Select) and arg_key == "expressions" ): - if not _projections_only_safely_added(previous_expression, this_expression): + if not _projections_only_safely_added( + previous_expression, this_expression + ): return False elif len(previous_value) != len(this_value): return False @@ -3188,7 +3320,9 @@ def render_model_defaults( for boolean in {"optimize_query", "allow_partials", "enabled"}: var = rendered_defaults.get(boolean) if var is not None and not isinstance(var, (exp.Boolean, bool)): - raise ConfigError(f"Expected boolean for '{var}', got '{type(var)}' instead") + raise ConfigError( + f"Expected boolean for '{var}', got '{type(var)}' instead" + ) # Validate the 'interval_unit' if present is an Interval Unit var = rendered_defaults.get("interval_unit") @@ -3257,7 +3391,9 @@ def render_expression( "post": _list_of_calls_to_exp, "audits": _list_of_calls_to_exp, "columns_to_types_": lambda value: exp.Schema( - expressions=[exp.ColumnDef(this=exp.to_column(c), kind=t) for c, t in value.items()] + expressions=[ + exp.ColumnDef(this=exp.to_column(c), kind=t) for c, t in value.items() + ] ), "column_descriptions_": lambda value: exp.Schema( expressions=[exp.to_column(c).eq(d) for c, d in value.items()] @@ -3271,11 +3407,17 @@ def render_expression( "allow_partials": exp.convert, "signals": lambda values: exp.tuple_( *( - exp.func( - name, *(exp.PropertyEQ(this=exp.var(k), expression=v) for k, v in args.items()) + ( + exp.func( + name, + *( + exp.PropertyEQ(this=exp.var(k), expression=v) + for k, v in args.items() + ), + ) + if name + else exp.Tuple(expressions=[exp.var(k).eq(v) for k, v in args.items()]) ) - if name - else exp.Tuple(expressions=[exp.var(k).eq(v) for k, v in args.items()]) for name, args in values ) ), @@ -3299,9 +3441,9 @@ def clickhouse_partition_func( ) -> exp.Expr: # `toMonday()` function accepts a Date or DateTime type column - col_type = (columns_to_types and columns_to_types.get(column.name)) or exp.DataType.build( - "UNKNOWN" - ) + col_type = ( + columns_to_types and columns_to_types.get(column.name) + ) or exp.DataType.build("UNKNOWN") col_type_is_conformable = col_type.is_type( exp.DataType.Type.DATE, exp.DataType.Type.DATE32, @@ -3317,7 +3459,9 @@ def clickhouse_partition_func( if col_type.is_type(exp.DataType.Type.UNKNOWN): return exp.func( "toMonday", - exp.cast(column, exp.DataType.build("DateTime64(9, 'UTC')", dialect="clickhouse")), + exp.cast( + column, exp.DataType.build("DateTime64(9, 'UTC')", dialect="clickhouse") + ), dialect="clickhouse", ) @@ -3325,7 +3469,9 @@ def clickhouse_partition_func( return exp.cast( exp.func( "toMonday", - exp.cast(column, exp.DataType.build("DateTime64(9, 'UTC')", dialect="clickhouse")), + exp.cast( + column, exp.DataType.build("DateTime64(9, 'UTC')", dialect="clickhouse") + ), dialect="clickhouse", ), col_type, diff --git a/sqlmesh/core/model/kind.py b/sqlmesh/core/model/kind.py index a8960fc3e1..3b493b7f3e 100644 --- a/sqlmesh/core/model/kind.py +++ b/sqlmesh/core/model/kind.py @@ -2,7 +2,6 @@ import typing as t from enum import Enum -from typing_extensions import Self from pydantic import Field from sqlglot import exp @@ -10,32 +9,20 @@ from sqlglot.optimizer.qualify_columns import quote_identifiers from sqlglot.optimizer.simplify import gen from sqlglot.time import format_time +from typing_extensions import Self from sqlmesh.core import dialect as d -from sqlmesh.core.model.common import ( - parse_properties, - properties_validator, - validate_extra_and_required_fields, -) +from sqlmesh.core.model.common import (parse_properties, properties_validator, + validate_extra_and_required_fields) from sqlmesh.core.model.seed import CsvSettings from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.pydantic import ( - PydanticModel, - SQLGlotBool, - SQLGlotColumn, - SQLGlotListOfFieldsOrStar, - SQLGlotListOfFields, - SQLGlotPositiveInt, - SQLGlotString, - SQLGlotCron, - ValidationInfo, - column_validator, - field_validator, - get_dialect, - validate_string, - validate_expression, -) - +from sqlmesh.utils.pydantic import (PydanticModel, SQLGlotBool, SQLGlotColumn, + SQLGlotCron, SQLGlotListOfFields, + SQLGlotListOfFieldsOrStar, + SQLGlotPositiveInt, SQLGlotString, + ValidationInfo, column_validator, + field_validator, get_dialect, + validate_expression, validate_string) if t.TYPE_CHECKING: from sqlmesh.core._typing import CustomMaterializationProperties @@ -105,7 +92,10 @@ def is_scd_type_2(self) -> bool: @property def is_scd_type_2_by_time(self) -> bool: - return self.model_kind_name in {ModelKindName.SCD_TYPE_2, ModelKindName.SCD_TYPE_2_BY_TIME} + return self.model_kind_name in { + ModelKindName.SCD_TYPE_2, + ModelKindName.SCD_TYPE_2_BY_TIME, + } @property def is_scd_type_2_by_column(self) -> bool: @@ -130,7 +120,9 @@ def is_symbolic(self) -> bool: @property def is_materialized(self) -> bool: - return self.model_kind_name is not None and not (self.is_symbolic or self.is_view) + return self.model_kind_name is not None and not ( + self.is_symbolic or self.is_view + ) @property def only_execution_time(self) -> bool: @@ -247,7 +239,9 @@ def _on_destructive_change_validator( ) -> t.Any: if v and not isinstance(v, OnDestructiveChange): return OnDestructiveChange( - v.this.upper() if isinstance(v, (exp.Identifier, exp.Literal)) else v.upper() + v.this.upper() + if isinstance(v, (exp.Identifier, exp.Literal)) + else v.upper() ) return v @@ -257,7 +251,9 @@ def _on_additive_change_validator( ) -> t.Any: if v and not isinstance(v, OnAdditiveChange): return OnAdditiveChange( - v.this.upper() if isinstance(v, (exp.Identifier, exp.Literal)) else v.upper() + v.this.upper() + if isinstance(v, (exp.Identifier, exp.Literal)) + else v.upper() ) return v @@ -266,9 +262,9 @@ def _on_additive_change_validator( _on_additive_change_validator ) -on_destructive_change_validator = field_validator("on_destructive_change", mode="before")( - _on_destructive_change_validator -) +on_destructive_change_validator = field_validator( + "on_destructive_change", mode="before" +)(_on_destructive_change_validator) class _ModelKind(PydanticModel, ModelKindMixin): @@ -330,7 +326,10 @@ def to_expression(self, dialect: str) -> exp.Expr: expressions=[ self.column, exp.Literal.string( - format_time(self.format, d.Dialect.get_or_raise(dialect).INVERSE_TIME_MAPPING) + format_time( + self.format, + d.Dialect.get_or_raise(dialect).INVERSE_TIME_MAPPING, + ) ), ] ) @@ -345,7 +344,9 @@ def create(cls, v: t.Any, dialect: str) -> Self: raise ConfigError("Time Column cannot be empty.") column_expr = v.expressions[0] column = ( - exp.column(column_expr) if isinstance(column_expr, exp.Identifier) else column_expr + exp.column(column_expr) + if isinstance(column_expr, exp.Identifier) + else column_expr ) format = v.expressions[1].name if len(v.expressions) > 1 else None elif isinstance(v, exp.Expr): @@ -369,7 +370,9 @@ def create(cls, v: t.Any, dialect: str) -> Self: else: raise ConfigError(f"Invalid time_column: '{v}'.") - column = quote_identifiers(normalize_identifiers(column, dialect=dialect), dialect=dialect) + column = quote_identifiers( + normalize_identifiers(column, dialect=dialect), dialect=dialect + ) column.meta["dialect"] = dialect return cls(column=column, format=format) @@ -381,7 +384,9 @@ def _kind_dialect_validator(cls: t.Type, v: t.Optional[str]) -> str: return v -kind_dialect_validator = field_validator("dialect", mode="before")(_kind_dialect_validator) +kind_dialect_validator = field_validator("dialect", mode="before")( + _kind_dialect_validator +) class _Incremental(_ModelKind): @@ -487,7 +492,12 @@ def to_expression( } ), *( - [_property("auto_restatement_intervals", self.auto_restatement_intervals)] + [ + _property( + "auto_restatement_intervals", + self.auto_restatement_intervals, + ) + ] if self.auto_restatement_intervals is not None else [] ), @@ -496,16 +506,22 @@ def to_expression( @property def data_hash_values(self) -> t.List[t.Optional[str]]: - return [*super().data_hash_values, gen(self.time_column.column), self.time_column.format] + return [ + *super().data_hash_values, + gen(self.time_column.column), + self.time_column.format, + ] @property def metadata_hash_values(self) -> t.List[t.Optional[str]]: return [ *super().metadata_hash_values, str(self.partition_by_time_column), - str(self.auto_restatement_intervals) - if self.auto_restatement_intervals is not None - else None, + ( + str(self.auto_restatement_intervals) + if self.auto_restatement_intervals is not None + else None + ), ] @@ -540,7 +556,9 @@ def _when_matched_validator( v = t.cast(exp.Whens, d.parse_one(v, into=exp.Whens, dialect=dialect)) v = validate_expression(v, dialect=dialect) - return t.cast(exp.Whens, v.transform(d.replace_merge_table_aliases, dialect=dialect)) + return t.cast( + exp.Whens, v.transform(d.replace_merge_table_aliases, dialect=dialect) + ) @field_validator("merge_filter", mode="before") def _merge_filter_validator( @@ -587,7 +605,9 @@ def to_expression( class IncrementalByPartitionKind(_Incremental): - name: t.Literal[ModelKindName.INCREMENTAL_BY_PARTITION] = ModelKindName.INCREMENTAL_BY_PARTITION + name: t.Literal[ModelKindName.INCREMENTAL_BY_PARTITION] = ( + ModelKindName.INCREMENTAL_BY_PARTITION + ) forward_only: t.Literal[True] = True disable_restatement: SQLGlotBool = False @@ -624,7 +644,9 @@ def to_expression( class IncrementalUnmanagedKind(_Incremental): - name: t.Literal[ModelKindName.INCREMENTAL_UNMANAGED] = ModelKindName.INCREMENTAL_UNMANAGED + name: t.Literal[ModelKindName.INCREMENTAL_UNMANAGED] = ( + ModelKindName.INCREMENTAL_UNMANAGED + ) insert_overwrite: SQLGlotBool = False forward_only: SQLGlotBool = True disable_restatement: SQLGlotBool = True @@ -722,7 +744,10 @@ def data_hash_values(self) -> t.List[t.Optional[str]]: csv_setting_values = (self.csv_settings or CsvSettings()).dict().values() return [ *super().data_hash_values, - *(v if isinstance(v, (str, type(None))) else str(v) for v in csv_setting_values), + *( + v if isinstance(v, (str, type(None))) else str(v) + for v in csv_setting_values + ), ] @property @@ -741,10 +766,14 @@ class FullKind(_ModelKind): class _SCDType2Kind(_Incremental): dialect: t.Optional[str] = Field(None, validate_default=True) unique_key: SQLGlotListOfFields - valid_from_name: SQLGlotColumn = Field(exp.column("valid_from"), validate_default=True) + valid_from_name: SQLGlotColumn = Field( + exp.column("valid_from"), validate_default=True + ) valid_to_name: SQLGlotColumn = Field(exp.column("valid_to"), validate_default=True) invalidate_hard_deletes: SQLGlotBool = False - time_data_type: exp.DataType = Field(exp.DataType.build("TIMESTAMP"), validate_default=True) + time_data_type: exp.DataType = Field( + exp.DataType.build("TIMESTAMP"), validate_default=True + ) batch_size: t.Optional[SQLGlotPositiveInt] = None forward_only: SQLGlotBool = True @@ -752,13 +781,15 @@ class _SCDType2Kind(_Incremental): _dialect_validator = kind_dialect_validator - _always_validate_column = field_validator("valid_from_name", "valid_to_name", mode="before")( - column_validator - ) + _always_validate_column = field_validator( + "valid_from_name", "valid_to_name", mode="before" + )(column_validator) @field_validator("time_data_type", mode="before") @classmethod - def _time_data_type_validator(cls, v: t.Union[str, exp.Expr], values: t.Any) -> exp.Expr: + def _time_data_type_validator( + cls, v: t.Union[str, exp.Expr], values: t.Any + ) -> exp.Expr: if isinstance(v, exp.Expr) and not isinstance(v, exp.DataType): v = v.name dialect = get_dialect(values) @@ -824,7 +855,9 @@ class SCDType2ByTimeKind(_SCDType2Kind): name: t.Literal[ModelKindName.SCD_TYPE_2, ModelKindName.SCD_TYPE_2_BY_TIME] = ( ModelKindName.SCD_TYPE_2_BY_TIME ) - updated_at_name: SQLGlotColumn = Field(exp.column("updated_at"), validate_default=True) + updated_at_name: SQLGlotColumn = Field( + exp.column("updated_at"), validate_default=True + ) updated_at_as_valid_from: SQLGlotBool = False _always_validate_updated_at = field_validator("updated_at_name", mode="before")( @@ -856,7 +889,9 @@ def to_expression( class SCDType2ByColumnKind(_SCDType2Kind): - name: t.Literal[ModelKindName.SCD_TYPE_2_BY_COLUMN] = ModelKindName.SCD_TYPE_2_BY_COLUMN + name: t.Literal[ModelKindName.SCD_TYPE_2_BY_COLUMN] = ( + ModelKindName.SCD_TYPE_2_BY_COLUMN + ) columns: SQLGlotListOfFieldsOrStar execution_time_as_valid_from: SQLGlotBool = False updated_at_name: t.Optional[SQLGlotColumn] = None @@ -883,9 +918,11 @@ def to_expression( *(expressions or []), *_properties( { - "columns": exp.Tuple(expressions=self.columns) - if isinstance(self.columns, list) - else self.columns, + "columns": ( + exp.Tuple(expressions=self.columns) + if isinstance(self.columns, list) + else self.columns + ), "execution_time_as_valid_from": self.execution_time_as_valid_from, } ), @@ -991,7 +1028,11 @@ def data_hash_values(self) -> t.List[t.Optional[str]]: return [ *super().data_hash_values, self.materialization, - gen(self.materialization_properties_) if self.materialization_properties_ else None, + ( + gen(self.materialization_properties_) + if self.materialization_properties_ + else None + ), str(self.lookback) if self.lookback is not None else None, ] @@ -1004,9 +1045,11 @@ def metadata_hash_values(self) -> t.List[t.Optional[str]]: str(self.forward_only), str(self.disable_restatement), self.auto_restatement_cron, - str(self.auto_restatement_intervals) - if self.auto_restatement_intervals is not None - else None, + ( + str(self.auto_restatement_intervals) + if self.auto_restatement_intervals is not None + else None + ), ] def to_expression( @@ -1078,7 +1121,9 @@ def model_kind_type_from_name(name: t.Optional[str]) -> t.Type[ModelKind]: return t.cast(t.Type[ModelKind], klass) -def create_model_kind(v: t.Any, dialect: str, defaults: t.Dict[str, t.Any]) -> ModelKind: +def create_model_kind( + v: t.Any, dialect: str, defaults: t.Dict[str, t.Any] +) -> ModelKind: if isinstance(v, _ModelKind): return t.cast(ModelKind, v) @@ -1124,7 +1169,8 @@ def create_model_kind(v: t.Any, dialect: str, defaults: t.Dict[str, t.Any]) -> M if kind_type == CustomKind: # load the custom materialization class and check if it uses a custom kind type - from sqlmesh.core.snapshot.evaluator import get_custom_materialization_type + from sqlmesh.core.snapshot.evaluator import \ + get_custom_materialization_type if "materialization" not in props: raise ConfigError( @@ -1151,12 +1197,16 @@ def create_model_kind(v: t.Any, dialect: str, defaults: t.Dict[str, t.Any]) -> M return model_kind_type_from_name(name)(name=name) # type: ignore -def _model_kind_validator(cls: t.Type, v: t.Any, info: t.Optional[ValidationInfo]) -> ModelKind: +def _model_kind_validator( + cls: t.Type, v: t.Any, info: t.Optional[ValidationInfo] +) -> ModelKind: dialect = get_dialect(info.data) if info else "" return create_model_kind(v, dialect, {}) -model_kind_validator: t.Callable = field_validator("kind", mode="before")(_model_kind_validator) +model_kind_validator: t.Callable = field_validator("kind", mode="before")( + _model_kind_validator +) def _property(name: str, value: t.Any) -> exp.Property: diff --git a/sqlmesh/core/model/meta.py b/sqlmesh/core/model/meta.py index 94956dff99..917a02a1f7 100644 --- a/sqlmesh/core/model/meta.py +++ b/sqlmesh/core/model/meta.py @@ -3,53 +3,39 @@ import typing as t from enum import Enum from functools import cached_property -from typing_extensions import Self from pydantic import Field from sqlglot import Dialect, exp, parse_one from sqlglot.helper import ensure_collection, ensure_list from sqlglot.optimizer.normalize_identifiers import normalize_identifiers +from typing_extensions import Self from sqlmesh.core import dialect as d from sqlmesh.core.config.common import VirtualEnvironmentMode -from sqlmesh.core.constants import LIQUID_CLUSTERING_KEYWORDS from sqlmesh.core.config.linter import LinterConfig +from sqlmesh.core.constants import LIQUID_CLUSTERING_KEYWORDS from sqlmesh.core.dialect import normalize_model_name -from sqlmesh.utils import classproperty -from sqlmesh.core.model.common import ( - bool_validator, - default_catalog_validator, - depends_on_validator, - properties_validator, - parse_properties, -) -from sqlmesh.core.model.kind import ( - CustomKind, - IncrementalByUniqueKeyKind, - ModelKind, - OnDestructiveChange, - SCDType2ByColumnKind, - SCDType2ByTimeKind, - TimeColumn, - ViewKind, - model_kind_validator, - OnAdditiveChange, -) +from sqlmesh.core.model.common import (bool_validator, + default_catalog_validator, + depends_on_validator, parse_properties, + properties_validator) +from sqlmesh.core.model.kind import (CustomKind, IncrementalByUniqueKeyKind, + ModelKind, OnAdditiveChange, + OnDestructiveChange, SCDType2ByColumnKind, + SCDType2ByTimeKind, TimeColumn, ViewKind, + model_kind_validator) from sqlmesh.core.node import _Node, str_or_exp_to_str from sqlmesh.core.reference import Reference +from sqlmesh.utils import classproperty from sqlmesh.utils.date import TimeLike from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.pydantic import ( - ValidationInfo, - field_validator, - list_of_fields_validator, - model_validator, - get_dialect, - validation_data, -) +from sqlmesh.utils.pydantic import (ValidationInfo, field_validator, + get_dialect, list_of_fields_validator, + model_validator, validation_data) if t.TYPE_CHECKING: - from sqlmesh.core._typing import CustomMaterializationProperties, SessionProperties + from sqlmesh.core._typing import (CustomMaterializationProperties, + SessionProperties) from sqlmesh.core.engine_adapter._typing import GrantsConfig FunctionCall = t.Tuple[str, t.Dict[str, exp.Expr]] @@ -98,7 +84,9 @@ class ModelMeta(_Node): clustered_by: t.List[exp.Expr] = [] default_catalog: t.Optional[str] = None depends_on_: t.Optional[t.Set[str]] = Field(default=None, alias="depends_on") - columns_to_types_: t.Optional[t.Dict[str, exp.DataType]] = Field(default=None, alias="columns") + columns_to_types_: t.Optional[t.Dict[str, exp.DataType]] = Field( + default=None, alias="columns" + ) column_descriptions_: t.Optional[t.Dict[str, str]] = Field( default=None, alias="column_descriptions" ) @@ -106,9 +94,15 @@ class ModelMeta(_Node): grains: t.List[exp.Expr] = [] references: t.List[exp.Expr] = [] physical_schema_override: t.Optional[str] = None - physical_properties_: t.Optional[exp.Tuple] = Field(default=None, alias="physical_properties") - virtual_properties_: t.Optional[exp.Tuple] = Field(default=None, alias="virtual_properties") - session_properties_: t.Optional[exp.Tuple] = Field(default=None, alias="session_properties") + physical_properties_: t.Optional[exp.Tuple] = Field( + default=None, alias="physical_properties" + ) + virtual_properties_: t.Optional[exp.Tuple] = Field( + default=None, alias="virtual_properties" + ) + session_properties_: t.Optional[exp.Tuple] = Field( + default=None, alias="session_properties" + ) allow_partials: bool = False signals: t.List[FunctionCall] = [] enabled: bool = True @@ -131,7 +125,10 @@ class ModelMeta(_Node): @field_validator("audits", "signals", mode="before") def _func_call_validator(cls, v: t.Any, field: t.Any) -> t.Any: - is_signal = getattr(field, "name" if hasattr(field, "name") else "field_name") == "signals" + is_signal = ( + getattr(field, "name" if hasattr(field, "name") else "field_name") + == "signals" + ) return d.extract_function_calls(v, allow_tuples=is_signal) @@ -159,13 +156,18 @@ def _normalize(value: t.Any) -> t.Any: value = _normalize(v) return value.name if isinstance(value, exp.Expr) else value if isinstance(v, (list, tuple)): - return [cls._validate_value_or_tuple(elm, data, normalize=normalize) for elm in v] + return [ + cls._validate_value_or_tuple(elm, data, normalize=normalize) + for elm in v + ] return v @field_validator("table_format", "storage_format", mode="before") def _format_validator(cls, v: t.Any, info: ValidationInfo) -> t.Optional[str]: - if isinstance(v, exp.Expr) and not (isinstance(v, (exp.Literal, exp.Identifier))): + if isinstance(v, exp.Expr) and not ( + isinstance(v, (exp.Literal, exp.Identifier)) + ): return v.sql(validation_data(info).get("dialect")) return str_or_exp_to_str(v) @@ -190,7 +192,9 @@ def _gateway_validator(cls, v: t.Any) -> t.Optional[str]: return gateway and gateway.lower() @field_validator("partitioned_by_", "clustered_by", mode="before") - def _partition_and_cluster_validator(cls, v: t.Any, info: ValidationInfo) -> t.List[exp.Expr]: + def _partition_and_cluster_validator( + cls, v: t.Any, info: ValidationInfo + ) -> t.List[exp.Expr]: field = info.field_name or "" dialect = (get_dialect(info) or "").lower() @@ -202,11 +206,11 @@ def _partition_and_cluster_validator(cls, v: t.Any, info: ValidationInfo) -> t.L # this branch gets hit when we are deserializing from json because `partitioned_by` is stored as a List[str] # however, we should only invoke this if the list contains strings because this validator is also # called by Python models which might pass a List[exp.Expression] - string_to_parse = ( - f"({','.join(v)})" # recreate the (a, b, c) part of "partitioned_by (a, b, c)" - ) + string_to_parse = f"({','.join(v)})" # recreate the (a, b, c) part of "partitioned_by (a, b, c)" parsed = parse_one( - string_to_parse, into=exp.PartitionedByProperty, dialect=get_dialect(info) + string_to_parse, + into=exp.PartitionedByProperty, + dialect=get_dialect(info), ) v = parsed.this.expressions if isinstance(parsed.this, exp.Schema) else v @@ -218,9 +222,12 @@ def _partition_and_cluster_validator(cls, v: t.Any, info: ValidationInfo) -> t.L # Restore keyword sentinels (AUTO/NONE) before list_of_fields_validator normalises # them into quoted columns. v = [ - exp.Var(this=item.upper()) - if isinstance(item, str) and item.upper() in LIQUID_CLUSTERING_KEYWORDS - else item + ( + exp.Var(this=item.upper()) + if isinstance(item, str) + and item.upper() in LIQUID_CLUSTERING_KEYWORDS + else item + ) for item in v ] @@ -251,7 +258,10 @@ def _partition_and_cluster_validator(cls, v: t.Any, info: ValidationInfo) -> t.L return expressions @field_validator( - "columns_to_types_", "derived_columns_to_types", mode="before", check_fields=False + "columns_to_types_", + "derived_columns_to_types", + mode="before", + check_fields=False, ) def _columns_validator( cls, v: t.Any, info: ValidationInfo @@ -266,7 +276,9 @@ def _columns_validator( raise ConfigError(f"Missing data type for column '{column.name}'.") expr.meta["dialect"] = dialect - columns_to_types[normalize_identifiers(column, dialect=dialect).name] = expr + columns_to_types[ + normalize_identifiers(column, dialect=dialect).name + ] = expr return columns_to_types @@ -293,7 +305,9 @@ def _columns_validator( ): sql_repr = expr.sql(dialect=dialect) try: - normalized = parse_one(sql_repr, read=dialect, into=exp.DataType) + normalized = parse_one( + sql_repr, read=dialect, into=exp.DataType + ) if normalized is not None: expr = normalized except Exception: @@ -324,7 +338,10 @@ def _column_descriptions_validator( raw_col_descriptions = ( vs if isinstance(vs, dict) - else {".".join([part.this for part in v.this.parts]): v.expression.name for v in vs} + else { + ".".join([part.this for part in v.this.parts]): v.expression.name + for v in vs + } ) col_descriptions = { @@ -395,7 +412,8 @@ def session_properties_validator(cls, v: t.Any, info: ValidationInfo) -> t.Any: if prop_name == "query_label": query_label = eq.right if not isinstance( - query_label, (exp.Array, exp.Tuple, exp.Paren, d.MacroFunc, d.MacroVar) + query_label, + (exp.Array, exp.Tuple, exp.Paren, d.MacroFunc, d.MacroVar), ): raise ConfigError( "Invalid value for `session_properties.query_label`. Must be an array or tuple." @@ -411,7 +429,10 @@ def session_properties_validator(cls, v: t.Any, info: ValidationInfo) -> t.Any: if not ( isinstance(label_tuple, exp.Tuple) and len(label_tuple.expressions) == 2 - and all(isinstance(label, exp.Literal) for label in label_tuple.expressions) + and all( + isinstance(label, exp.Literal) + for label in label_tuple.expressions + ) ): raise ConfigError( "Invalid entry in `session_properties.query_label`. Must be tuples of string literals with length 2." @@ -505,7 +526,9 @@ def _root_validator(self) -> Self: ): name = field[:-1] if field.endswith("_") else field raise ValueError(f"{name} field cannot be set for {kind.name} models") - if kind.is_incremental_by_partition and not getattr(self, "partitioned_by_", None): + if kind.is_incremental_by_partition and not getattr( + self, "partitioned_by_", None + ): raise ValueError(f"partitioned_by field is required for {kind.name} models") # needs to be in a mode=after model validator so that the field validators have run to convert from Expression -> str @@ -535,7 +558,8 @@ def time_column(self) -> t.Optional[TimeColumn]: @property def unique_key(self) -> t.List[exp.Expr]: if isinstance( - self.kind, (SCDType2ByTimeKind, SCDType2ByColumnKind, IncrementalByUniqueKeyKind) + self.kind, + (SCDType2ByTimeKind, SCDType2ByColumnKind, IncrementalByUniqueKeyKind), ): return self.kind.unique_key return [] @@ -572,14 +596,18 @@ def batch_concurrency(self) -> t.Optional[int]: def physical_properties(self) -> t.Dict[str, exp.Expr]: """A dictionary of properties that will be applied to the physical layer. It replaces table_properties which is deprecated.""" if self.physical_properties_: - return {e.this.name: e.expression for e in self.physical_properties_.expressions} + return { + e.this.name: e.expression for e in self.physical_properties_.expressions + } return {} @cached_property def virtual_properties(self) -> t.Dict[str, exp.Expr]: """A dictionary of properties that will be applied to the virtual layer.""" if self.virtual_properties_: - return {e.this.name: e.expression for e in self.virtual_properties_.expressions} + return { + e.this.name: e.expression for e in self.virtual_properties_.expressions + } return {} @property @@ -614,17 +642,25 @@ def grants(self) -> t.Optional[GrantsConfig]: grants_dict[permission_name] = grantee_list except ConfigError as e: permission_name = ( - eq_expr.left.name if hasattr(eq_expr.left, "name") else str(eq_expr.left) + eq_expr.left.name + if hasattr(eq_expr.left, "name") + else str(eq_expr.left) + ) + raise ConfigError( + f"Invalid grants configuration for '{permission_name}': {e}" ) - raise ConfigError(f"Invalid grants configuration for '{permission_name}': {e}") return grants_dict if grants_dict else None @property def all_references(self) -> t.List[Reference]: """All references including grains.""" - return [Reference(model_name=self.name, expression=e, unique=True) for e in self.grains] + [ - Reference(model_name=self.name, expression=e, unique=True) for e in self.references + return [ + Reference(model_name=self.name, expression=e, unique=True) + for e in self.grains + ] + [ + Reference(model_name=self.name, expression=e, unique=True) + for e in self.references ] @property @@ -634,7 +670,9 @@ def on(self) -> t.List[str]: on: t.List[str] = [] for expr in [ref.expression for ref in self.all_references if ref.unique]: if isinstance(expr, exp.Tuple): - on.extend([key.this.sql(dialect=self.dialect) for key in expr.expressions]) + on.extend( + [key.this.sql(dialect=self.dialect) for key in expr.expressions] + ) else: # Handle a single Column or Paren expression on.append(expr.this.sql(dialect=self.dialect)) @@ -706,7 +744,9 @@ def flatten_expr(expr: exp.Expr) -> None: for elem in expr.expressions: flatten_expr(elem) elif isinstance(expr, (exp.Tuple, exp.Paren)): - expressions = [expr.unnest()] if isinstance(expr, exp.Paren) else expr.expressions + expressions = ( + [expr.unnest()] if isinstance(expr, exp.Paren) else expr.expressions + ) for elem in expressions: flatten_expr(elem) else: diff --git a/sqlmesh/core/model/schema.py b/sqlmesh/core/model/schema.py index e29cacade0..202a1d6051 100644 --- a/sqlmesh/core/model/schema.py +++ b/sqlmesh/core/model/schema.py @@ -7,11 +7,9 @@ from sqlglot.errors import SchemaError from sqlglot.schema import MappingSchema -from sqlmesh.core.model.cache import ( - load_optimized_query_and_mapping, - optimized_query_cache_pool, - OptimizedQueryCache, -) +from sqlmesh.core.model.cache import (OptimizedQueryCache, + load_optimized_query_and_mapping, + optimized_query_cache_pool) if t.TYPE_CHECKING: from sqlmesh.core.model.definition import Model @@ -87,7 +85,9 @@ def process_models(completed_model: t.Optional[Model] = None) -> None: for future in as_completed(futures): try: futures.remove(future) - fqn, entry_name, data_hash, metadata_hash, mapping_schema = future.result() + fqn, entry_name, data_hash, metadata_hash, mapping_schema = ( + future.result() + ) model = models[fqn] model._data_hash = data_hash model._metadata_hash = metadata_hash diff --git a/sqlmesh/core/model/seed.py b/sqlmesh/core/model/seed.py index ff12085690..21fada6ecf 100644 --- a/sqlmesh/core/model/seed.py +++ b/sqlmesh/core/model/seed.py @@ -37,7 +37,9 @@ class CsvSettings(PydanticModel): na_values: t.Optional[NaValues] = None keep_default_na: t.Optional[bool] = None - @field_validator("doublequote", "skipinitialspace", "keep_default_na", mode="before") + @field_validator( + "doublequote", "skipinitialspace", "keep_default_na", mode="before" + ) @classmethod def _bool_validator(cls, v: t.Any) -> t.Optional[bool]: if v is None: @@ -45,7 +47,12 @@ def _bool_validator(cls, v: t.Any) -> t.Optional[bool]: return parse_bool(v) @field_validator( - "delimiter", "quotechar", "escapechar", "lineterminator", "encoding", mode="before" + "delimiter", + "quotechar", + "escapechar", + "lineterminator", + "encoding", + mode="before", ) @classmethod def _str_validator(cls, v: t.Any) -> t.Optional[str]: @@ -83,7 +90,9 @@ def _na_values_validator(cls, v: t.Any) -> t.Optional[NaValues]: return [e.to_py() for e in expressions] except ValueError as e: - logger.warning(f"Failed to coerce na_values '{v}', proceeding with defaults. {str(e)}") + logger.warning( + f"Failed to coerce na_values '{v}', proceeding with defaults. {str(e)}" + ) return None @@ -107,7 +116,9 @@ def column_hashes(self) -> t.Dict[str, str]: for column_name in df.columns } - def read(self, batch_size: t.Optional[int] = None) -> t.Generator[pd.DataFrame, None, None]: + def read( + self, batch_size: t.Optional[int] = None + ) -> t.Generator[pd.DataFrame, None, None]: df = self._get_df() batch_size = batch_size or df.size @@ -145,7 +156,9 @@ class Seed(PydanticModel): content: str - def reader(self, dialect: str = "", settings: t.Optional[CsvSettings] = None) -> CsvSeedReader: + def reader( + self, dialect: str = "", settings: t.Optional[CsvSettings] = None + ) -> CsvSeedReader: return CsvSeedReader(self.content, dialect, settings or CsvSettings()) diff --git a/sqlmesh/core/node.py b/sqlmesh/core/node.py index 3ee97becc9..35ab2e2d82 100644 --- a/sqlmesh/core/node.py +++ b/sqlmesh/core/node.py @@ -12,13 +12,8 @@ from sqlmesh.utils.cron import CroniterCache from sqlmesh.utils.date import TimeLike, to_datetime, validate_date_range from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.pydantic import ( - PydanticModel, - SQLGlotCron, - field_validator, - model_validator, - PRIVATE_FIELDS, -) +from sqlmesh.utils.pydantic import (PRIVATE_FIELDS, PydanticModel, SQLGlotCron, + field_validator, model_validator) if t.TYPE_CHECKING: from sqlmesh.core._typing import Self @@ -51,12 +46,16 @@ def from_cron(klass, cron: str) -> IntervalUnit: if not interval_seconds: samples = [croniter.get_next() for _ in range(5)] - interval_seconds = int(min(b - a for a, b in zip(samples, samples[1:])).total_seconds()) + interval_seconds = int( + min(b - a for a, b in zip(samples, samples[1:])).total_seconds() + ) for unit, seconds in INTERVAL_SECONDS.items(): if seconds <= interval_seconds: return unit - raise ConfigError(f"Invalid cron '{cron}': must run at a frequency of 5 minutes or slower.") + raise ConfigError( + f"Invalid cron '{cron}': must run at a frequency of 5 minutes or slower." + ) @property def is_date_granularity(self) -> bool: @@ -80,7 +79,11 @@ def is_hour(self) -> bool: @property def is_minute(self) -> bool: - return self in (IntervalUnit.FIVE_MINUTE, IntervalUnit.QUARTER_HOUR, IntervalUnit.HALF_HOUR) + return self in ( + IntervalUnit.FIVE_MINUTE, + IntervalUnit.QUARTER_HOUR, + IntervalUnit.HALF_HOUR, + ) @property def cron_expr(self) -> str: @@ -315,7 +318,9 @@ class _Node(DbtInfoMixin, PydanticModel): end: t.Optional[TimeLike] = None cron: SQLGlotCron = "@daily" cron_tz: t.Optional[zoneinfo.ZoneInfo] = None - interval_unit_: t.Optional[IntervalUnit] = Field(alias="interval_unit", default=None) + interval_unit_: t.Optional[IntervalUnit] = Field( + alias="interval_unit", default=None + ) tags: t.List[str] = [] stamp: t.Optional[str] = None dbt_node_info_: t.Optional[DbtNodeInfo] = Field(alias="dbt_node_info", default=None) @@ -360,7 +365,9 @@ def _date_validator(cls, v: t.Any) -> t.Optional[TimeLike]: if isinstance(v, exp.Expr): v = v.name if v and not to_datetime(v): - raise ConfigError(f"'{v}' needs to be time-like: https://pypi.org/project/dateparser") + raise ConfigError( + f"'{v}' needs to be time-like: https://pypi.org/project/dateparser" + ) return v @field_validator("owner", "description", "stamp", mode="before") @@ -370,7 +377,9 @@ def _string_expr_validator(cls, v: t.Any) -> t.Optional[str]: @field_validator("interval_unit_", mode="before") @classmethod - def _interval_unit_validator(cls, v: t.Any) -> t.Optional[t.Union[IntervalUnit, str]]: + def _interval_unit_validator( + cls, v: t.Any + ) -> t.Optional[t.Union[IntervalUnit, str]]: if isinstance(v, IntervalUnit): return v v = str_or_exp_to_str(v) @@ -452,7 +461,10 @@ def is_metadata_only_change(self, previous: _Node) -> bool: Returns: True if this node is a metadata only change, False otherwise. """ - return self.data_hash == previous.data_hash and self.metadata_hash != previous.metadata_hash + return ( + self.data_hash == previous.data_hash + and self.metadata_hash != previous.metadata_hash + ) def is_data_change(self, previous: _Node) -> bool: """Determines if this node is a data change in relation to the `previous` node. @@ -464,7 +476,8 @@ def is_data_change(self, previous: _Node) -> bool: True if this node is a data change, False otherwise. """ return ( - self.data_hash != previous.data_hash or self.metadata_hash != previous.metadata_hash + self.data_hash != previous.data_hash + or self.metadata_hash != previous.metadata_hash ) and not self.is_metadata_only_change(previous) def croniter(self, value: TimeLike) -> CroniterCache: @@ -511,7 +524,9 @@ def cron_floor(self, value: TimeLike, estimate: bool = False) -> datetime: Returns: The timestamp floor. """ - return self.croniter(self.cron_next(value, estimate=estimate)).get_prev(estimate=True) + return self.croniter(self.cron_next(value, estimate=estimate)).get_prev( + estimate=True + ) def text_diff(self, other: Node, rendered: bool = False) -> str: """Produce a text diff against another node. diff --git a/sqlmesh/core/notification_target.py b/sqlmesh/core/notification_target.py index fba6e36f66..f8fc4b8fd2 100644 --- a/sqlmesh/core/notification_target.py +++ b/sqlmesh/core/notification_target.py @@ -10,7 +10,8 @@ from sqlmesh.core.console import Console, get_console from sqlmesh.integrations import slack -from sqlmesh.utils.errors import AuditError, ConfigError, MissingDependencyError +from sqlmesh.utils.errors import (AuditError, ConfigError, + MissingDependencyError) from sqlmesh.utils.pydantic import PydanticModel if t.TYPE_CHECKING: @@ -78,7 +79,9 @@ class BaseNotificationTarget(PydanticModel, frozen=True): type_: str notify_on: t.FrozenSet[NotificationEvent] = frozenset() - def send(self, notification_status: NotificationStatus, msg: str, **kwargs: t.Any) -> None: + def send( + self, notification_status: NotificationStatus, msg: str, **kwargs: t.Any + ) -> None: """Sends notification with the provided message. Args: @@ -120,7 +123,10 @@ def notify_run_start(self, environment: str, *args: t.Any, **kwargs: t.Any) -> N Args: environment: The target environment of the run. """ - self.send(NotificationStatus.INFO, f"SQLMesh run started for environment `{environment}`.") + self.send( + NotificationStatus.INFO, + f"SQLMesh run started for environment `{environment}`.", + ) def notify_run_end(self, environment: str, *args: t.Any, **kwargs: t.Any) -> None: """Notify when a SQLMesh run ends. @@ -129,7 +135,8 @@ def notify_run_end(self, environment: str, *args: t.Any, **kwargs: t.Any) -> Non environment: The target environment of the run. """ self.send( - NotificationStatus.SUCCESS, f"SQLMesh run finished for environment `{environment}`." + NotificationStatus.SUCCESS, + f"SQLMesh run finished for environment `{environment}`.", ) def notify_migration_start(self, *args: t.Any, **kwargs: t.Any) -> None: @@ -164,7 +171,9 @@ def notify_run_failure(self, exc: str, *args: t.Any, **kwargs: t.Any) -> None: """ self.send(NotificationStatus.FAILURE, "SQLMesh run failed.", exc=exc) - def notify_audit_failure(self, audit_error: AuditError, *args: t.Any, **kwargs: t.Any) -> None: + def notify_audit_failure( + self, audit_error: AuditError, *args: t.Any, **kwargs: t.Any + ) -> None: """Notify in the case of an audit failure. Args: @@ -190,7 +199,9 @@ class BaseTextBasedNotificationTarget(BaseNotificationTarget): A base class for unstructured notification targets (e.g.: console, email, etc.) """ - def send_text_message(self, notification_status: NotificationStatus, msg: str) -> None: + def send_text_message( + self, notification_status: NotificationStatus, msg: str + ) -> None: """Send the notification message as text.""" def send( @@ -207,7 +218,9 @@ def send( elif exc: error = exc - self.send_text_message(notification_status, msg if error is None else f"{msg}\n{error}") + self.send_text_message( + notification_status, msg if error is None else f"{msg}\n{error}" + ) class ConsoleNotificationTarget(BaseTextBasedNotificationTarget): @@ -224,7 +237,9 @@ def console(self) -> Console: self._console = get_console() return self._console - def send_text_message(self, notification_status: NotificationStatus, msg: str) -> None: + def send_text_message( + self, notification_status: NotificationStatus, msg: str + ) -> None: if notification_status.is_success: self.console.log_success(msg) elif notification_status.is_failure: @@ -251,7 +266,9 @@ def send( } composed = slack.message().add_primary_blocks( - slack.header_block(f"{status_emoji[notification_status]} SQLMesh Notification"), + slack.header_block( + f"{status_emoji[notification_status]} SQLMesh Notification" + ), slack.context_block(f"*Status:* `{notification_status.value}`"), slack.divider_block(), slack.text_section_block(f"*Message*: {msg}"), @@ -274,7 +291,8 @@ def send( *details, slack.divider_block(), slack.context_block( - f"*SQLMesh Version:* {_sqlmesh_version()}", f"*Python Version:* {sys.version}" + f"*SQLMesh Version:* {_sqlmesh_version()}", + f"*Python Version:* {sys.version}", ), ) @@ -440,7 +458,9 @@ class NotificationTargetManager: def __init__( self, - notification_targets: t.Dict[NotificationEvent, t.Set[NotificationTarget]] | None = None, + notification_targets: ( + t.Dict[NotificationEvent, t.Set[NotificationTarget]] | None + ) = None, user_notification_targets: t.Dict[str, t.Set[NotificationTarget]] | None = None, username: str | None = None, ) -> None: @@ -454,7 +474,9 @@ def notify(self, event: NotificationEvent, *args: t.Any, **kwargs: t.Any) -> Non self.notify_user(event, self.username, *args, **kwargs) else: for notification_target in self.notification_targets.get(event, set()): - notify_func = self._get_notification_function(notification_target, event) + notify_func = self._get_notification_function( + notification_target, event + ) notify_func(*args, **kwargs) def notify_user( @@ -464,7 +486,9 @@ def notify_user( notification_targets = self.user_notification_targets.get(username, set()) for notification_target in notification_targets: if event in notification_target.notify_on: - notify_func = self._get_notification_function(notification_target, event) + notify_func = self._get_notification_function( + notification_target, event + ) notify_func(*args, **kwargs) def _get_notification_function( diff --git a/sqlmesh/core/plan/__init__.py b/sqlmesh/core/plan/__init__.py index 8b3ba63e55..106d4ec1c2 100644 --- a/sqlmesh/core/plan/__init__.py +++ b/sqlmesh/core/plan/__init__.py @@ -1,12 +1,9 @@ from sqlmesh.core.plan.builder import PlanBuilder as PlanBuilder -from sqlmesh.core.plan.definition import ( - Plan as Plan, - EvaluatablePlan as EvaluatablePlan, - PlanStatus as PlanStatus, - SnapshotIntervals as SnapshotIntervals, -) -from sqlmesh.core.plan.evaluator import ( - BuiltInPlanEvaluator as BuiltInPlanEvaluator, - PlanEvaluator as PlanEvaluator, -) +from sqlmesh.core.plan.definition import EvaluatablePlan as EvaluatablePlan +from sqlmesh.core.plan.definition import Plan as Plan +from sqlmesh.core.plan.definition import PlanStatus as PlanStatus +from sqlmesh.core.plan.definition import SnapshotIntervals as SnapshotIntervals +from sqlmesh.core.plan.evaluator import \ + BuiltInPlanEvaluator as BuiltInPlanEvaluator +from sqlmesh.core.plan.evaluator import PlanEvaluator as PlanEvaluator from sqlmesh.core.plan.explainer import PlanExplainer as PlanExplainer diff --git a/sqlmesh/core/plan/builder.py b/sqlmesh/core/plan/builder.py index a6307a9ffd..27f5206874 100644 --- a/sqlmesh/core/plan/builder.py +++ b/sqlmesh/core/plan/builder.py @@ -4,49 +4,30 @@ import re import typing as t from collections import defaultdict -from functools import cached_property from datetime import datetime +from functools import cached_property - +from sqlmesh.core.config import (AutoCategorizationMode, CategorizerConfig, + EnvironmentSuffixTarget) from sqlmesh.core.console import PlanBuilderConsole, get_console -from sqlmesh.core.config import ( - AutoCategorizationMode, - CategorizerConfig, - EnvironmentSuffixTarget, -) from sqlmesh.core.context_diff import ContextDiff from sqlmesh.core.environment import EnvironmentNamingInfo -from sqlmesh.core.plan.common import should_force_rebuild, is_breaking_kind_change -from sqlmesh.core.plan.definition import ( - Plan, - SnapshotMapping, - UserProvidedFlags, - earliest_interval_start, -) -from sqlmesh.core.schema_diff import ( - get_schema_differ, - has_drop_alteration, - has_additive_alteration, - TableAlterOperation, -) -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Snapshot, - SnapshotChangeCategory, -) +from sqlmesh.core.plan.common import (is_breaking_kind_change, + should_force_rebuild) +from sqlmesh.core.plan.definition import (Plan, SnapshotMapping, + UserProvidedFlags, + earliest_interval_start) +from sqlmesh.core.schema_diff import (TableAlterOperation, get_schema_differ, + has_additive_alteration, + has_drop_alteration) +from sqlmesh.core.snapshot import (DeployabilityIndex, Snapshot, + SnapshotChangeCategory) from sqlmesh.core.snapshot.categorizer import categorize_change from sqlmesh.core.snapshot.definition import Interval, SnapshotId from sqlmesh.utils import columns_to_types_all_known, random_id from sqlmesh.utils.dag import DAG -from sqlmesh.utils.date import ( - TimeLike, - now, - to_datetime, - yesterday_ds, - to_timestamp, - time_like_to_str, - is_relative, -) +from sqlmesh.utils.date import (TimeLike, is_relative, now, time_like_to_str, + to_datetime, to_timestamp, yesterday_ds) from sqlmesh.utils.errors import NoChangesPlanError, PlanError logger = logging.getLogger(__name__) @@ -164,7 +145,9 @@ def __init__( self._categorizer_config = categorizer_config or CategorizerConfig() self._auto_categorization_enabled = auto_categorization_enabled self._include_unmodified = include_unmodified - self._restate_models = set(restate_models) if restate_models is not None else None + self._restate_models = ( + set(restate_models) if restate_models is not None else None + ) self._restate_all_snapshots = restate_all_snapshots self._effective_from = effective_from @@ -200,16 +183,20 @@ def __init__( self._start = default_start or yesterday_ds() self._plan_id: str = random_id() - self._model_fqn_to_snapshot = {s.name: s for s in self._context_diff.snapshots.values()} + self._model_fqn_to_snapshot = { + s.name: s for s in self._context_diff.snapshots.values() + } self.override_start = start is not None self.override_end = end is not None - self.environment_naming_info = EnvironmentNamingInfo.from_environment_catalog_mapping( - environment_catalog_mapping or {}, - name=self._context_diff.environment, - suffix_target=environment_suffix_target, - normalize_name=self._context_diff.normalize_environment_name, - gateway_managed=self._context_diff.gateway_managed_virtual_layer, + self.environment_naming_info = ( + EnvironmentNamingInfo.from_environment_catalog_mapping( + environment_catalog_mapping or {}, + name=self._context_diff.environment, + suffix_target=environment_suffix_target, + normalize_name=self._context_diff.normalize_environment_name, + gateway_managed=self._context_diff.gateway_managed_virtual_layer, + ) ) self._latest_plan: t.Optional[Plan] = None @@ -223,14 +210,18 @@ def is_start_and_end_allowed(self) -> bool: def start(self) -> t.Optional[TimeLike]: if self._start and is_relative(self._start): # only do this for relative expressions otherwise inclusive date strings like '2020-01-01' can be turned into exclusive timestamps eg '2020-01-01 00:00:00' - return to_datetime(self._start, relative_base=to_datetime(self.execution_time)) + return to_datetime( + self._start, relative_base=to_datetime(self.execution_time) + ) return self._start @property def end(self) -> t.Optional[TimeLike]: if self._end and is_relative(self._end): # only do this for relative expressions otherwise inclusive date strings like '2020-01-01' can be turned into exclusive timestamps eg '2020-01-01 00:00:00' - return to_datetime(self._end, relative_base=to_datetime(self.execution_time)) + return to_datetime( + self._end, relative_base=to_datetime(self.execution_time) + ) return self._end @cached_property @@ -269,7 +260,9 @@ def set_effective_from(self, effective_from: t.Optional[TimeLike]) -> PlanBuilde self._latest_plan = None return self - def set_choice(self, snapshot: Snapshot, choice: SnapshotChangeCategory) -> PlanBuilder: + def set_choice( + self, snapshot: Snapshot, choice: SnapshotChangeCategory + ) -> PlanBuilder: """Sets a snapshot version based on the user choice. Args: @@ -284,7 +277,9 @@ def set_choice(self, snapshot: Snapshot, choice: SnapshotChangeCategory) -> Plan not self._context_diff.directly_modified(snapshot.name) and snapshot.snapshot_id not in self._context_diff.added ): - raise PlanError(f"Only directly modified models can be categorized ({snapshot.name}).") + raise PlanError( + f"Only directly modified models can be categorized ({snapshot.name})." + ) self._choices[snapshot.snapshot_id] = choice self._latest_plan = None @@ -308,7 +303,9 @@ def build(self) -> Plan: self._apply_effective_from() dag = self._build_dag() - directly_modified, indirectly_modified = self._build_directly_and_indirectly_modified(dag) + directly_modified, indirectly_modified = ( + self._build_directly_and_indirectly_modified(dag) + ) self._check_destructive_additive_changes(directly_modified) self._categorize_snapshots(dag, indirectly_modified) @@ -326,7 +323,9 @@ def build(self) -> Plan: restatements = self._build_restatements( dag, - earliest_interval_start(self._context_diff.snapshots.values(), self.execution_time), + earliest_interval_start( + self._context_diff.snapshots.values(), self.execution_time + ), ) models_to_backfill = self._build_models_to_backfill(dag, restatements) @@ -426,7 +425,9 @@ def _build_restatements( # Add restate snapshots and their downstream snapshots for model_fqn in restate_models: if model_fqn not in self._model_fqn_to_snapshot: - raise PlanError(f"Cannot restate model '{model_fqn}'. Model does not exist.") + raise PlanError( + f"Cannot restate model '{model_fqn}'. Model does not exist." + ) # Get restatement intervals for all restated snapshots and make sure that if an incremental snapshot expands it's # restatement range that it's downstream dependencies all expand their restatement ranges as well. @@ -439,7 +440,9 @@ def _build_restatements( # Since we are traversing the graph in topological order and the largest interval range is pushed down # the graph we just have to check our immediate parents in the graph and not the whole upstream graph. restating_parents = [ - self._context_diff.snapshots[s] for s in snapshot.parents if s in restatements + self._context_diff.snapshots[s] + for s in snapshot.parents + if s in restatements ] if not restating_parents and snapshot.name not in restate_models: @@ -452,7 +455,9 @@ def _build_restatements( "Run the restatement against the production environment instead to restate this model." ) continue - elif (not self._is_dev or not snapshot.is_paused) and snapshot.disable_restatement: + elif ( + not self._is_dev or not snapshot.is_paused + ) and snapshot.disable_restatement: self._console.log_warning( f"Cannot restate model '{snapshot.name}'. " "Restatement is disabled for this model to prevent possible data loss. " @@ -464,10 +469,14 @@ def _build_restatements( continue possible_intervals = { - restatements[p.snapshot_id] for p in restating_parents if p.is_incremental + restatements[p.snapshot_id] + for p in restating_parents + if p.is_incremental } removal_start = ( - self._forward_only_preview_start(snapshot, start, end) if is_preview else start + self._forward_only_preview_start(snapshot, start, end) + if is_preview + else start ) possible_intervals.add( snapshot.get_removal_interval( @@ -484,7 +493,9 @@ def _build_restatements( # We may be tasked with restating a time range smaller than the target snapshot interval unit # For example, restating an hour of Hourly Model A, which has a downstream dependency of Daily Model B # we need to ensure the whole affected day in Model B is restated - floored_snapshot_start = snapshot.node.interval_unit.cron_floor(snapshot_start) + floored_snapshot_start = snapshot.node.interval_unit.cron_floor( + snapshot_start + ) floored_snapshot_end = snapshot.node.interval_unit.cron_floor(snapshot_end) if to_timestamp(floored_snapshot_end) < snapshot_end: snapshot_start = to_timestamp(floored_snapshot_start) @@ -549,10 +560,12 @@ def _build_models_to_backfill( backfill_models = ( self._backfill_models if self._backfill_models is not None - else [r.name for r in restatements] - # Only backfill models explicitly marked for restatement. - if self._restate_models - else None + else ( + [r.name for r in restatements] + # Only backfill models explicitly marked for restatement. + if self._restate_models + else None + ) ) if backfill_models is None: return None @@ -571,7 +584,9 @@ def _adjust_snapshot_intervals(self) -> None: for new, old in self._context_diff.modified_snapshots.values(): if not new.is_model or not old.is_model: continue - is_same_version = old.version_get_or_generate() == new.version_get_or_generate() + is_same_version = ( + old.version_get_or_generate() == new.version_get_or_generate() + ) if is_same_version and should_force_rebuild(old, new): # If the difference between 2 snapshots requires a full rebuild, # then clear the intervals for the new snapshot. @@ -584,7 +599,9 @@ def _adjust_snapshot_intervals(self) -> None: if new.is_forward_only: new.dev_intervals = new.intervals.copy() - def _check_destructive_additive_changes(self, directly_modified: t.Set[SnapshotId]) -> None: + def _check_destructive_additive_changes( + self, directly_modified: t.Set[SnapshotId] + ) -> None: for s_id in sorted(directly_modified): if s_id.name not in self._context_diff.modified_snapshots: continue @@ -593,11 +610,13 @@ def _check_destructive_additive_changes(self, directly_modified: t.Set[SnapshotI needs_destructive_check = snapshot.needs_destructive_check( self._allow_destructive_models ) - needs_additive_check = snapshot.needs_additive_check(self._allow_additive_models) - # should we raise/warn if this snapshot has/inherits a destructive change? - should_raise_or_warn = (self._is_forward_only_change(s_id) or self._forward_only) and ( - needs_destructive_check or needs_additive_check + needs_additive_check = snapshot.needs_additive_check( + self._allow_additive_models ) + # should we raise/warn if this snapshot has/inherits a destructive change? + should_raise_or_warn = ( + self._is_forward_only_change(s_id) or self._forward_only + ) and (needs_destructive_check or needs_additive_check) if not should_raise_or_warn or not snapshot.is_model: continue @@ -608,9 +627,9 @@ def _check_destructive_additive_changes(self, directly_modified: t.Set[SnapshotI old_columns_to_types = old.model.columns_to_types or {} new_columns_to_types = new.model.columns_to_types or {} - if columns_to_types_all_known(old_columns_to_types) and columns_to_types_all_known( - new_columns_to_types - ): + if columns_to_types_all_known( + old_columns_to_types + ) and columns_to_types_all_known(new_columns_to_types): alter_operations = t.cast( t.List[TableAlterOperation], get_schema_differ(snapshot.model.dialect).compare_columns( @@ -645,7 +664,9 @@ def _check_destructive_additive_changes(self, directly_modified: t.Set[SnapshotI error=not snapshot.model.on_additive_change.is_warn, ) if snapshot.model.on_additive_change.is_error: - raise PlanError("Plan requires an additive change to a forward-only model.") + raise PlanError( + "Plan requires an additive change to a forward-only model." + ) def _categorize_snapshots( self, dag: DAG[SnapshotId], indirectly_modified: SnapshotMapping @@ -677,7 +698,9 @@ def _categorize_snapshots( if s_id in self._context_diff.added: snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only) elif s_id.name in self._context_diff.modified_snapshots: - self._categorize_snapshot(snapshot, forward_only, dag, indirectly_modified) + self._categorize_snapshot( + snapshot, forward_only, dag, indirectly_modified + ) def _categorize_snapshot( self, @@ -707,7 +730,9 @@ def _categorize_snapshot( break if s_id_with_missing_columns is None: - change_category = categorize_change(new, old, config=self._categorizer_config) + change_category = categorize_change( + new, old, config=self._categorizer_config + ) if change_category is not None: snapshot.categorize_as(change_category, forward_only) else: @@ -715,14 +740,18 @@ def _categorize_snapshot( new.model.source_type, AutoCategorizationMode.OFF ) if mode == AutoCategorizationMode.FULL: - snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only) + snapshot.categorize_as( + SnapshotChangeCategory.BREAKING, forward_only + ) elif self._context_diff.indirectly_modified(snapshot.name): if snapshot.is_materialized_view and not forward_only: # We categorize changes as breaking to allow for instantaneous switches in a virtual layer. # Otherwise, there might be a potentially long downtime during MVs recreation. # In the case of forward-only changes this optimization is not applicable because we want to continue # using the same (existing) table version. - snapshot.categorize_as(SnapshotChangeCategory.INDIRECT_BREAKING, forward_only) + snapshot.categorize_as( + SnapshotChangeCategory.INDIRECT_BREAKING, forward_only + ) return all_upstream_forward_only = set() @@ -744,9 +773,14 @@ def _categorize_snapshot( forward_only = True if direct_parent_categories.intersection( - {SnapshotChangeCategory.BREAKING, SnapshotChangeCategory.INDIRECT_BREAKING} + { + SnapshotChangeCategory.BREAKING, + SnapshotChangeCategory.INDIRECT_BREAKING, + } ): - snapshot.categorize_as(SnapshotChangeCategory.INDIRECT_BREAKING, forward_only) + snapshot.categorize_as( + SnapshotChangeCategory.INDIRECT_BREAKING, forward_only + ) elif not direct_parent_categories: snapshot.categorize_as( self._get_orphaned_indirect_change_category(snapshot), forward_only @@ -754,7 +788,9 @@ def _categorize_snapshot( elif all_upstream_categories == {SnapshotChangeCategory.METADATA}: snapshot.categorize_as(SnapshotChangeCategory.METADATA, forward_only) else: - snapshot.categorize_as(SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only) + snapshot.categorize_as( + SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only + ) else: # Metadata updated. snapshot.categorize_as(SnapshotChangeCategory.METADATA, forward_only) @@ -769,7 +805,9 @@ def _get_orphaned_indirect_change_category( This function is used to infer the correct change category for such downstream snapshots based on change categories of their parents. """ - previous_snapshot = self._context_diff.modified_snapshots[indirect_snapshot.name][1] + previous_snapshot = self._context_diff.modified_snapshots[ + indirect_snapshot.name + ][1] previous_parent_snapshot_ids = {p.name: p for p in previous_snapshot.parents} current_parent_snapshots = [ @@ -783,7 +821,9 @@ def _get_orphaned_indirect_change_category( if current_parent_snapshot.name not in previous_parent_snapshot_ids: # This is a new parent so falling back to INDIRECT_BREAKING return SnapshotChangeCategory.INDIRECT_BREAKING - pevious_parent_snapshot_id = previous_parent_snapshot_ids[current_parent_snapshot.name] + pevious_parent_snapshot_id = previous_parent_snapshot_ids[ + current_parent_snapshot.name + ] if current_parent_snapshot.snapshot_id == pevious_parent_snapshot_id: # There were no new versions of this parent since the previous version of this snapshot, @@ -805,7 +845,10 @@ def _get_orphaned_indirect_change_category( return SnapshotChangeCategory.INDIRECT_BREAKING if previous_parent_categories.intersection( - {SnapshotChangeCategory.BREAKING, SnapshotChangeCategory.INDIRECT_BREAKING} + { + SnapshotChangeCategory.BREAKING, + SnapshotChangeCategory.INDIRECT_BREAKING, + } ): # One of the new parents in the chain was breaking so this indirect snapshot is breaking return SnapshotChangeCategory.INDIRECT_BREAKING @@ -830,7 +873,9 @@ def _get_orphaned_indirect_change_category( def _apply_effective_from(self) -> None: if self._effective_from: if not self._forward_only: - raise PlanError("Effective date can only be set for a forward-only plan.") + raise PlanError( + "Effective date can only be set for a forward-only plan." + ) if to_datetime(self._effective_from) > now(): raise PlanError("Effective date cannot be in the future.") @@ -838,7 +883,10 @@ def _apply_effective_from(self) -> None: if ( snapshot.evaluatable and not snapshot.disable_restatement - and (not snapshot.full_history_restatement_only or not snapshot.is_incremental) + and ( + not snapshot.full_history_restatement_only + or not snapshot.is_incremental + ) ): snapshot.effective_from = self._effective_from @@ -854,7 +902,9 @@ def _is_forward_only_change(self, s_id: SnapshotId) -> bool: if snapshot.is_model and is_breaking_kind_change(old, snapshot): return False return ( - snapshot.is_model and snapshot.model.forward_only and bool(snapshot.previous_versions) + snapshot.is_model + and snapshot.model.forward_only + and bool(snapshot.previous_versions) ) def _is_new_snapshot(self, snapshot: Snapshot) -> bool: @@ -862,7 +912,9 @@ def _is_new_snapshot(self, snapshot: Snapshot) -> bool: return snapshot.snapshot_id in self._context_diff.new_snapshots def _ensure_valid_date_range(self) -> None: - if (self.override_start or self.override_end) and not self.is_start_and_end_allowed: + if ( + self.override_start or self.override_end + ) and not self.is_start_and_end_allowed: raise PlanError( "The start and end dates can't be set for a production plan without restatements." ) @@ -892,7 +944,10 @@ def _ensure_valid_date_range(self) -> None: ) for model_name in models_to_check: if snapshot := self._model_fqn_to_snapshot.get(model_name): - if snapshot.node.start is None or to_datetime(snapshot.node.start) > end_ts: + if ( + snapshot.node.start is None + or to_datetime(snapshot.node.start) > end_ts + ): raise PlanError( f"Model '{model_name}': Start date / time '({time_like_to_str(start_ts)})' can't be greater than end date / time '({time_like_to_str(end_ts)})'.\n" f"Set the `start` attribute in your project config model defaults to avoid this issue." @@ -901,7 +956,9 @@ def _ensure_valid_date_range(self) -> None: def _ensure_no_broken_references(self) -> None: for snapshot in self._context_diff.snapshots.values(): broken_references = { - x.name for x in self._context_diff.removed_snapshots.values() if not x.is_external + x.name + for x in self._context_diff.removed_snapshots.values() + if not x.is_external } & {x for x in snapshot.node.depends_on} if broken_references: broken_references_msg = ", ".join(f"'{x}'" for x in broken_references) diff --git a/sqlmesh/core/plan/common.py b/sqlmesh/core/plan/common.py index bece17639c..d3a1c387e2 100644 --- a/sqlmesh/core/plan/common.py +++ b/sqlmesh/core/plan/common.py @@ -1,11 +1,13 @@ from __future__ import annotations -import typing as t + import logging +import typing as t from dataclasses import dataclass, field -from sqlmesh.core.state_sync import StateReader -from sqlmesh.core.snapshot import Snapshot, SnapshotId, SnapshotIdAndVersion, SnapshotNameVersion +from sqlmesh.core.snapshot import (Snapshot, SnapshotId, SnapshotIdAndVersion, + SnapshotNameVersion) from sqlmesh.core.snapshot.definition import Interval +from sqlmesh.core.state_sync import StateReader from sqlmesh.utils.dag import DAG from sqlmesh.utils.date import now_timestamp @@ -119,7 +121,9 @@ def identify_restatement_intervals_across_snapshot_versions( affected_snapshot_names = [ x - for x in ([restate_snapshot_name] + env_dag.downstream(restate_snapshot_name)) + for x in ( + [restate_snapshot_name] + env_dag.downstream(restate_snapshot_name) + ) if x not in disable_restatement_models ] @@ -131,25 +135,33 @@ def identify_restatement_intervals_across_snapshot_versions( if affected_snapshot.name_version in prod_name_versions: continue - clear_request = snapshot_intervals_to_clear.get(affected_snapshot.snapshot_id) + clear_request = snapshot_intervals_to_clear.get( + affected_snapshot.snapshot_id + ) if not clear_request: clear_request = SnapshotIntervalClearRequest( snapshot=affected_snapshot.id_and_version, interval=interval ) - snapshot_intervals_to_clear[affected_snapshot.snapshot_id] = clear_request + snapshot_intervals_to_clear[affected_snapshot.snapshot_id] = ( + clear_request + ) clear_request.environment_names |= set([env.name]) # snapshot_intervals_to_clear now contains the entire hierarchy of affected snapshots based # on building the DAG for each environment and including downstream snapshots # but, what if there are affected snapshots that arent part of any environment? - unique_snapshot_names = set(snapshot_id.name for snapshot_id in snapshot_intervals_to_clear) + unique_snapshot_names = set( + snapshot_id.name for snapshot_id in snapshot_intervals_to_clear + ) current_ts = current_ts or now_timestamp() all_matching_non_prod_snapshots = { s.snapshot_id: s for s in state_reader.get_snapshots_by_names( - snapshot_names=unique_snapshot_names, current_ts=current_ts, exclude_expired=True + snapshot_names=unique_snapshot_names, + current_ts=current_ts, + exclude_expired=True, ) # Don't clear intervals for a snapshot if it shares the same physical version with prod. # Otherwise, prod will be affected by what should be a dev operation @@ -178,9 +190,13 @@ def identify_restatement_intervals_across_snapshot_versions( for remaining_snapshot_id in remaining_snapshot_ids: remaining_snapshot = all_matching_non_prod_snapshots[remaining_snapshot_id] - snapshot_intervals_to_clear[remaining_snapshot_id] = SnapshotIntervalClearRequest( - snapshot=remaining_snapshot, - interval=snapshot_name_to_widest_interval[remaining_snapshot_id.name], + snapshot_intervals_to_clear[remaining_snapshot_id] = ( + SnapshotIntervalClearRequest( + snapshot=remaining_snapshot, + interval=snapshot_name_to_widest_interval[ + remaining_snapshot_id.name + ], + ) ) # for any affected full_history_restatement_only snapshots, we need to widen the intervals being restated to diff --git a/sqlmesh/core/plan/definition.py b/sqlmesh/core/plan/definition.py index 866299eff8..00b728fb26 100644 --- a/sqlmesh/core/plan/definition.py +++ b/sqlmesh/core/plan/definition.py @@ -5,27 +5,21 @@ from datetime import datetime from enum import Enum from functools import cached_property + from pydantic import Field from sqlmesh.core.context_diff import ContextDiff -from sqlmesh.core.environment import Environment, EnvironmentNamingInfo, EnvironmentStatements -from sqlmesh.utils.metaprogramming import Executable # noqa +from sqlmesh.core.environment import (Environment, EnvironmentNamingInfo, + EnvironmentStatements) from sqlmesh.core.node import IntervalUnit -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Intervals, - Snapshot, - earliest_start_date, - merge_intervals, - missing_intervals, -) -from sqlmesh.core.snapshot.definition import ( - Interval, - SnapshotId, - SnapshotTableInfo, - format_intervals, -) +from sqlmesh.core.snapshot import (DeployabilityIndex, Intervals, Snapshot, + earliest_start_date, merge_intervals, + missing_intervals) +from sqlmesh.core.snapshot.definition import (Interval, SnapshotId, + SnapshotTableInfo, + format_intervals) from sqlmesh.utils.date import TimeLike, now, to_datetime, to_timestamp +from sqlmesh.utils.metaprogramming import Executable # noqa from sqlmesh.utils.pydantic import PydanticModel SnapshotMapping = t.Dict[SnapshotId, t.Set[SnapshotId]] @@ -153,17 +147,25 @@ def snapshots(self) -> t.Dict[SnapshotId, Snapshot]: return self.context_diff.snapshots @cached_property - def modified_snapshots(self) -> t.Dict[SnapshotId, t.Union[Snapshot, SnapshotTableInfo]]: + def modified_snapshots( + self, + ) -> t.Dict[SnapshotId, t.Union[Snapshot, SnapshotTableInfo]]: """Returns the modified (either directly or indirectly) snapshots.""" return { - **{s_id: self.context_diff.snapshots[s_id] for s_id in sorted(self.directly_modified)}, + **{ + s_id: self.context_diff.snapshots[s_id] + for s_id in sorted(self.directly_modified) + }, **{ s_id: self.context_diff.snapshots[s_id] for downstream_s_ids in self.indirectly_modified.values() for s_id in sorted(downstream_s_ids) }, **self.context_diff.removed_snapshots, - **{s_id: self.context_diff.snapshots[s_id] for s_id in sorted(self.metadata_updated)}, + **{ + s_id: self.context_diff.snapshots[s_id] + for s_id in sorted(self.metadata_updated) + }, } @cached_property @@ -187,7 +189,11 @@ def missing_intervals(self) -> t.List[SnapshotIntervals]: intervals = [ SnapshotIntervals(snapshot_id=snapshot.snapshot_id, intervals=missing) for snapshot, missing in missing_intervals( - [s for s in self.snapshots.values() if self.is_selected_for_backfill(s.name)], + [ + s + for s in self.snapshots.values() + if self.is_selected_for_backfill(s.name) + ], start=self.provided_start or self._earliest_interval_start, end=self.provided_end, execution_time=self.execution_time, @@ -226,10 +232,16 @@ def environment(self) -> Environment: ], } elif not self.include_unmodified: - promotable_snapshot_ids = self.context_diff.promotable_snapshot_ids.copy() + promotable_snapshot_ids = ( + self.context_diff.promotable_snapshot_ids.copy() + ) promoted_snapshot_ids = ( - [s.snapshot_id for s in snapshots if s.snapshot_id in promotable_snapshot_ids] + [ + s.snapshot_id + for s in snapshots + if s.snapshot_id in promotable_snapshot_ids + ] if promotable_snapshot_ids is not None else None ) @@ -282,7 +294,8 @@ def to_evaluatable(self) -> EvaluatablePlan: ignore_cron=self.ignore_cron, directly_modified_snapshots=sorted(self.directly_modified), indirectly_modified_snapshots={ - s.name: sorted(snapshot_ids) for s, snapshot_ids in self.indirectly_modified.items() + s.name: sorted(snapshot_ids) + for s, snapshot_ids in self.indirectly_modified.items() }, metadata_updated_snapshots=sorted(self.metadata_updated), removed_snapshots=sorted(self.context_diff.removed_snapshots), diff --git a/sqlmesh/core/plan/evaluator.py b/sqlmesh/core/plan/evaluator.py index f2f432a97e..f934a3819b 100644 --- a/sqlmesh/core/plan/evaluator.py +++ b/sqlmesh/core/plan/evaluator.py @@ -17,31 +17,28 @@ import abc import logging import typing as t + from sqlmesh.core import analytics from sqlmesh.core import constants as c from sqlmesh.core.console import Console, get_console -from sqlmesh.core.environment import EnvironmentNamingInfo, execute_environment_statements +from sqlmesh.core.environment import (EnvironmentNamingInfo, + execute_environment_statements) from sqlmesh.core.macros import RuntimeStage -from sqlmesh.core.snapshot.definition import to_view_mapping, SnapshotTableInfo from sqlmesh.core.plan import stages +from sqlmesh.core.plan.common import \ + identify_restatement_intervals_across_snapshot_versions from sqlmesh.core.plan.definition import EvaluatablePlan from sqlmesh.core.scheduler import Scheduler -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Snapshot, - SnapshotEvaluator, - SnapshotIntervals, - SnapshotId, - SnapshotInfoLike, - SnapshotCreationFailedError, -) -from sqlmesh.utils import to_snake_case +from sqlmesh.core.snapshot import (DeployabilityIndex, Snapshot, + SnapshotCreationFailedError, + SnapshotEvaluator, SnapshotId, + SnapshotInfoLike, SnapshotIntervals) +from sqlmesh.core.snapshot.definition import SnapshotTableInfo, to_view_mapping from sqlmesh.core.state_sync import StateSync -from sqlmesh.core.plan.common import identify_restatement_intervals_across_snapshot_versions -from sqlmesh.utils import CorrelationId +from sqlmesh.utils import CorrelationId, to_snake_case from sqlmesh.utils.concurrency import NodeExecutionFailedError -from sqlmesh.utils.errors import PlanError, ConflictingPlanError, SQLMeshError from sqlmesh.utils.date import now, to_timestamp +from sqlmesh.utils.errors import ConflictingPlanError, PlanError, SQLMeshError logger = logging.getLogger(__name__) @@ -71,7 +68,9 @@ def __init__( self, state_sync: StateSync, snapshot_evaluator: SnapshotEvaluator, - create_scheduler: t.Callable[[t.Iterable[Snapshot], SnapshotEvaluator], Scheduler], + create_scheduler: t.Callable[ + [t.Iterable[Snapshot], SnapshotEvaluator], Scheduler + ], default_catalog: t.Optional[str], console: t.Optional[Console] = None, ): @@ -101,7 +100,9 @@ def evaluate( ) try: - plan_stages = stages.build_plan_stages(plan, self.state_sync, self.default_catalog) + plan_stages = stages.build_plan_stages( + plan, self.state_sync, self.default_catalog + ) self._evaluate_stages(plan_stages, plan) except Exception as e: analytics.collector.on_plan_apply_end(plan_id=plan.plan_id, error=e) @@ -124,7 +125,9 @@ def _evaluate_stages( handler = getattr(self, handler_name) handler(stage, plan) - def visit_before_all_stage(self, stage: stages.BeforeAllStage, plan: EvaluatablePlan) -> None: + def visit_before_all_stage( + self, stage: stages.BeforeAllStage, plan: EvaluatablePlan + ) -> None: execute_environment_statements( adapter=self.snapshot_evaluator.adapter, environment_statements=stage.statements, @@ -138,7 +141,9 @@ def visit_before_all_stage(self, stage: stages.BeforeAllStage, plan: Evaluatable selected_models=plan.selected_models, ) - def visit_after_all_stage(self, stage: stages.AfterAllStage, plan: EvaluatablePlan) -> None: + def visit_after_all_stage( + self, stage: stages.AfterAllStage, plan: EvaluatablePlan + ) -> None: execute_environment_statements( adapter=self.snapshot_evaluator.adapter, environment_statements=stage.statements, @@ -165,7 +170,9 @@ def visit_create_snapshot_records_stage( def visit_physical_layer_update_stage( self, stage: stages.PhysicalLayerUpdateStage, plan: EvaluatablePlan ) -> None: - skip_message = "" if plan.restatements else "\nSKIP: No physical layer updates to perform" + skip_message = ( + "" if plan.restatements else "\nSKIP: No physical layer updates to perform" + ) snapshots_to_create = stage.snapshots if not snapshots_to_create: @@ -203,7 +210,8 @@ def visit_physical_layer_update_stage( finally: if not progress_stopped: self.console.stop_creation_progress( - success=completion_status is not None and completion_status.is_success + success=completion_status is not None + and completion_status.is_success ) def visit_physical_layer_schema_creation_stage( @@ -216,15 +224,21 @@ def visit_physical_layer_schema_creation_stage( except Exception as ex: raise PlanError("Plan application failed.") from ex - def visit_backfill_stage(self, stage: stages.BackfillStage, plan: EvaluatablePlan) -> None: + def visit_backfill_stage( + self, stage: stages.BackfillStage, plan: EvaluatablePlan + ) -> None: if plan.empty_backfill: intervals_to_add = [] for snapshot in stage.all_snapshots.values(): - if not snapshot.evaluatable or not plan.is_selected_for_backfill(snapshot.name): + if not snapshot.evaluatable or not plan.is_selected_for_backfill( + snapshot.name + ): # Skip snapshots that are not evaluatable or not selected for backfill. continue intervals = [ - snapshot.inclusive_exclusive(plan.start, plan.end, strict=False, expand=False) + snapshot.inclusive_exclusive( + plan.start, plan.end, strict=False, expand=False + ) ] is_deployable = stage.deployability_index.is_deployable(snapshot) intervals_to_add.append( @@ -245,7 +259,9 @@ def visit_backfill_stage(self, stage: stages.BackfillStage, plan: EvaluatablePla self.console.log_success("SKIP: No model batches to execute") return - scheduler = self.create_scheduler(stage.all_snapshots.values(), self.snapshot_evaluator) + scheduler = self.create_scheduler( + stage.all_snapshots.values(), self.snapshot_evaluator + ) errors, _ = scheduler.run_merged_intervals( merged_intervals=stage.snapshot_to_intervals, deployability_index=stage.deployability_index, @@ -365,7 +381,9 @@ def visit_environment_record_update_stage( ) -> None: self.state_sync.promote( plan.environment, - no_gaps_snapshot_names=stage.no_gaps_snapshot_names if plan.no_gaps else set(), + no_gaps_snapshot_names=( + stage.no_gaps_snapshot_names if plan.no_gaps else set() + ), environment_statements=plan.environment_statements, ) @@ -383,7 +401,9 @@ def visit_migrate_schemas_stage( except NodeExecutionFailedError as ex: raise PlanError(str(ex.__cause__) if ex.__cause__ else str(ex)) - def visit_unpause_stage(self, stage: stages.UnpauseStage, plan: EvaluatablePlan) -> None: + def visit_unpause_stage( + self, stage: stages.UnpauseStage, plan: EvaluatablePlan + ) -> None: self.state_sync.unpause_snapshots(stage.promoted_snapshots, plan.end) def visit_virtual_layer_update_stage( @@ -409,10 +429,15 @@ def visit_virtual_layer_update_stage( ) if stage.demoted_environment_naming_info: self._demote_snapshots( - [stage.all_snapshots[s.snapshot_id] for s in stage.demoted_snapshots], + [ + stage.all_snapshots[s.snapshot_id] + for s in stage.demoted_snapshots + ], stage.demoted_environment_naming_info, deployability_index=stage.deployability_index, - on_complete=lambda s: self.console.update_promotion_progress(s, False), + on_complete=lambda s: self.console.update_promotion_progress( + s, False + ), snapshots=stage.all_snapshots, ) @@ -472,7 +497,9 @@ def _demote_snapshots( on_complete=on_complete, ) - def _update_intervals_for_new_snapshots(self, snapshots: t.Collection[Snapshot]) -> None: + def _update_intervals_for_new_snapshots( + self, snapshots: t.Collection[Snapshot] + ) -> None: snapshots_intervals: t.List[SnapshotIntervals] = [] for snapshot in snapshots: if snapshot.is_forward_only: diff --git a/sqlmesh/core/plan/explainer.py b/sqlmesh/core/plan/explainer.py index f0a1e44aff..e643b264b9 100644 --- a/sqlmesh/core/plan/explainer.py +++ b/sqlmesh/core/plan/explainer.py @@ -1,37 +1,34 @@ from __future__ import annotations import abc -import typing as t import logging -from dataclasses import dataclass +import typing as t from collections import defaultdict +from dataclasses import dataclass from rich.console import Console as RichConsole from rich.tree import Tree from sqlglot.dialects.dialect import DialectType + from sqlmesh.core import constants as c from sqlmesh.core.console import Console, TerminalConsole, get_console from sqlmesh.core.environment import EnvironmentNamingInfo +from sqlmesh.core.plan import stages from sqlmesh.core.plan.common import ( SnapshotIntervalClearRequest, - identify_restatement_intervals_across_snapshot_versions, -) + identify_restatement_intervals_across_snapshot_versions) from sqlmesh.core.plan.definition import EvaluatablePlan, SnapshotIntervals -from sqlmesh.core.plan import stages -from sqlmesh.core.plan.evaluator import ( - PlanEvaluator, -) +from sqlmesh.core.plan.evaluator import PlanEvaluator +from sqlmesh.core.snapshot.definition import (SnapshotIdAndVersion, + SnapshotInfoMixin, + model_display_name) from sqlmesh.core.state_sync import StateReader -from sqlmesh.core.snapshot.definition import ( - SnapshotInfoMixin, - SnapshotIdAndVersion, - model_display_name, -) -from sqlmesh.utils import Verbosity, rich as srich, to_snake_case +from sqlmesh.utils import Verbosity +from sqlmesh.utils import rich as srich +from sqlmesh.utils import to_snake_case from sqlmesh.utils.date import to_ts from sqlmesh.utils.errors import SQLMeshError - logger = logging.getLogger(__name__) @@ -51,16 +48,22 @@ def evaluate( plan: EvaluatablePlan, circuit_breaker: t.Optional[t.Callable[[], bool]] = None, ) -> None: - plan_stages = stages.build_plan_stages(plan, self.state_reader, self.default_catalog) + plan_stages = stages.build_plan_stages( + plan, self.state_reader, self.default_catalog + ) explainer_console = _get_explainer_console( self.console, plan.environment, self.default_catalog ) # add extra metadata that's only needed at this point for better --explain output plan_stages = [ - ExplainableRestatementStage.from_restatement_stage(stage, self.state_reader, plan) - if isinstance(stage, stages.RestatementStage) - else stage + ( + ExplainableRestatementStage.from_restatement_stage( + stage, self.state_reader, plan + ) + if isinstance(stage, stages.RestatementStage) + else stage + ) for stage in plan_stages ] @@ -90,17 +93,23 @@ def from_restatement_stage( state_reader: StateReader, plan: EvaluatablePlan, ) -> ExplainableRestatementStage: - all_restatement_intervals = identify_restatement_intervals_across_snapshot_versions( - state_reader=state_reader, - prod_restatements=plan.restatements, - disable_restatement_models=plan.disabled_restatement_models, - loaded_snapshots={s.snapshot_id: s for s in stage.all_snapshots.values()}, + all_restatement_intervals = ( + identify_restatement_intervals_across_snapshot_versions( + state_reader=state_reader, + prod_restatements=plan.restatements, + disable_restatement_models=plan.disabled_restatement_models, + loaded_snapshots={ + s.snapshot_id: s for s in stage.all_snapshots.values() + }, + ) ) # Group the interval clear requests by snapshot name to make them easier to write to the console snapshot_intervals_to_clear = defaultdict(list) for clear_request in all_restatement_intervals.values(): - snapshot_intervals_to_clear[clear_request.snapshot.name].append(clear_request) + snapshot_intervals_to_clear[clear_request.snapshot.name].append( + clear_request + ) return cls( snapshot_intervals_to_clear=snapshot_intervals_to_clear, @@ -145,9 +154,13 @@ def visit_before_all_stage(self, stage: stages.BeforeAllStage) -> Tree: def visit_after_all_stage(self, stage: stages.AfterAllStage) -> Tree: return Tree("[bold]Execute after all statements[/bold]") - def visit_physical_layer_update_stage(self, stage: stages.PhysicalLayerUpdateStage) -> Tree: + def visit_physical_layer_update_stage( + self, stage: stages.PhysicalLayerUpdateStage + ) -> Tree: snapshots = [ - s for s in stage.snapshots if s.snapshot_id in stage.snapshots_with_missing_intervals + s + for s in stage.snapshots + if s.snapshot_id in stage.snapshots_with_missing_intervals ] if not snapshots: return Tree("[bold]SKIP: No physical layer updates to perform[/bold]") @@ -174,7 +187,9 @@ def visit_physical_layer_update_stage(self, stage: stages.PhysicalLayerUpdateSta if snapshot.is_view: create_tree = Tree("Create view if it doesn't exist") elif ( - snapshot.is_forward_only and snapshot.previous_versions and not snapshot.is_managed + snapshot.is_forward_only + and snapshot.previous_versions + and not snapshot.is_managed ): prod_table = snapshot.table_name(True) create_tree = Tree( @@ -184,7 +199,9 @@ def visit_physical_layer_update_stage(self, stage: stages.PhysicalLayerUpdateSta create_tree = Tree("Create table if it doesn't exist") if not is_deployable: - create_tree.add("[orange1]preview[/orange1]: data will NOT be reused in production") + create_tree.add( + "[orange1]preview[/orange1]: data will NOT be reused in production" + ) model_tree.add(create_tree) if snapshot.is_model and snapshot.model.post_statements: @@ -200,7 +217,9 @@ def visit_audit_only_run_stage(self, stage: stages.AuditOnlyRunStage) -> Tree: tree.add(display_name) return tree - def visit_explainable_restatement_stage(self, stage: ExplainableRestatementStage) -> Tree: + def visit_explainable_restatement_stage( + self, stage: ExplainableRestatementStage + ) -> Tree: return self.visit_restatement_stage(stage) def visit_restatement_stage( @@ -216,7 +235,10 @@ def visit_restatement_stage( ): for name, clear_requests in snapshot_intervals.items(): display_name = model_display_name( - name, self.environment_naming_info, self.default_catalog, self.dialect + name, + self.environment_naming_info, + self.default_catalog, + self.dialect, ) interval_start = min(cr.interval[0] for cr in clear_requests) interval_end = max(cr.interval[1] for cr in clear_requests) @@ -224,10 +246,16 @@ def visit_restatement_stage( if not interval_start or not interval_end: continue - node = tree.add(f"{display_name} [{to_ts(interval_start)} - {to_ts(interval_end)}]") + node = tree.add( + f"{display_name} [{to_ts(interval_start)} - {to_ts(interval_end)}]" + ) all_environment_names = sorted( - set(env_name for cr in clear_requests for env_name in cr.environment_names) + set( + env_name + for cr in clear_requests + for env_name in cr.environment_names + ) ) node.add("in environments: " + ", ".join(all_environment_names)) @@ -300,7 +328,9 @@ def visit_migrate_schemas_stage(self, stage: stages.MigrateSchemasStage) -> Tree tree.add(f"{display_name} -> {table_name}") return tree - def visit_virtual_layer_update_stage(self, stage: stages.VirtualLayerUpdateStage) -> Tree: + def visit_virtual_layer_update_stage( + self, stage: stages.VirtualLayerUpdateStage + ) -> Tree: tree = Tree( f"[bold]Update the virtual layer for environment '{self.environment_naming_info.name}'[/bold]" ) @@ -309,14 +339,18 @@ def visit_virtual_layer_update_stage(self, stage: stages.VirtualLayerUpdateStage ) for snapshot in stage.promoted_snapshots: display_name = self._display_name(snapshot) - table_name = snapshot.table_name(stage.deployability_index.is_representative(snapshot)) + table_name = snapshot.table_name( + stage.deployability_index.is_representative(snapshot) + ) promote_tree.add(f"{display_name} -> {table_name}") demote_tree = Tree( "[bold]Delete views in the virtual layer for models that were removed[/bold]" ) for snapshot in stage.demoted_snapshots: - display_name = self._display_name(snapshot, stage.demoted_environment_naming_info) + display_name = self._display_name( + snapshot, stage.demoted_environment_naming_info + ) demote_tree.add(display_name) if stage.promoted_snapshots: @@ -349,10 +383,13 @@ def _display_name( environment_naming_info: t.Optional[EnvironmentNamingInfo] = None, ) -> str: return snapshot.display_name( - environment_naming_info=environment_naming_info or self.environment_naming_info, - default_catalog=self.default_catalog - if self.verbosity < Verbosity.VERY_VERBOSE - else None, + environment_naming_info=environment_naming_info + or self.environment_naming_info, + default_catalog=( + self.default_catalog + if self.verbosity < Verbosity.VERY_VERBOSE + else None + ), dialect=self.dialect, ) diff --git a/sqlmesh/core/plan/stages.py b/sqlmesh/core/plan/stages.py index 729e1705b4..b080b7c8e9 100644 --- a/sqlmesh/core/plan/stages.py +++ b/sqlmesh/core/plan/stages.py @@ -1,19 +1,17 @@ import typing as t - from dataclasses import dataclass + from sqlmesh.core import constants as c -from sqlmesh.core.environment import EnvironmentStatements, EnvironmentNamingInfo, Environment +from sqlmesh.core.environment import (Environment, EnvironmentNamingInfo, + EnvironmentStatements) from sqlmesh.core.plan.common import should_force_rebuild from sqlmesh.core.plan.definition import EvaluatablePlan +from sqlmesh.core.scheduler import (SnapshotToIntervals, + merged_missing_intervals) +from sqlmesh.core.snapshot.definition import (DeployabilityIndex, Snapshot, + SnapshotId, SnapshotTableInfo, + snapshots_to_dag) from sqlmesh.core.state_sync import StateReader -from sqlmesh.core.scheduler import merged_missing_intervals, SnapshotToIntervals -from sqlmesh.core.snapshot.definition import ( - DeployabilityIndex, - Snapshot, - SnapshotTableInfo, - SnapshotId, - snapshots_to_dag, -) from sqlmesh.utils.errors import PlanError @@ -253,7 +251,9 @@ def build(self, plan: EvaluatablePlan) -> t.List[PlanStage]: dag = snapshots_to_dag(snapshots.values()) all_selected_for_backfill_snapshots = { - s.snapshot_id for s in snapshots.values() if plan.is_selected_for_backfill(s.name) + s.snapshot_id + for s in snapshots.values() + if plan.is_selected_for_backfill(s.name) } existing_environment = self.state_reader.get_environment(plan.environment.name) @@ -271,11 +271,15 @@ def build(self, plan: EvaluatablePlan) -> t.List[PlanStage]: if (deployability_index.is_representative(s) or s.is_seed) and plan.is_selected_for_backfill(s.name) } - after_promote_snapshots = all_selected_for_backfill_snapshots - before_promote_snapshots + after_promote_snapshots = ( + all_selected_for_backfill_snapshots - before_promote_snapshots + ) deployability_index = DeployabilityIndex.all_deployable() snapshot_ids_with_schema_migration = [ - s.snapshot_id for s in snapshots.values() if s.requires_schema_migration_in_prod + s.snapshot_id + for s in snapshots.values() + if s.requires_schema_migration_in_prod ] # Include all upstream dependencies of snapshots that require schema migration to make sure # the upstream tables are created before the schema updates are applied @@ -289,7 +293,9 @@ def build(self, plan: EvaluatablePlan) -> t.List[PlanStage]: plan, snapshots_by_name, deployability_index ) needs_backfill = ( - not plan.empty_backfill and not plan.skip_backfill and bool(snapshots_to_intervals) + not plan.empty_backfill + and not plan.skip_backfill + and bool(snapshots_to_intervals) ) missing_intervals_before_promote: SnapshotToIntervals = {} missing_intervals_after_promote: SnapshotToIntervals = {} @@ -317,7 +323,8 @@ def build(self, plan: EvaluatablePlan) -> t.List[PlanStage]: if snapshots_to_create: stages.append( PhysicalLayerSchemaCreationStage( - snapshots=snapshots_to_create, deployability_index=deployability_index + snapshots=snapshots_to_create, + deployability_index=deployability_index, ) ) if not needs_backfill: @@ -333,7 +340,9 @@ def build(self, plan: EvaluatablePlan) -> t.List[PlanStage]: audit_only_snapshots = self._get_audit_only_snapshots(new_snapshots) if audit_only_snapshots: - stages.append(AuditOnlyRunStage(snapshots=list(audit_only_snapshots.values()))) + stages.append( + AuditOnlyRunStage(snapshots=list(audit_only_snapshots.values())) + ) if missing_intervals_before_promote: stages.append( @@ -383,7 +392,11 @@ def build(self, plan: EvaluatablePlan) -> t.List[PlanStage]: ) ) - if not plan.is_dev and not plan.ensure_finalized_snapshots and promoted_snapshots: + if ( + not plan.is_dev + and not plan.ensure_finalized_snapshots + and promoted_snapshots + ): # Only unpause at this point if we don't have to use the finalized snapshots # for subsequent plan applications. Otherwise, unpause right before updating # the virtual layer. @@ -505,11 +518,21 @@ def _get_virtual_layer_update_stage( ) -> t.Optional[VirtualLayerUpdateStage]: def _should_update_virtual_layer(snapshot: SnapshotTableInfo) -> bool: # Skip virtual layer update for snapshots with virtual environment support disabled - virtual_environment_enabled = is_dev or snapshot.virtual_environment_mode.is_full - return snapshot.is_model and not snapshot.is_symbolic and virtual_environment_enabled + virtual_environment_enabled = ( + is_dev or snapshot.virtual_environment_mode.is_full + ) + return ( + snapshot.is_model + and not snapshot.is_symbolic + and virtual_environment_enabled + ) - promoted_snapshots = {s for s in promoted_snapshots if _should_update_virtual_layer(s)} - demoted_snapshots = {s for s in demoted_snapshots if _should_update_virtual_layer(s)} + promoted_snapshots = { + s for s in promoted_snapshots if _should_update_virtual_layer(s) + } + demoted_snapshots = { + s for s in demoted_snapshots if _should_update_virtual_layer(s) + } if not promoted_snapshots and not demoted_snapshots: return None @@ -524,11 +547,14 @@ def _should_update_virtual_layer(snapshot: SnapshotTableInfo) -> bool: def _get_promoted_demoted_snapshots( self, plan: EvaluatablePlan, existing_environment: t.Optional[Environment] ) -> t.Tuple[ - t.Set[SnapshotTableInfo], t.Set[SnapshotTableInfo], t.Optional[EnvironmentNamingInfo] + t.Set[SnapshotTableInfo], + t.Set[SnapshotTableInfo], + t.Optional[EnvironmentNamingInfo], ]: if existing_environment: new_table_infos = { - table_info.name: table_info for table_info in plan.environment.promoted_snapshots + table_info.name: table_info + for table_info in plan.environment.promoted_snapshots } existing_table_infos = { table_info.name: table_info @@ -541,9 +567,9 @@ def _get_promoted_demoted_snapshots( and existing_table_info.qualified_view_name.for_environment( existing_environment.naming_info ) - != new_table_infos[existing_table_info.name].qualified_view_name.for_environment( - plan.environment.naming_info - ) + != new_table_infos[ + existing_table_info.name + ].qualified_view_name.for_environment(plan.environment.naming_info) } missing_model_names = set(existing_table_infos) - { s.name for s in plan.environment.promoted_snapshots @@ -555,11 +581,15 @@ def _get_promoted_demoted_snapshots( demoted_snapshots = set() promoted_snapshots = set(plan.environment.promoted_snapshots) - if existing_environment and plan.environment.can_partially_promote(existing_environment): + if existing_environment and plan.environment.can_partially_promote( + existing_environment + ): promoted_snapshots -= set(existing_environment.promoted_snapshots) demoted_environment_naming_info = ( - existing_environment.naming_info if demoted_snapshots and existing_environment else None + existing_environment.naming_info + if demoted_snapshots and existing_environment + else None ) return ( @@ -607,10 +637,13 @@ def _get_audit_only_snapshots( # Bulk load all the previous snapshots previous_snapshot_ids = [ - s.previous_version.snapshot_id(s.name) for s in metadata_snapshots if s.previous_version + s.previous_version.snapshot_id(s.name) + for s in metadata_snapshots + if s.previous_version ] previous_snapshots = { - s.name: s for s in self.state_reader.get_snapshots(previous_snapshot_ids).values() + s.name: s + for s in self.state_reader.get_snapshots(previous_snapshot_ids).values() } # Check if any of the snapshots have modifications to the audits field by comparing the hashes diff --git a/sqlmesh/core/reference.py b/sqlmesh/core/reference.py index 9e93ce7b38..ef99d35282 100644 --- a/sqlmesh/core/reference.py +++ b/sqlmesh/core/reference.py @@ -76,7 +76,9 @@ def add_model(self, model: Model) -> None: self._ref_models.setdefault(ref.name, set()) self._ref_models[ref.name].add(model.name) - def models_for_column(self, source: str, column: str, max_depth: int = 3) -> t.List[str]: + def models_for_column( + self, source: str, column: str, max_depth: int = 3 + ) -> t.List[str]: """Find all the models with a column that join to a source within max_depth. Args: @@ -99,7 +101,9 @@ def models_for_column(self, source: str, column: str, max_depth: int = 3) -> t.L return sorted(models) - def find_path(self, source: str, target: str, max_depth: int = 3) -> t.List[Reference]: + def find_path( + self, source: str, target: str, max_depth: int = 3 + ) -> t.List[Reference]: """Find a path from source model to target model with max depth. Args: diff --git a/sqlmesh/core/renderer.py b/sqlmesh/core/renderer.py index 9f403cbcb4..93efc481da 100644 --- a/sqlmesh/core/renderer.py +++ b/sqlmesh/core/renderer.py @@ -6,7 +6,7 @@ from functools import partial from pathlib import Path -from sqlglot import exp, Dialect +from sqlglot import Dialect, exp from sqlglot.errors import SqlglotError from sqlglot.helper import ensure_list from sqlglot.optimizer.annotate_types import annotate_types @@ -16,20 +16,10 @@ from sqlmesh.core import constants as c from sqlmesh.core import dialect as d from sqlmesh.core.macros import MacroEvaluator, RuntimeStage -from sqlmesh.utils.date import ( - TimeLike, - date_dict, - make_inclusive, - to_datetime, - make_ts_exclusive, - to_tstz, -) -from sqlmesh.utils.errors import ( - ConfigError, - ParsetimeAdapterCallError, - SQLMeshError, - raise_config_error, -) +from sqlmesh.utils.date import (TimeLike, date_dict, make_inclusive, + make_ts_exclusive, to_datetime, to_tstz) +from sqlmesh.utils.errors import (ConfigError, ParsetimeAdapterCallError, + SQLMeshError, raise_config_error) from sqlmesh.utils.jinja import JinjaMacroRegistry, extract_error_details from sqlmesh.utils.metaprogramming import Executable, prepare_env @@ -129,7 +119,9 @@ def _render( ) views.append( snapshot.display_name( - environment_naming_info, self._default_catalog, self._dialect + environment_naming_info, + self._default_catalog, + self._dialect, ) ) if schemas: @@ -139,7 +131,9 @@ def _render( this_model = kwargs.pop("this_model", None) - this_snapshot = (snapshots or {}).get(self._model_fqn) if self._model_fqn else None + this_snapshot = ( + (snapshots or {}).get(self._model_fqn) if self._model_fqn else None + ) if not this_model and self._model_fqn: this_model = self._resolve_table( self._model_fqn, @@ -211,7 +205,9 @@ def _resolve_table(table: str | exp.Table) -> str: jinja_env_kwargs = { **{ **render_kwargs, - **_prepare_python_env_for_jinja(macro_evaluator, self._python_env), + **_prepare_python_env_for_jinja( + macro_evaluator, self._python_env + ), **variables, }, "snapshots": snapshots or {}, @@ -238,13 +234,19 @@ def _resolve_table(table: str | exp.Table) -> str: if ref.event_time_filter: ref.event_time_filter["start"] = render_kwargs["start_tstz"] ref.event_time_filter["end"] = to_tstz( - make_ts_exclusive(render_kwargs["end_tstz"], dialect=self._dialect) + make_ts_exclusive( + render_kwargs["end_tstz"], dialect=self._dialect + ) ) - jinja_env = self._jinja_macro_registry.build_environment(**jinja_env_kwargs) + jinja_env = self._jinja_macro_registry.build_environment( + **jinja_env_kwargs + ) expressions = [] - rendered_expression = jinja_env.from_string(self._expression.name).render() + rendered_expression = jinja_env.from_string( + self._expression.name + ).render() logger.debug( f"Rendered Jinja expression for model '{self._model_fqn}' at '{self._path}': '{rendered_expression}'" ) @@ -252,7 +254,8 @@ def _resolve_table(table: str | exp.Table) -> str: raise except Exception as ex: raise ConfigError( - f"Could not render jinja for '{self._path}'.\n" + extract_error_details(ex) + f"Could not render jinja for '{self._path}'.\n" + + extract_error_details(ex) ) from ex if rendered_expression.strip(): @@ -263,7 +266,9 @@ def _resolve_table(table: str | exp.Table) -> str: if tokens: try: expressions = [ - e for e in dialect.parser().parse(tokens, rendered_expression) if e + e + for e in dialect.parser().parse(tokens, rendered_expression) + if e ] if not expressions: @@ -287,7 +292,9 @@ def _resolve_table(table: str | exp.Table) -> str: for expression in expressions: try: - transformed_expressions = ensure_list(macro_evaluator.transform(expression)) + transformed_expressions = ensure_list( + macro_evaluator.transform(expression) + ) except Exception as ex: raise_config_error( f"Failed to resolve macros for\n\n{expression.sql(dialect=self._dialect, pretty=True)}\n\n{ex}\n", @@ -298,12 +305,16 @@ def _resolve_table(table: str | exp.Table) -> str: with self._normalize_and_quote(expression) as expression: if hasattr(expression, "selects"): for select in expression.selects: - if not isinstance(select, exp.Alias) and select.output_name not in ( + if not isinstance( + select, exp.Alias + ) and select.output_name not in ( "*", "", ): alias = exp.alias_( - select, select.output_name, quoted=self._quote_identifiers + select, + select.output_name, + quoted=self._quote_identifiers, ) comments = alias.this.comments if comments: @@ -316,7 +327,9 @@ def _resolve_table(table: str | exp.Table) -> str: # We dont cache here if columns_to_type was called in a macro. # This allows the model's query to be re-rendered so that the # MacroEvaluator can resolve columns_to_types calls and provide true schemas. - if should_cache and (not self.schema.empty or not macro_evaluator.columns_to_types_called): + if should_cache and ( + not self.schema.empty or not macro_evaluator.columns_to_types_called + ): self._cache = resolved_expressions return resolved_expressions @@ -331,9 +344,14 @@ def _resolve_table( deployability_index: t.Optional[DeployabilityIndex] = None, ) -> exp.Table: table = exp.replace_tables( - t.cast(exp.Table, exp.maybe_parse(table_name, into=exp.Table, dialect=self._dialect)), + t.cast( + exp.Table, + exp.maybe_parse(table_name, into=exp.Table, dialect=self._dialect), + ), { - **self._to_table_mapping((snapshots or {}).values(), deployability_index), + **self._to_table_mapping( + (snapshots or {}).values(), deployability_index + ), **(table_mapping or {}), }, dialect=self._dialect, @@ -342,7 +360,9 @@ def _resolve_table( # We quote the table here to mimic the behavior of _resolve_tables, otherwise we may end # up normalizing twice, because _to_table_mapping returns the mapped names unquoted. return ( - d.quote_identifiers(table, dialect=self._dialect) if self._quote_identifiers else table + d.quote_identifiers(table, dialect=self._dialect) + if self._quote_identifiers + else table ) def _resolve_tables( @@ -405,7 +425,9 @@ def _expand(node: exp.Expr) -> exp.Expr: alias=node.alias or model.view_name, copy=False, ) - logger.warning("Failed to expand the nested model '%s'", name) + logger.warning( + "Failed to expand the nested model '%s'", name + ) return node expression = expression.transform(_expand, copy=False) # type: ignore @@ -421,7 +443,10 @@ def _expand(node: exp.Expr) -> exp.Expr: def _normalize_and_quote(self, query: E) -> t.Iterator[E]: if self._normalize_identifiers: with d.normalize_and_quote( - query, self._dialect, self._default_catalog, quote=self._quote_identifiers + query, + self._dialect, + self._default_catalog, + quote=self._quote_identifiers, ) as query: yield query else: @@ -431,7 +456,9 @@ def _should_cache(self, runtime_stage: RuntimeStage, *args: t.Any) -> bool: return runtime_stage == RuntimeStage.LOADING and not any(args) def _to_table_mapping( - self, snapshots: t.Iterable[Snapshot], deployability_index: t.Optional[DeployabilityIndex] + self, + snapshots: t.Iterable[Snapshot], + deployability_index: t.Optional[DeployabilityIndex], ) -> t.Dict[str, str]: from sqlmesh.core.snapshot import to_table_mapping @@ -509,7 +536,9 @@ def render_statements( f"Rendering `{expression.sql(dialect=dialect)}` did not return an expression" ) else: - rendered_statements.extend(expr.sql(dialect=dialect) for expr in rendered) + rendered_statements.extend( + expr.sql(dialect=dialect) for expr in rendered + ) return rendered_statements @@ -588,7 +617,9 @@ def render( if isinstance(self._expression, d.JinjaQuery): return None - raise ConfigError(f"Failed to render query at '{self._path}':\n{self._expression}") + raise ConfigError( + f"Failed to render query at '{self._path}':\n{self._expression}" + ) if len(expressions) > 1: raise ConfigError(f"Too many statements in query:\n{self._expression}") @@ -599,7 +630,8 @@ def render( return None if not isinstance(query, exp.Query): raise_config_error( - f"Model query needs to be a SELECT or a UNION, got {query}.", self._path + f"Model query needs to be a SELECT or a UNION, got {query}.", + self._path, ) raise @@ -646,9 +678,7 @@ def update_cache( def _optimize_query(self, query: exp.Query, all_deps: t.Set[str]) -> exp.Query: from sqlmesh.core.linter.rules.builtin import ( - AmbiguousOrInvalidColumn, - InvalidSelectStarExpansion, - ) + AmbiguousOrInvalidColumn, InvalidSelectStarExpansion) # We don't want to normalize names in the schema because that's handled by the optimizer original = query @@ -661,7 +691,11 @@ def _optimize_query(self, query: exp.Query, all_deps: t.Set[str]) -> exp.Query: should_optimize = False missing_deps.add(dep) - if self._model_fqn and not should_optimize and any(s.is_star for s in query.selects): + if ( + self._model_fqn + and not should_optimize + and any(s.is_star for s in query.selects) + ): deps = ", ".join(f"'{dep}'" for dep in sorted(missing_deps)) self._violated_rules[InvalidSelectStarExpansion] = deps diff --git a/sqlmesh/core/scheduler.py b/sqlmesh/core/scheduler.py index 5eb0ff40ff..9045b4962d 100644 --- a/sqlmesh/core/scheduler.py +++ b/sqlmesh/core/scheduler.py @@ -1,56 +1,40 @@ from __future__ import annotations -from dataclasses import dataclass + import abc import logging -import typing as t import time +import typing as t +from dataclasses import dataclass from datetime import datetime + from sqlglot import exp + from sqlmesh.core import constants as c from sqlmesh.core.console import Console, get_console -from sqlmesh.core.environment import EnvironmentNamingInfo, execute_environment_statements +from sqlmesh.core.environment import (EnvironmentNamingInfo, + execute_environment_statements) from sqlmesh.core.macros import RuntimeStage from sqlmesh.core.model.definition import AuditResult from sqlmesh.core.node import IntervalUnit -from sqlmesh.core.notification_target import ( - NotificationEvent, - NotificationTargetManager, -) -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Snapshot, - SnapshotId, - SnapshotIdBatch, - SnapshotEvaluator, - apply_auto_restatements, - earliest_start_date, - missing_intervals, - merge_intervals, - snapshots_to_dag, - Intervals, -) -from sqlmesh.core.snapshot.definition import check_ready_intervals -from sqlmesh.core.snapshot.definition import ( - Interval, - expand_range, - parent_snapshots_by_name, -) +from sqlmesh.core.notification_target import (NotificationEvent, + NotificationTargetManager) +from sqlmesh.core.snapshot import (DeployabilityIndex, Intervals, Snapshot, + SnapshotEvaluator, SnapshotId, + SnapshotIdBatch, apply_auto_restatements, + earliest_start_date, merge_intervals, + missing_intervals, snapshots_to_dag) +from sqlmesh.core.snapshot.definition import (Interval, check_ready_intervals, + expand_range, + parent_snapshots_by_name) from sqlmesh.core.state_sync import StateSync from sqlmesh.utils import CompletionStatus -from sqlmesh.utils.concurrency import concurrent_apply_to_dag, NodeExecutionFailedError +from sqlmesh.utils.concurrency import (NodeExecutionFailedError, + concurrent_apply_to_dag) from sqlmesh.utils.dag import DAG -from sqlmesh.utils.date import ( - TimeLike, - now_timestamp, - validate_date_range, -) -from sqlmesh.utils.errors import ( - AuditError, - NodeAuditsErrors, - CircuitBreakerError, - SQLMeshError, - SignalEvalError, -) +from sqlmesh.utils.date import TimeLike, now_timestamp, validate_date_range +from sqlmesh.utils.errors import (AuditError, CircuitBreakerError, + NodeAuditsErrors, SignalEvalError, + SQLMeshError) if t.TYPE_CHECKING: from sqlmesh.core.context import ExecutionContext @@ -78,7 +62,12 @@ class EvaluateNode(SchedulingUnit): def __lt__(self, other: SchedulingUnit) -> bool: if not isinstance(other, EvaluateNode): return super().__lt__(other) - return (self.__class__.__name__, self.snapshot_name, self.interval, self.batch_index) < ( + return ( + self.__class__.__name__, + self.snapshot_name, + self.interval, + self.batch_index, + ) < ( other.__class__.__name__, other.snapshot_name, other.interval, @@ -125,8 +114,12 @@ def __init__( ): self.state_sync = state_sync self.snapshots = {s.snapshot_id: s for s in snapshots} - self.snapshots_by_name = {snapshot.name: snapshot for snapshot in self.snapshots.values()} - self.snapshot_per_version = _resolve_one_snapshot_per_version(self.snapshots.values()) + self.snapshots_by_name = { + snapshot.name: snapshot for snapshot in self.snapshots.values() + } + self.snapshot_per_version = _resolve_one_snapshot_per_version( + self.snapshots.values() + ) self.default_catalog = default_catalog self.snapshot_evaluator = snapshot_evaluator self.max_workers = max_workers @@ -184,7 +177,9 @@ def merged_missing_intervals( # to correctly infer start dates. if selected_snapshots is not None: snapshots_to_intervals = { - s: i for s, i in snapshots_to_intervals.items() if s.name in selected_snapshots + s: i + for s, i in snapshots_to_intervals.items() + if s.name in selected_snapshots } return snapshots_to_intervals @@ -252,7 +247,11 @@ def evaluate( ) self.state_sync.add_interval( - snapshot, start, end, is_dev=not is_deployable, last_altered_ts=now_timestamp() + snapshot, + start, + end, + is_dev=not is_deployable, + last_altered_ts=now_timestamp(), ) return audit_results @@ -347,7 +346,9 @@ def batch_intervals( [ i for interval in intervals - for i in _expand_range_as_interval(*interval, snapshot.node.interval_unit) + for i in _expand_range_as_interval( + *interval, snapshot.node.interval_unit + ) ], ) for snapshot, intervals in merged_intervals.items() @@ -401,7 +402,9 @@ def batch_intervals( next_batch: t.List[t.Tuple[int, int]] = [] for interval in interval_diff( - intervals, merge_intervals(unready), uninterrupted=snapshot.depends_on_past + intervals, + merge_intervals(unready), + uninterrupted=snapshot.depends_on_past, ): if (batch_size and len(next_batch) >= batch_size) or ( next_batch and interval[0] != next_batch[-1][-1] @@ -436,7 +439,9 @@ def run_merged_intervals( audit_only: bool = False, auto_restatement_triggers: t.Dict[SnapshotId, t.List[SnapshotId]] = {}, is_restatement: bool = False, - ) -> t.Tuple[t.List[NodeExecutionFailedError[SchedulingUnit]], t.List[SchedulingUnit]]: + ) -> t.Tuple[ + t.List[NodeExecutionFailedError[SchedulingUnit]], t.List[SchedulingUnit] + ]: """Runs precomputed batches of missing intervals. Args: @@ -456,7 +461,9 @@ def run_merged_intervals( """ execution_time = execution_time or now_timestamp() - selected_snapshots = [self.snapshots[sid] for sid in (selected_snapshot_ids or set())] + selected_snapshots = [ + self.snapshots[sid] for sid in (selected_snapshot_ids or set()) + ] if not selected_snapshots: selected_snapshots = list(merged_intervals) @@ -514,7 +521,9 @@ def run_merged_intervals( } dag = self._dag( - batched_intervals, snapshot_dag=snapshot_dag, snapshots_to_create=snapshots_to_create + batched_intervals, + snapshot_dag=snapshot_dag, + snapshots_to_create=snapshots_to_create, ) def run_node(node: SchedulingUnit) -> None: @@ -549,7 +558,8 @@ def run_node(node: SchedulingUnit) -> None: else: # If batch_index > 0, then the target table must exist since the first batch would have created it target_table_exists = ( - snapshot.snapshot_id not in snapshots_to_create or node.batch_index > 0 + snapshot.snapshot_id not in snapshots_to_create + or node.batch_index > 0 ) audit_results = self.evaluate( snapshot=snapshot, @@ -568,10 +578,17 @@ def run_node(node: SchedulingUnit) -> None: evaluation_duration_ms = now_timestamp() - execution_start_ts finally: num_audits = len(audit_results) - num_audits_failed = sum(1 for result in audit_results if result.count) + num_audits_failed = sum( + 1 for result in audit_results if result.count + ) - execution_stats = self.snapshot_evaluator.execution_tracker.get_execution_stats( - SnapshotIdBatch(snapshot_id=snapshot.snapshot_id, batch_id=node.batch_index) + execution_stats = ( + self.snapshot_evaluator.execution_tracker.get_execution_stats( + SnapshotIdBatch( + snapshot_id=snapshot.snapshot_id, + batch_id=node.batch_index, + ) + ) ) self.console.update_snapshot_evaluation_progress( @@ -606,7 +623,9 @@ def run_node(node: SchedulingUnit) -> None: self.console.stop_evaluation_progress(success=not errors) skipped_snapshots = { - i.snapshot_name for i in skipped_intervals if isinstance(i, EvaluateNode) + i.snapshot_name + for i in skipped_intervals + if isinstance(i, EvaluateNode) } self.console.log_skipped_models(skipped_snapshots) for skipped in skipped_snapshots: @@ -701,7 +720,9 @@ def _dag( snapshots_to_create.remove(snapshot.snapshot_id) for i, interval in enumerate(intervals): - node = EvaluateNode(snapshot_name=snapshot.name, interval=interval, batch_index=i) + node = EvaluateNode( + snapshot_name=snapshot.name, interval=interval, batch_index=i + ) if create_node: dag.add(node, [create_node]) @@ -835,12 +856,15 @@ def _run_or_audit( all_auto_restatement_triggers: t.Dict[SnapshotId, t.List[SnapshotId]] = {} if auto_restatement_enabled: - auto_restated_intervals, all_auto_restatement_triggers = apply_auto_restatements( - self.snapshots, execution_time + auto_restated_intervals, all_auto_restatement_triggers = ( + apply_auto_restatements(self.snapshots, execution_time) ) self.state_sync.add_snapshots_intervals(auto_restated_intervals) self.state_sync.update_auto_restatements( - {s.name_version: s.next_auto_restatement_ts for s in self.snapshots.values()} + { + s.name_version: s.next_auto_restatement_ts + for s in self.snapshots.values() + } ) merged_intervals = self.merged_missing_intervals( @@ -860,7 +884,9 @@ def _run_or_audit( auto_restatement_triggers: t.Dict[SnapshotId, t.List[SnapshotId]] = {} if all_auto_restatement_triggers: - merged_intervals_snapshots = {snapshot.snapshot_id for snapshot in merged_intervals} + merged_intervals_snapshots = { + snapshot.snapshot_id for snapshot in merged_intervals + } auto_restatement_triggers = { s_id: all_auto_restatement_triggers.get(s_id, []) for s_id in merged_intervals_snapshots @@ -920,7 +946,9 @@ def _audit_snapshot( query=t.cast(exp.Query, audit_result.query), adapter_dialect=self.snapshot_evaluator.adapter.dialect, ) - self.notification_target_manager.notify(NotificationEvent.AUDIT_FAILURE, error) + self.notification_target_manager.notify( + NotificationEvent.AUDIT_FAILURE, error + ) if is_deployable and snapshot.node.owner: self.notification_target_manager.notify_user( NotificationEvent.AUDIT_FAILURE, snapshot.node.owner, error @@ -980,7 +1008,9 @@ def _check_ready_intervals( environment_naming_info or EnvironmentNamingInfo(), ) - for signal_idx, (signal_name, kwargs) in enumerate(signals.signals_to_kwargs.items()): + for signal_idx, (signal_name, kwargs) in enumerate( + signals.signals_to_kwargs.items() + ): # Capture intervals before signal check for display intervals_to_check = merge_intervals(intervals) @@ -1182,7 +1212,8 @@ def _resolve_one_snapshot_per_version( else: prev_snapshot = snapshot_per_version[key] if snapshot.unpaused_ts and ( - not prev_snapshot.unpaused_ts or snapshot.created_ts > prev_snapshot.created_ts + not prev_snapshot.unpaused_ts + or snapshot.created_ts > prev_snapshot.created_ts ): snapshot_per_version[key] = snapshot diff --git a/sqlmesh/core/schema_diff.py b/sqlmesh/core/schema_diff.py index ecf38b18a8..8e2eea4e9f 100644 --- a/sqlmesh/core/schema_diff.py +++ b/sqlmesh/core/schema_diff.py @@ -3,8 +3,8 @@ import abc import logging import typing as t -from dataclasses import dataclass from collections import defaultdict +from dataclasses import dataclass from enum import Enum from pydantic import Field @@ -12,8 +12,8 @@ from sqlglot.helper import ensure_list, seq_get from sqlmesh.utils import columns_to_types_to_struct -from sqlmesh.utils.pydantic import PydanticModel from sqlmesh.utils.errors import SQLMeshError +from sqlmesh.utils.pydantic import PydanticModel if t.TYPE_CHECKING: from sqlmesh.core._typing import TableName @@ -83,7 +83,9 @@ class TableAlterTypedColumnOperation(TableAlterColumnOperation, abc.ABC): @property def column_def(self) -> exp.ColumnDef: if not self.column_type: - raise SQLMeshError("Tried to access column type when it shouldn't be needed") + raise SQLMeshError( + "Tried to access column type when it shouldn't be needed" + ) return exp.ColumnDef( this=self.column, kind=self.column_type, @@ -249,7 +251,11 @@ def first(cls) -> TableAlterColumnPosition: def last( cls, after: t.Optional[t.Union[str, exp.Identifier]] = None ) -> TableAlterColumnPosition: - return cls(is_first=False, is_last=True, after=exp.to_identifier(after) if after else None) + return cls( + is_first=False, + is_last=True, + after=exp.to_identifier(after) if after else None, + ) @classmethod def middle(cls, after: t.Union[str, exp.Identifier]) -> TableAlterColumnPosition: @@ -366,7 +372,9 @@ class SchemaDiffer(PydanticModel): precision_increase_allowed_types: t.Optional[t.Set[exp.DType]] = None support_coercing_compatible_types: bool = False drop_cascade: bool = False - parameterized_type_defaults: t.Dict[exp.DType, t.List[t.Tuple[t.Union[int, float], ...]]] = {} + parameterized_type_defaults: t.Dict[ + exp.DType, t.List[t.Tuple[t.Union[int, float], ...]] + ] = {} max_parameter_length: t.Dict[exp.DType, t.Union[int, float]] = {} types_with_unlimited_length: t.Dict[exp.DType, t.Set[exp.DType]] = {} treat_alter_data_type_as_destructive: bool = False @@ -378,7 +386,9 @@ def coerceable_types(self) -> t.Dict[exp.DataType, t.Set[exp.DataType]]: if not self._coerceable_types: if not self.support_coercing_compatible_types or not self.compatible_types: return self.coerceable_types_ - coerceable_types: t.Dict[exp.DataType, t.Set[exp.DataType]] = defaultdict(set) + coerceable_types: t.Dict[exp.DataType, t.Set[exp.DataType]] = defaultdict( + set + ) coerceable_types.update(self.coerceable_types_) for source_type, target_types in self.compatible_types.items(): for target_type in target_types: @@ -386,7 +396,9 @@ def coerceable_types(self) -> t.Dict[exp.DataType, t.Set[exp.DataType]]: self._coerceable_types = coerceable_types return self._coerceable_types - def _is_compatible_type(self, current_type: exp.DataType, new_type: exp.DataType) -> bool: + def _is_compatible_type( + self, current_type: exp.DataType, new_type: exp.DataType + ) -> bool: # types are identical or both types are parameterized and new has higher precision # - default parameter values are automatically provided if not present if current_type == new_type or ( @@ -398,11 +410,16 @@ def _is_compatible_type(self, current_type: exp.DataType, new_type: exp.DataType if current_type in self.compatible_types: return new_type in self.compatible_types[current_type] # new type is un-parameterized and has unlimited length, current type is compatible - if not new_type.expressions and new_type.this in self.types_with_unlimited_length: + if ( + not new_type.expressions + and new_type.this in self.types_with_unlimited_length + ): return current_type.this in self.types_with_unlimited_length[new_type.this] return False - def _is_coerceable_type(self, current_type: exp.DataType, new_type: exp.DataType) -> bool: + def _is_coerceable_type( + self, current_type: exp.DataType, new_type: exp.DataType + ) -> bool: if current_type in self.coerceable_types: is_coerceable = new_type in self.coerceable_types[current_type] if is_coerceable: @@ -421,7 +438,9 @@ def _is_precision_increase_allowed(self, current_type: exp.DataType) -> bool: or current_type.this in self.precision_increase_allowed_types ) - def _is_precision_increase(self, current_type: exp.DataType, new_type: exp.DataType) -> bool: + def _is_precision_increase( + self, current_type: exp.DataType, new_type: exp.DataType + ) -> bool: if current_type.this == new_type.this and not current_type.is_type( *exp.DataType.NESTED_TYPES ): @@ -431,7 +450,9 @@ def _is_precision_increase(self, current_type: exp.DataType, new_type: exp.DataT if len(current_params) != len(new_params): return False - return all(new >= current for current, new in zip(current_params, new_params)) + return all( + new >= current for current, new in zip(current_params, new_params) + ) return False def get_type_parameters(self, type: exp.DataType) -> t.List[t.Union[int, float]]: @@ -495,7 +516,9 @@ def _drop_operation( ) -> t.List[TableAlterColumnOperation]: columns = ensure_list(columns) operations: t.List[TableAlterColumnOperation] = [] - column_pos, column_kwarg = self._get_matching_kwarg(columns[-1].name, struct, pos) + column_pos, column_kwarg = self._get_matching_kwarg( + columns[-1].name, struct, pos + ) if column_pos is None or not column_kwarg: raise SQLMeshError( f"Cannot drop column '{columns[-1].name}' from table '{table_name}' - column not found. " @@ -517,7 +540,9 @@ def _requires_drop_alteration( self, current_struct: exp.DataType, new_struct: exp.DataType ) -> bool: for current_pos, current_kwarg in enumerate(current_struct.expressions.copy()): - new_pos, _ = self._get_matching_kwarg(current_kwarg, new_struct, current_pos) + new_pos, _ = self._get_matching_kwarg( + current_kwarg, new_struct, current_pos + ) if new_pos is None: return True return False @@ -532,8 +557,12 @@ def _resolve_drop_operation( ) -> t.List[TableAlterColumnOperation]: operations = [] for current_pos, current_kwarg in enumerate(current_struct.expressions.copy()): - new_pos, _ = self._get_matching_kwarg(current_kwarg, new_struct, current_pos) - columns = parent_columns + [TableAlterColumn.from_struct_kwarg(current_kwarg)] + new_pos, _ = self._get_matching_kwarg( + current_kwarg, new_struct, current_pos + ) + columns = parent_columns + [ + TableAlterColumn.from_struct_kwarg(current_kwarg) + ] if new_pos is None: operations.extend( self._drop_operation( @@ -553,7 +582,9 @@ def _add_operation( is_part_of_destructive_change: bool = False, ) -> t.List[TableAlterColumnOperation]: if self.support_positional_add: - col_pos = TableAlterColumnPosition.create(new_pos, current_struct.expressions) + col_pos = TableAlterColumnPosition.create( + new_pos, current_struct.expressions + ) current_struct.expressions.insert(new_pos, new_kwarg) else: col_pos = None @@ -580,12 +611,21 @@ def _resolve_add_operations( ) -> t.List[TableAlterColumnOperation]: operations = [] for new_pos, new_kwarg in enumerate(new_struct.expressions): - possible_current_pos, _ = self._get_matching_kwarg(new_kwarg, current_struct, new_pos) + possible_current_pos, _ = self._get_matching_kwarg( + new_kwarg, current_struct, new_pos + ) if possible_current_pos is None: - columns = parent_columns + [TableAlterColumn.from_struct_kwarg(new_kwarg)] + columns = parent_columns + [ + TableAlterColumn.from_struct_kwarg(new_kwarg) + ] operations.extend( self._add_operation( - columns, new_pos, new_kwarg, current_struct, root_struct, table_name + columns, + new_pos, + new_kwarg, + current_struct, + root_struct, + table_name, ) ) return operations @@ -630,11 +670,18 @@ def _alter_operation( return [] new_array_type = new_type.expressions[0] current_array_type = current_type.expressions[0] - if new_array_type.this == current_array_type.this == exp.DataType.Type.STRUCT: + if ( + new_array_type.this + == current_array_type.this + == exp.DataType.Type.STRUCT + ): if self.nested_support.is_ignore: return [] - if self.nested_support.is_all or not self._requires_drop_alteration( - current_array_type, new_array_type + if ( + self.nested_support.is_all + or not self._requires_drop_alteration( + current_array_type, new_array_type + ) ): return self._get_operations( columns, @@ -694,14 +741,18 @@ def _resolve_alter_operations( ) -> t.List[TableAlterColumnOperation]: operations = [] for current_pos, current_kwarg in enumerate(current_struct.expressions.copy()): - _, new_kwarg = self._get_matching_kwarg(current_kwarg, new_struct, current_pos) + _, new_kwarg = self._get_matching_kwarg( + current_kwarg, new_struct, current_pos + ) if new_kwarg is None: if ignore_destructive: continue raise ValueError("Cannot alter a column that is being dropped") _, new_type = _get_name_and_type(new_kwarg) _, current_type = _get_name_and_type(current_kwarg) - columns = parent_columns + [TableAlterColumn.from_struct_kwarg(current_kwarg)] + columns = parent_columns + [ + TableAlterColumn.from_struct_kwarg(current_kwarg) + ] if new_type == current_type: continue operations.extend( @@ -827,7 +878,9 @@ def get_additive_changes( return [x for x in alter_operations if x.is_additive] -def get_dropped_column_names(alter_expressions: t.List[TableAlterOperation]) -> t.List[str]: +def get_dropped_column_names( + alter_expressions: t.List[TableAlterOperation], +) -> t.List[str]: return [ op.column.alias_or_name for op in alter_expressions @@ -835,7 +888,9 @@ def get_dropped_column_names(alter_expressions: t.List[TableAlterOperation]) -> ] -def get_additive_column_names(alter_expressions: t.List[TableAlterOperation]) -> t.List[str]: +def get_additive_column_names( + alter_expressions: t.List[TableAlterOperation], +) -> t.List[str]: return [ op.column.alias_or_name for op in alter_expressions @@ -856,11 +911,9 @@ def get_schema_differ( Returns: The SchemaDiffer instance configured for the given dialect. """ - from sqlmesh.core.engine_adapter import ( - DIALECT_TO_ENGINE_ADAPTER, - DIALECT_ALIASES, - EngineAdapter, - ) + from sqlmesh.core.engine_adapter import (DIALECT_ALIASES, + DIALECT_TO_ENGINE_ADAPTER, + EngineAdapter) dialect = dialect.lower() dialect = DIALECT_ALIASES.get(dialect, dialect) diff --git a/sqlmesh/core/schema_loader.py b/sqlmesh/core/schema_loader.py index 4803bdd606..045f5a4d76 100644 --- a/sqlmesh/core/schema_loader.py +++ b/sqlmesh/core/schema_loader.py @@ -52,7 +52,9 @@ def create_external_models_file( external_model_fqns.add(dep) # Make sure we don't convert internal models into external ones. - existing_model_fqns = state_reader.nodes_exist(external_model_fqns, exclude_external=True) + existing_model_fqns = state_reader.nodes_exist( + external_model_fqns, exclude_external=True + ) if existing_model_fqns: existing_model_fqns_str = ", ".join(existing_model_fqns) get_console().log_warning( diff --git a/sqlmesh/core/selector.py b/sqlmesh/core/selector.py index 54b89d2680..d29941f739 100644 --- a/sqlmesh/core/selector.py +++ b/sqlmesh/core/selector.py @@ -1,30 +1,32 @@ from __future__ import annotations +import abc import fnmatch import typing as t -from pathlib import Path from itertools import zip_longest -import abc +from pathlib import Path from sqlglot import exp -from sqlglot.errors import ParseError -from sqlglot.tokens import Token, TokenType, Tokenizer as BaseTokenizer from sqlglot.dialects.dialect import Dialect, DialectType +from sqlglot.errors import ParseError from sqlglot.helper import seq_get +from sqlglot.tokens import Token +from sqlglot.tokens import Tokenizer as BaseTokenizer +from sqlglot.tokens import TokenType from sqlmesh.core import constants as c +from sqlmesh.core.audit import StandaloneAudit from sqlmesh.core.dialect import normalize_model_name from sqlmesh.core.environment import Environment from sqlmesh.core.model import update_model_schemas -from sqlmesh.core.audit import StandaloneAudit from sqlmesh.utils import UniqueKeyDict from sqlmesh.utils.dag import DAG -from sqlmesh.utils.git import GitClient from sqlmesh.utils.errors import SQLMeshError - +from sqlmesh.utils.git import GitClient if t.TYPE_CHECKING: from typing_extensions import Literal as Lit # noqa + from sqlmesh.core.model import Model from sqlmesh.core.node import Node from sqlmesh.core.state_sync import StateReader @@ -157,7 +159,9 @@ def _load_env_models( ensure_finalized_snapshots: bool = False, ) -> t.Dict[str, "Model"]: """Loads models from the target environment, falling back to the fallback environment if needed.""" - target_env = self._state_reader.get_environment(Environment.sanitize_name(target_env_name)) + target_env = self._state_reader.get_environment( + Environment.sanitize_name(target_env_name) + ) if target_env and target_env.expired: target_env = None @@ -176,12 +180,16 @@ def _load_env_models( ) return { s.name: s.model - for s in self._state_reader.get_snapshots(environment_snapshot_infos).values() + for s in self._state_reader.get_snapshots( + environment_snapshot_infos + ).values() if s.is_model } def expand_model_selections( - self, model_selections: t.Iterable[str], models: t.Optional[t.Dict[str, Node]] = None + self, + model_selections: t.Iterable[str], + models: t.Optional[t.Dict[str, Node]] = None, ) -> t.Set[str]: """Expands a set of model selections into a set of model fqns that can be looked up in the Context. @@ -226,9 +234,13 @@ def evaluate(node: exp.Expr) -> t.Set[str]: git_modified_files = { *self._git_client.list_untracked_files(), *self._git_client.list_uncommitted_changed_files(), - *self._git_client.list_committed_changed_files(target_branch=target_branch), + *self._git_client.list_committed_changed_files( + target_branch=target_branch + ), + } + return { + m.fqn for m in all_models.values() if m._path in git_modified_files } - return {m.fqn for m in all_models.values() if m._path in git_modified_files} if isinstance(node, Tag): pattern = node.name.lower() @@ -269,7 +281,9 @@ def _model_name(self, model: Node) -> str: pass @abc.abstractmethod - def _pattern_to_model_fqns(self, pattern: str, all_models: t.Dict[str, Node]) -> t.Set[str]: + def _pattern_to_model_fqns( + self, pattern: str, all_models: t.Dict[str, Node] + ) -> t.Set[str]: """Given a pattern, return the keys of the matching models from :all_models""" pass @@ -285,7 +299,9 @@ class NativeSelector(Selector): def _model_name(self, model: Node) -> str: return model.name - def _pattern_to_model_fqns(self, pattern: str, all_models: t.Dict[str, Node]) -> t.Set[str]: + def _pattern_to_model_fqns( + self, pattern: str, all_models: t.Dict[str, Node] + ) -> t.Set[str]: fqn = normalize_model_name(pattern, self._default_catalog, self._dialect) return {fqn} if fqn in all_models else set() @@ -304,9 +320,13 @@ class DbtSelector(Selector): def _model_name(self, model: Node) -> str: if dbt_fqn := model.dbt_fqn: return dbt_fqn - raise SQLMeshError("dbt node information must be populated to use dbt selectors") + raise SQLMeshError( + "dbt node information must be populated to use dbt selectors" + ) - def _pattern_to_model_fqns(self, pattern: str, all_models: t.Dict[str, Node]) -> t.Set[str]: + def _pattern_to_model_fqns( + self, pattern: str, all_models: t.Dict[str, Node] + ) -> t.Set[str]: # a pattern like "staging.customers" should match a model called "jaffle_shop.staging.customers" # but not a model called "jaffle_shop.customers.staging" # also a pattern like "aging" should not match "staging" so we need to consider components; not substrings @@ -366,7 +386,9 @@ def _matches_resource_type(self, resource_type: str, model: Node) -> bool: return resource_type == "test" if resource_type == "model": - return model.is_model and not model.kind.is_external and not model.kind.is_seed + return ( + model.is_model and not model.kind.is_external and not model.kind.is_seed + ) if resource_type == "source": return model.kind.is_external if resource_type == "seed": @@ -433,7 +455,9 @@ def _next() -> t.Optional[Token]: def _error(msg: str) -> str: return f"{msg} at index {i}: {selector}" - def _match(token_type: TokenType, raise_unmatched: bool = False) -> t.Optional[Token]: + def _match( + token_type: TokenType, raise_unmatched: bool = False + ) -> t.Optional[Token]: token = _curr() if token and token.token_type == token_type: return _advance() diff --git a/sqlmesh/core/signal.py b/sqlmesh/core/signal.py index 554dd60a39..fd31be2cec 100644 --- a/sqlmesh/core/signal.py +++ b/sqlmesh/core/signal.py @@ -1,14 +1,14 @@ from __future__ import annotations import typing as t + from sqlmesh.utils import UniqueKeyDict, registry_decorator from sqlmesh.utils.errors import MissingSourceError if t.TYPE_CHECKING: from sqlmesh.core.context import ExecutionContext - from sqlmesh.core.snapshot.definition import Snapshot + from sqlmesh.core.snapshot.definition import DeployabilityIndex, Snapshot from sqlmesh.utils.date import DatetimeRanges - from sqlmesh.core.snapshot.definition import DeployabilityIndex class signal(registry_decorator): @@ -57,7 +57,9 @@ def freshness( if context.is_restatement or not adapter.SUPPORTS_METADATA_TABLE_LAST_MODIFIED_TS: return True - deployability_index = context.deployability_index or DeployabilityIndex.all_deployable() + deployability_index = ( + context.deployability_index or DeployabilityIndex.all_deployable() + ) last_altered_ts = ( snapshot.last_altered_ts @@ -71,7 +73,9 @@ def freshness( parent_snapshots = {context.snapshots[p.name] for p in snapshot.parents} upstream_parent_snapshots = {p for p in parent_snapshots if not p.is_external} - external_parents = snapshot.node.depends_on - {p.name for p in upstream_parent_snapshots} + external_parents = snapshot.node.depends_on - { + p.name for p in upstream_parent_snapshots + } if context.parent_intervals: # At least one upstream sqlmesh model has intervals to compute (i.e is fresh), diff --git a/sqlmesh/core/snapshot/__init__.py b/sqlmesh/core/snapshot/__init__.py index 65e5c2a822..6923da1c7a 100644 --- a/sqlmesh/core/snapshot/__init__.py +++ b/sqlmesh/core/snapshot/__init__.py @@ -1,35 +1,53 @@ -from sqlmesh.core.snapshot.definition import ( - DeployabilityIndex as DeployabilityIndex, - Intervals as Intervals, - Node as Node, - QualifiedViewName as QualifiedViewName, - Snapshot as Snapshot, - SnapshotIdAndVersion as SnapshotIdAndVersion, - SnapshotChangeCategory as SnapshotChangeCategory, - SnapshotDataVersion as SnapshotDataVersion, - SnapshotFingerprint as SnapshotFingerprint, - SnapshotId as SnapshotId, - SnapshotIdBatch as SnapshotIdBatch, - SnapshotIdLike as SnapshotIdLike, - SnapshotIdAndVersionLike as SnapshotIdAndVersionLike, - SnapshotInfoLike as SnapshotInfoLike, - SnapshotIntervals as SnapshotIntervals, - SnapshotNameVersion as SnapshotNameVersion, - SnapshotNameVersionLike as SnapshotNameVersionLike, - SnapshotTableCleanupTask as SnapshotTableCleanupTask, - SnapshotTableInfo as SnapshotTableInfo, - apply_auto_restatements as apply_auto_restatements, - earliest_start_date as earliest_start_date, - fingerprint_from_node as fingerprint_from_node, - has_paused_forward_only as has_paused_forward_only, - merge_intervals as merge_intervals, - missing_intervals as missing_intervals, - snapshots_to_dag as snapshots_to_dag, - start_date as start_date, - table_name as table_name, - to_table_mapping as to_table_mapping, -) -from sqlmesh.core.snapshot.evaluator import ( - SnapshotEvaluator as SnapshotEvaluator, - SnapshotCreationFailedError as SnapshotCreationFailedError, -) +from sqlmesh.core.snapshot.definition import \ + DeployabilityIndex as DeployabilityIndex +from sqlmesh.core.snapshot.definition import Intervals as Intervals +from sqlmesh.core.snapshot.definition import Node as Node +from sqlmesh.core.snapshot.definition import \ + QualifiedViewName as QualifiedViewName +from sqlmesh.core.snapshot.definition import Snapshot as Snapshot +from sqlmesh.core.snapshot.definition import \ + SnapshotChangeCategory as SnapshotChangeCategory +from sqlmesh.core.snapshot.definition import \ + SnapshotDataVersion as SnapshotDataVersion +from sqlmesh.core.snapshot.definition import \ + SnapshotFingerprint as SnapshotFingerprint +from sqlmesh.core.snapshot.definition import SnapshotId as SnapshotId +from sqlmesh.core.snapshot.definition import \ + SnapshotIdAndVersion as SnapshotIdAndVersion +from sqlmesh.core.snapshot.definition import \ + SnapshotIdAndVersionLike as SnapshotIdAndVersionLike +from sqlmesh.core.snapshot.definition import SnapshotIdBatch as SnapshotIdBatch +from sqlmesh.core.snapshot.definition import SnapshotIdLike as SnapshotIdLike +from sqlmesh.core.snapshot.definition import \ + SnapshotInfoLike as SnapshotInfoLike +from sqlmesh.core.snapshot.definition import \ + SnapshotIntervals as SnapshotIntervals +from sqlmesh.core.snapshot.definition import \ + SnapshotNameVersion as SnapshotNameVersion +from sqlmesh.core.snapshot.definition import \ + SnapshotNameVersionLike as SnapshotNameVersionLike +from sqlmesh.core.snapshot.definition import \ + SnapshotTableCleanupTask as SnapshotTableCleanupTask +from sqlmesh.core.snapshot.definition import \ + SnapshotTableInfo as SnapshotTableInfo +from sqlmesh.core.snapshot.definition import \ + apply_auto_restatements as apply_auto_restatements +from sqlmesh.core.snapshot.definition import \ + earliest_start_date as earliest_start_date +from sqlmesh.core.snapshot.definition import \ + fingerprint_from_node as fingerprint_from_node +from sqlmesh.core.snapshot.definition import \ + has_paused_forward_only as has_paused_forward_only +from sqlmesh.core.snapshot.definition import merge_intervals as merge_intervals +from sqlmesh.core.snapshot.definition import \ + missing_intervals as missing_intervals +from sqlmesh.core.snapshot.definition import \ + snapshots_to_dag as snapshots_to_dag +from sqlmesh.core.snapshot.definition import start_date as start_date +from sqlmesh.core.snapshot.definition import table_name as table_name +from sqlmesh.core.snapshot.definition import \ + to_table_mapping as to_table_mapping +from sqlmesh.core.snapshot.evaluator import \ + SnapshotCreationFailedError as SnapshotCreationFailedError +from sqlmesh.core.snapshot.evaluator import \ + SnapshotEvaluator as SnapshotEvaluator diff --git a/sqlmesh/core/snapshot/cache.py b/sqlmesh/core/snapshot/cache.py index d46b5f0620..28c7b30da5 100644 --- a/sqlmesh/core/snapshot/cache.py +++ b/sqlmesh/core/snapshot/cache.py @@ -2,18 +2,15 @@ import logging import typing as t - from pathlib import Path -from sqlmesh.core.model.cache import ( - OptimizedQueryCache, - optimized_query_cache_pool, - load_optimized_query, -) + from sqlmesh.core import constants as c +from sqlmesh.core.model.cache import (OptimizedQueryCache, + load_optimized_query, + optimized_query_cache_pool) from sqlmesh.core.snapshot.definition import Snapshot, SnapshotId from sqlmesh.utils.cache import FileCache - logger = logging.getLogger(__name__) @@ -77,7 +74,8 @@ def get_or_load( self._optimized_query_cache.with_optimized_query(snapshot.model) except Exception: logger.exception( - "Failed to cache optimized query for snapshot %s", snapshot.snapshot_id + "Failed to cache optimized query for snapshot %s", + snapshot.snapshot_id, ) self.put(snapshot) diff --git a/sqlmesh/core/snapshot/categorizer.py b/sqlmesh/core/snapshot/categorizer.py index 78ea7466ed..03da20acb2 100644 --- a/sqlmesh/core/snapshot/categorizer.py +++ b/sqlmesh/core/snapshot/categorizer.py @@ -66,5 +66,7 @@ def categorize_change( return default_category return ( - SnapshotChangeCategory.BREAKING if breaking_change else SnapshotChangeCategory.NON_BREAKING + SnapshotChangeCategory.BREAKING + if breaking_change + else SnapshotChangeCategory.NON_BREAKING ) diff --git a/sqlmesh/core/snapshot/definition.py b/sqlmesh/core/snapshot/definition.py index 0c9635a7c2..282b881c70 100644 --- a/sqlmesh/core/snapshot/definition.py +++ b/sqlmesh/core/snapshot/definition.py @@ -1,11 +1,11 @@ from __future__ import annotations +import logging import sys import typing as t from collections import defaultdict from datetime import datetime, timedelta from enum import IntEnum -import logging from functools import cached_property, lru_cache from pathlib import Path @@ -13,48 +13,34 @@ from sqlglot import exp from sqlglot.optimizer.normalize_identifiers import normalize_identifiers -from sqlmesh.core.config.common import ( - TableNamingConvention, - VirtualEnvironmentMode, - EnvironmentSuffixTarget, -) from sqlmesh.core import constants as c from sqlmesh.core.audit import StandaloneAudit +from sqlmesh.core.config.common import (EnvironmentSuffixTarget, + TableNamingConvention, + VirtualEnvironmentMode) from sqlmesh.core.macros import call_macro -from sqlmesh.core.model import Model, ModelKindMixin, ModelKindName, ViewKind, CustomKind +from sqlmesh.core.model import (CustomKind, Model, ModelKindMixin, + ModelKindName, ViewKind) from sqlmesh.core.model.definition import _Model from sqlmesh.core.node import IntervalUnit, NodeType from sqlmesh.utils import sanitize_name, unique from sqlmesh.utils.dag import DAG -from sqlmesh.utils.date import ( - TimeLike, - is_date, - make_inclusive, - make_exclusive, - make_inclusive_end, - now, - now_timestamp, - time_like_to_str, - to_date, - to_datetime, - to_ds, - to_timestamp, - to_ts, - validate_date_range, - yesterday, -) -from sqlmesh.utils.errors import SQLMeshError, SignalEvalError -from sqlmesh.utils.metaprogramming import ( - format_evaluated_code_exception, - Executable, -) +from sqlmesh.utils.date import (TimeLike, is_date, make_exclusive, + make_inclusive, make_inclusive_end, now, + now_timestamp, time_like_to_str, to_date, + to_datetime, to_ds, to_timestamp, to_ts, + validate_date_range, yesterday) +from sqlmesh.utils.errors import SignalEvalError, SQLMeshError from sqlmesh.utils.hashing import hash_data, md5 +from sqlmesh.utils.metaprogramming import (Executable, + format_evaluated_code_exception) from sqlmesh.utils.pydantic import PydanticModel, field_validator if t.TYPE_CHECKING: from sqlglot.dialects.dialect import DialectType - from sqlmesh.core.environment import EnvironmentNamingInfo + from sqlmesh.core.context import ExecutionContext + from sqlmesh.core.environment import EnvironmentNamingInfo Interval = t.Tuple[int, int] Intervals = t.List[Interval] @@ -224,7 +210,9 @@ def remove_pending_restatement_interval(self, start: int, end: int) -> None: def is_empty(self) -> bool: return ( - not self.intervals and not self.dev_intervals and not self.pending_restatement_intervals + not self.intervals + and not self.dev_intervals + and not self.pending_restatement_intervals ) def _add_interval(self, start: int, end: int, interval_attr: str) -> None: @@ -237,7 +225,11 @@ def _update_last_altered_ts( ) -> None: if last_altered_ts: existing_last_altered_ts = getattr(self, last_altered_attr) - setattr(self, last_altered_attr, max(existing_last_altered_ts or 0, last_altered_ts)) + setattr( + self, + last_altered_attr, + max(existing_last_altered_ts or 0, last_altered_ts), + ) def _remove_interval(self, start: int, end: int, interval_attr: str) -> None: target_intervals = getattr(self, interval_attr) @@ -252,8 +244,12 @@ class SnapshotDataVersion(PydanticModel, frozen=True): change_category: t.Optional[SnapshotChangeCategory] = None physical_schema_: t.Optional[str] = Field(default=None, alias="physical_schema") dev_table_suffix: str - table_naming_convention: TableNamingConvention = Field(default=TableNamingConvention.default) - virtual_environment_mode: VirtualEnvironmentMode = Field(default=VirtualEnvironmentMode.default) + table_naming_convention: TableNamingConvention = Field( + default=TableNamingConvention.default + ) + virtual_environment_mode: VirtualEnvironmentMode = Field( + default=VirtualEnvironmentMode.default + ) def snapshot_id(self, name: str) -> SnapshotId: return SnapshotId(name=name, identifier=self.fingerprint.to_identifier()) @@ -284,24 +280,37 @@ class QualifiedViewName(PydanticModel, frozen=True): table: str def for_environment( - self, environment_naming_info: EnvironmentNamingInfo, dialect: DialectType = None + self, + environment_naming_info: EnvironmentNamingInfo, + dialect: DialectType = None, ) -> str: - return exp.table_name(self.table_for_environment(environment_naming_info, dialect=dialect)) + return exp.table_name( + self.table_for_environment(environment_naming_info, dialect=dialect) + ) def table_for_environment( - self, environment_naming_info: EnvironmentNamingInfo, dialect: DialectType = None + self, + environment_naming_info: EnvironmentNamingInfo, + dialect: DialectType = None, ) -> exp.Table: return exp.table_( self.table_name_for_environment(environment_naming_info, dialect=dialect), db=self.schema_for_environment(environment_naming_info, dialect=dialect), - catalog=self.catalog_for_environment(environment_naming_info, dialect=dialect), + catalog=self.catalog_for_environment( + environment_naming_info, dialect=dialect + ), ) def catalog_for_environment( - self, environment_naming_info: EnvironmentNamingInfo, dialect: DialectType = None + self, + environment_naming_info: EnvironmentNamingInfo, + dialect: DialectType = None, ) -> t.Optional[str]: catalog_name: t.Optional[str] = None - if environment_naming_info.is_dev and environment_naming_info.suffix_target.is_catalog: + if ( + environment_naming_info.is_dev + and environment_naming_info.suffix_target.is_catalog + ): catalog_name = f"{self.catalog}__{environment_naming_info.name}" elif environment_naming_info.catalog_name_override: catalog_name = environment_naming_info.catalog_name_override @@ -316,7 +325,9 @@ def catalog_for_environment( return self.catalog def schema_for_environment( - self, environment_naming_info: EnvironmentNamingInfo, dialect: DialectType = None + self, + environment_naming_info: EnvironmentNamingInfo, + dialect: DialectType = None, ) -> str: normalize = environment_naming_info.normalize_name @@ -327,7 +338,10 @@ def schema_for_environment( if normalize: schema = normalize_identifiers(schema, dialect=dialect).name - if environment_naming_info.is_dev and environment_naming_info.suffix_target.is_schema: + if ( + environment_naming_info.is_dev + and environment_naming_info.suffix_target.is_schema + ): env_name = environment_naming_info.name if normalize: env_name = normalize_identifiers(env_name, dialect=dialect).name @@ -337,10 +351,15 @@ def schema_for_environment( return schema def table_name_for_environment( - self, environment_naming_info: EnvironmentNamingInfo, dialect: DialectType = None + self, + environment_naming_info: EnvironmentNamingInfo, + dialect: DialectType = None, ) -> str: table = self.table - if environment_naming_info.is_dev and environment_naming_info.suffix_target.is_table: + if ( + environment_naming_info.is_dev + and environment_naming_info.suffix_target.is_table + ): env_name = environment_naming_info.name if environment_naming_info.normalize_name: env_name = normalize_identifiers(env_name, dialect=dialect).name @@ -413,7 +432,10 @@ def virtual_environment_mode(self) -> VirtualEnvironmentMode: @property def is_forward_only(self) -> bool: - return self.forward_only or self.change_category == SnapshotChangeCategory.FORWARD_ONLY + return ( + self.forward_only + or self.change_category == SnapshotChangeCategory.FORWARD_ONLY + ) @property def is_metadata(self) -> bool: @@ -456,10 +478,17 @@ def display_name( This is just used for presenting information back to the user and `qualified_view_name` should be used when wanting a view name in all other cases. """ - return display_name(self, environment_naming_info, default_catalog, dialect=dialect) + return display_name( + self, environment_naming_info, default_catalog, dialect=dialect + ) - def data_hash_matches(self, other: t.Optional[SnapshotInfoMixin | SnapshotDataVersion]) -> bool: - return other is not None and self.fingerprint.data_hash == other.fingerprint.data_hash + def data_hash_matches( + self, other: t.Optional[SnapshotInfoMixin | SnapshotDataVersion] + ) -> bool: + return ( + other is not None + and self.fingerprint.data_hash == other.fingerprint.data_hash + ) def _table_name(self, version: str, is_deployable: bool) -> str: """Full table name pointing to the materialized location of the snapshot. @@ -541,7 +570,10 @@ def __lt__(self, other: SnapshotTableInfo) -> bool: return self.name < other.name def __eq__(self, other: t.Any) -> bool: - return isinstance(other, SnapshotTableInfo) and self.fingerprint == other.fingerprint + return ( + isinstance(other, SnapshotTableInfo) + and self.fingerprint == other.fingerprint + ) def __hash__(self) -> int: return hash((self.__class__, self.name, self.fingerprint)) @@ -761,7 +793,9 @@ def hydrate_with_intervals_by_version( """ intervals_by_name_version = defaultdict(list) for interval in intervals: - intervals_by_name_version[(interval.name, interval.version)].append(interval) + intervals_by_name_version[(interval.name, interval.version)].append( + interval + ) result = [] for snapshot in snapshots: @@ -839,7 +873,9 @@ def __hash__(self) -> int: def __lt__(self, other: Snapshot) -> bool: return self.name < other.name - def add_interval(self, start: TimeLike, end: TimeLike, is_dev: bool = False) -> None: + def add_interval( + self, start: TimeLike, end: TimeLike, is_dev: bool = False + ) -> None: """Add a newly processed time interval to the snapshot. The actual stored intervals are [start_ts, end_ts) or start epoch timestamp inclusive and end epoch @@ -856,7 +892,9 @@ def add_interval(self, start: TimeLike, end: TimeLike, is_dev: bool = False) -> f"Attempted to add an Invalid interval ({start}, {end}) to snapshot {self.snapshot_id}" ) - start_ts, end_ts = self.inclusive_exclusive(start, end, strict=False, expand=False) + start_ts, end_ts = self.inclusive_exclusive( + start, end, strict=False, expand=False + ) if start_ts >= end_ts: # Skipping partial interval. @@ -907,7 +945,9 @@ def get_removal_interval( removal_interval = self.inclusive_exclusive(start, end, strict) if not is_preview and self.full_history_restatement_only and self.intervals: - expanded_removal_interval = self.inclusive_exclusive(self.intervals[0][0], end, strict) + expanded_removal_interval = self.inclusive_exclusive( + self.intervals[0][0], end, strict + ) requested_start, requested_end = removal_interval expanded_start, expanded_end = expanded_removal_interval @@ -957,7 +997,9 @@ def inclusive_exclusive( end, self.node.interval_unit, strict=strict, - allow_partial=self.allow_partials if allow_partial is None else allow_partial, + allow_partial=( + self.allow_partials if allow_partial is None else allow_partial + ), expand=expand, ) @@ -968,7 +1010,9 @@ def merge_intervals(self, other: t.Union[Snapshot, SnapshotIntervals]) -> None: other: The target snapshot to inherit intervals from. """ effective_from_ts = self.normalized_effective_from_ts or 0 - apply_effective_from = effective_from_ts > 0 and self.identifier != other.identifier + apply_effective_from = ( + effective_from_ts > 0 and self.identifier != other.identifier + ) for start, end in other.intervals: # If the effective_from is set, then intervals that come after it must come from # the current snapshots. @@ -1042,22 +1086,29 @@ def missing_intervals( if ( not is_date(end) and not self.allow_partials - and to_timestamp(end) - to_timestamp(start) < self.node.interval_unit.milliseconds + and to_timestamp(end) - to_timestamp(start) + < self.node.interval_unit.milliseconds ): return [] deployability_index = deployability_index or DeployabilityIndex.all_deployable() intervals = ( - self.intervals if deployability_index.is_representative(self) else self.dev_intervals + self.intervals + if deployability_index.is_representative(self) + else self.dev_intervals ) if not self.evaluatable or (self.is_seed and intervals): return [] - start_ts, end_ts = (to_timestamp(ts) for ts in self.inclusive_exclusive(start, end)) + start_ts, end_ts = ( + to_timestamp(ts) for ts in self.inclusive_exclusive(start, end) + ) interval_unit = self.node.interval_unit - execution_time_ts = to_timestamp(execution_time) if execution_time else now_timestamp() + execution_time_ts = ( + to_timestamp(execution_time) if execution_time else now_timestamp() + ) upper_bound_ts = ( execution_time_ts if ignore_cron @@ -1075,7 +1126,9 @@ def missing_intervals( if self.is_model: lookback = self.model.lookback - model_end_ts = to_timestamp(make_exclusive(self.model.end)) if self.model.end else None + model_end_ts = ( + to_timestamp(make_exclusive(self.model.end)) if self.model.end else None + ) return compute_missing_intervals( interval_unit, @@ -1118,16 +1171,18 @@ def check_ready_intervals( ) return intervals - def categorize_as(self, category: SnapshotChangeCategory, forward_only: bool = False) -> None: + def categorize_as( + self, category: SnapshotChangeCategory, forward_only: bool = False + ) -> None: """Assigns the given category to this snapshot. Args: category: The change category to assign to this snapshot. forward_only: Whether or not this snapshot is applied going forward in production. """ - assert category != SnapshotChangeCategory.FORWARD_ONLY, ( - "FORWARD_ONLY change category is deprecated" - ) + assert ( + category != SnapshotChangeCategory.FORWARD_ONLY + ), "FORWARD_ONLY change category is deprecated" self.dev_version_ = self.fingerprint.to_version() is_no_rebuild = forward_only or category in ( @@ -1152,7 +1207,9 @@ def categorize_as(self, category: SnapshotChangeCategory, forward_only: bool = F previous_version = self.previous_version self.physical_schema_ = previous_version.physical_schema self.table_naming_convention = previous_version.table_naming_convention - if self.is_materialized and (category.is_indirect_non_breaking or category.is_metadata): + if self.is_materialized and ( + category.is_indirect_non_breaking or category.is_metadata + ): # Reuse the dev table for indirect non-breaking changes. self.dev_version_ = ( previous_version.data_version.dev_version @@ -1200,9 +1257,13 @@ def is_valid_start( if not self.intervals: # The start date must be aligned by the interval unit. - snapshot_start_ts = to_timestamp(interval_unit.cron_floor(snapshot_start)) + snapshot_start_ts = to_timestamp( + interval_unit.cron_floor(snapshot_start) + ) if snapshot_start_ts < to_timestamp(snapshot_start): - snapshot_start_ts = to_timestamp(interval_unit.cron_next(snapshot_start_ts)) + snapshot_start_ts = to_timestamp( + interval_unit.cron_next(snapshot_start_ts) + ) return snapshot_start_ts >= start_ts # Make sure that if there are missing intervals for this snapshot that they all occur at or after the # provided start_ts. Otherwise we know that we are doing a non-contiguous load and therefore this is not @@ -1249,7 +1310,9 @@ def needs_additive_check( and self.name not in allow_additive_snapshots ) - def get_next_auto_restatement_interval(self, execution_time: TimeLike) -> t.Optional[Interval]: + def get_next_auto_restatement_interval( + self, execution_time: TimeLike + ) -> t.Optional[Interval]: """Returns the next auto restatement interval for the snapshot. Args: @@ -1268,7 +1331,9 @@ def get_next_auto_restatement_interval(self, execution_time: TimeLike) -> t.Opti execution_time_ts = to_timestamp(execution_time) next_auto_restatement_ts = self.next_auto_restatement_ts or to_timestamp( - self.model.auto_restatement_croniter(self.created_ts).get_next(estimate=False) + self.model.auto_restatement_croniter(self.created_ts).get_next( + estimate=False + ) ) if execution_time_ts < next_auto_restatement_ts: return None @@ -1286,7 +1351,9 @@ def get_next_auto_restatement_interval(self, execution_time: TimeLike) -> t.Opti ) return (auto_restatement_start_ts, auto_restatement_end_ts) - def update_next_auto_restatement_ts(self, execution_time: TimeLike) -> t.Optional[int]: + def update_next_auto_restatement_ts( + self, execution_time: TimeLike + ) -> t.Optional[int]: """Updates the next auto restatement timestamp. Args: @@ -1303,7 +1370,9 @@ def update_next_auto_restatement_ts(self, execution_time: TimeLike) -> t.Optiona self.next_auto_restatement_ts = None else: self.next_auto_restatement_ts = to_timestamp( - self.model.auto_restatement_croniter(execution_time).get_next(estimate=False) + self.model.auto_restatement_croniter(execution_time).get_next( + estimate=False + ) ) return self.next_auto_restatement_ts @@ -1318,7 +1387,9 @@ def apply_pending_restatement_intervals(self) -> None: time_like_to_str(pending_restatement_interval[1]), self.snapshot_id, ) - self.intervals = remove_interval(self.intervals, *pending_restatement_interval) + self.intervals = remove_interval( + self.intervals, *pending_restatement_interval + ) def is_directly_modified(self, other: Snapshot) -> bool: """Returns whether or not this snapshot is directly modified in relation to the other snapshot.""" @@ -1402,7 +1473,9 @@ def snapshot_intervals(self) -> SnapshotIntervals: def is_materialized_view(self) -> bool: """Returns whether or not this snapshot's model represents a materialized view.""" return ( - self.is_model and isinstance(self.model.kind, ViewKind) and self.model.kind.materialized + self.is_model + and isinstance(self.model.kind, ViewKind) + and self.model.kind.materialized ) @property @@ -1509,7 +1582,12 @@ def expiration_ts(self) -> int: @property def supports_schema_migration_in_prod(self) -> bool: """Returns whether or not this snapshot supports schema migration when deployed to production.""" - return self.is_paused and self.is_model and not self.is_symbolic and not self.is_seed + return ( + self.is_paused + and self.is_model + and not self.is_symbolic + and not self.is_seed + ) @property def requires_schema_migration_in_prod(self) -> bool: @@ -1534,14 +1612,20 @@ def custom_materialization(self) -> t.Optional[str]: @property def virtual_environment_mode(self) -> VirtualEnvironmentMode: return ( - self.model.virtual_environment_mode if self.is_model else VirtualEnvironmentMode.default + self.model.virtual_environment_mode + if self.is_model + else VirtualEnvironmentMode.default ) def _ensure_categorized(self) -> None: if not self.change_category: - raise SQLMeshError(f"Snapshot {self.snapshot_id} has not been categorized yet.") + raise SQLMeshError( + f"Snapshot {self.snapshot_id} has not been categorized yet." + ) if not self.version: - raise SQLMeshError(f"Snapshot {self.snapshot_id} has not been versioned yet.") + raise SQLMeshError( + f"Snapshot {self.snapshot_id} has not been versioned yet." + ) def __getstate__(self) -> t.Dict[t.Any, t.Any]: state = super().__getstate__() @@ -1578,7 +1662,9 @@ class DeployabilityIndex(PydanticModel, frozen=True): @field_validator("indexed_ids", "representative_shared_version_ids", mode="before") @classmethod - def _snapshot_ids_set_validator(cls, v: t.Any) -> t.Optional[t.FrozenSet[t.Tuple[str, str]]]: + def _snapshot_ids_set_validator( + cls, v: t.Any + ) -> t.Optional[t.FrozenSet[t.Tuple[str, str]]]: if v is None: return v # Transforming into strings because the serialization of sets of objects / lists is broken in Pydantic. @@ -1636,13 +1722,19 @@ def with_deployable(self, snapshot: SnapshotIdLike) -> DeployabilityIndex: """Creates a new index with the given snapshot marked as deployable.""" return self._add_snapshot(snapshot, True) - def _add_snapshot(self, snapshot: SnapshotIdLike, deployable: bool) -> DeployabilityIndex: + def _add_snapshot( + self, snapshot: SnapshotIdLike, deployable: bool + ) -> DeployabilityIndex: snapshot_id = {self._snapshot_id_key(snapshot.snapshot_id)} indexed_ids = self.indexed_ids if self.is_opposite_index: - indexed_ids = indexed_ids - snapshot_id if deployable else indexed_ids | snapshot_id + indexed_ids = ( + indexed_ids - snapshot_id if deployable else indexed_ids | snapshot_id + ) else: - indexed_ids = indexed_ids | snapshot_id if deployable else indexed_ids - snapshot_id + indexed_ids = ( + indexed_ids | snapshot_id if deployable else indexed_ids - snapshot_id + ) return DeployabilityIndex( indexed_ids=indexed_ids, @@ -1694,18 +1786,24 @@ def create( if this_deployable: is_forward_only_model = ( - snapshot.is_model and snapshot.model.forward_only and not snapshot.is_metadata + snapshot.is_model + and snapshot.model.forward_only + and not snapshot.is_metadata ) has_auto_restatement = ( - snapshot.is_model and snapshot.model.auto_restatement_cron is not None + snapshot.is_model + and snapshot.model.auto_restatement_cron is not None ) snapshot_start = start_override_per_model.get( - node.name, start_date(snapshot, snapshots.values(), cache=start_date_cache) + node.name, + start_date(snapshot, snapshots.values(), cache=start_date_cache), ) is_valid_start = ( - snapshot.is_valid_start(start, snapshot_start) if start is not None else True + snapshot.is_valid_start(start, snapshot_start) + if start is not None + else True ) children_deployable = is_valid_start and not has_auto_restatement @@ -1737,7 +1835,9 @@ def create( children_deployability_mapping[node] = children_deployable deployable_ids = { - snapshot_id for snapshot_id, deployable in deployability_mapping.items() if deployable + snapshot_id + for snapshot_id, deployable in deployability_mapping.items() + if deployable } non_deployable_ids = set(snapshots) - deployable_ids @@ -1786,13 +1886,17 @@ def table_name( # Therefore, a model with 3-part naming like "foo.bar.baz" gets passed as (name="bar.baz", catalog="foo") to this function # This is why there is no TableNamingConvention.CATALOG_AND_SCHEMA_AND_TABLE table_parts = table.parts - parts_to_consider = 2 if naming_convention == TableNamingConvention.SCHEMA_AND_TABLE else 1 + parts_to_consider = ( + 2 if naming_convention == TableNamingConvention.SCHEMA_AND_TABLE else 1 + ) # in case the parsed table name has less parts than what the naming convention says we should be considering parts_to_consider = min(len(table_parts), parts_to_consider) # bigquery projects usually have "-" in them which is illegal in the table name, so we aggressively prune - name = "__".join(sanitize_name(part.name) for part in table_parts[-parts_to_consider:]) + name = "__".join( + sanitize_name(part.name) for part in table_parts[-parts_to_consider:] + ) full_name = f"{name}__{version}" @@ -1891,7 +1995,9 @@ def fingerprint_from_node( parent_data_hash = hash_data(sorted(p.to_version() for p in parents)) parent_metadata_hash = hash_data( - sorted(h for p in parents for h in (p.metadata_hash, p.parent_metadata_hash)) + sorted( + h for p in parents for h in (p.metadata_hash, p.parent_metadata_hash) + ) ) cache[node.fqn] = SnapshotFingerprint( @@ -1960,7 +2066,9 @@ def format_intervals(intervals: Intervals, unit: t.Optional[IntervalUnit]) -> st ) -def remove_interval(intervals: Intervals, remove_start: int, remove_end: int) -> Intervals: +def remove_interval( + intervals: Intervals, remove_start: int, remove_end: int +) -> Intervals: """Remove an interval from a list of intervals. Assumes that the correct start and end intervals have been passed in. Use `get_remove_interval` method of `Snapshot` to get the correct start/end given the snapshot's information. @@ -1996,7 +2104,9 @@ def to_table_mapping( ) -> t.Dict[str, str]: deployability_index = deployability_index or DeployabilityIndex.all_deployable() return { - snapshot.name: snapshot.table_name(deployability_index.is_representative(snapshot)) + snapshot.name: snapshot.table_name( + deployability_index.is_representative(snapshot) + ) for snapshot in snapshots if snapshot.version and not snapshot.is_embedded and snapshot.is_model } @@ -2068,7 +2178,9 @@ def missing_intervals( restated_interval = restatements.get(snapshot.snapshot_id) if restated_interval: - snapshot_start_date, snapshot_end_date = (to_datetime(i) for i in restated_interval) + snapshot_start_date, snapshot_end_date = ( + to_datetime(i) for i in restated_interval + ) snapshot = snapshot.copy() snapshot.intervals = snapshot.intervals.copy() snapshot.remove_interval(restated_interval) @@ -2083,14 +2195,18 @@ def missing_intervals( snapshot_start_date = max( to_datetime(snapshot_start_date), - to_datetime(start_date(snapshot, snapshots, cache, relative_to=snapshot_end_date)), + to_datetime( + start_date(snapshot, snapshots, cache, relative_to=snapshot_end_date) + ), ) if snapshot_start_date > to_datetime(snapshot_end_date): continue missing_interval_end_date = snapshot_end_date node_end_date = snapshot.node.end - if node_end_date and (to_datetime(node_end_date) < to_datetime(snapshot_end_date)): + if node_end_date and ( + to_datetime(node_end_date) < to_datetime(snapshot_end_date) + ): missing_interval_end_date = node_end_date intervals = snapshot.missing_intervals( @@ -2108,7 +2224,9 @@ def missing_intervals( @lru_cache(maxsize=16384) -def expand_range(start_ts: int, end_ts: int, interval_unit: IntervalUnit) -> t.List[int]: +def expand_range( + start_ts: int, end_ts: int, interval_unit: IntervalUnit +) -> t.List[int]: croniter = interval_unit.croniter(start_ts) timestamps = [start_ts] @@ -2288,7 +2406,9 @@ def start_date( Start datetime object. """ cache = {} if cache is None else cache - key = f"{snapshot.name}_{to_timestamp(relative_to)}" if relative_to else snapshot.name + key = ( + f"{snapshot.name}_{to_timestamp(relative_to)}" if relative_to else snapshot.name + ) if key in cache: return cache[key] if snapshot.node.start: @@ -2347,7 +2467,9 @@ def apply_auto_restatements( if not snapshot.is_model or snapshot.model.disable_restatement: continue - next_auto_restated_interval = snapshot.get_next_auto_restatement_interval(execution_time) + next_auto_restated_interval = snapshot.get_next_auto_restatement_interval( + execution_time + ) auto_restated_intervals = [ auto_restated_intervals_per_snapshot[parent_s_id] for parent_s_id in snapshot.parents @@ -2379,8 +2501,12 @@ def apply_auto_restatements( auto_restated_interval_start = sys.maxsize auto_restated_interval_end = -sys.maxsize for interval in auto_restated_intervals: - auto_restated_interval_start = min(auto_restated_interval_start, interval[0]) - auto_restated_interval_end = max(auto_restated_interval_end, interval[1]) + auto_restated_interval_start = min( + auto_restated_interval_start, interval[0] + ) + auto_restated_interval_end = max( + auto_restated_interval_end, interval[1] + ) interval_to_remove_start = snapshot.node.interval_unit.cron_floor( auto_restated_interval_start @@ -2394,7 +2520,9 @@ def apply_auto_restatements( ) removal_interval = snapshot.get_removal_interval( - interval_to_remove_start, interval_to_remove_end, execution_time=execution_time + interval_to_remove_start, + interval_to_remove_end, + execution_time=execution_time, ) auto_restated_intervals_per_snapshot[s_id] = removal_interval @@ -2462,7 +2590,9 @@ def check_ready_intervals( checked_intervals: Intervals = [] for interval_batch in _contiguous_intervals(intervals): - batch = [(to_datetime(start), to_datetime(end)) for start, end in interval_batch] + batch = [ + (to_datetime(start), to_datetime(end)) for start, end in interval_batch + ] try: ready_intervals = call_macro( @@ -2486,14 +2616,20 @@ def check_ready_intervals( raise SignalEvalError(f"Unknown interval {i} for signal") batch = ready_intervals else: - raise SignalEvalError(f"Expected bool | list, got {type(ready_intervals)} for signal") + raise SignalEvalError( + f"Expected bool | list, got {type(ready_intervals)} for signal" + ) - checked_intervals.extend((to_timestamp(start), to_timestamp(end)) for start, end in batch) + checked_intervals.extend( + (to_timestamp(start), to_timestamp(end)) for start, end in batch + ) return checked_intervals -def get_next_model_interval_start(snapshots: t.Iterable[Snapshot]) -> t.Optional[datetime]: +def get_next_model_interval_start( + snapshots: t.Iterable[Snapshot], +) -> t.Optional[datetime]: now_dt = now() starts = [ diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index 11b3fd1f33..57a94ee5a2 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -24,66 +24,48 @@ import abc import logging -import typing as t import sys +import typing as t from collections import defaultdict from contextlib import contextmanager from functools import reduce from sqlglot import exp, select from sqlglot.executor import execute -from tenacity import retry, stop_after_attempt, wait_exponential, retry_if_not_exception_type +from tenacity import (retry, retry_if_not_exception_type, stop_after_attempt, + wait_exponential) from sqlmesh.core import constants as c from sqlmesh.core import dialect as d from sqlmesh.core.audit import Audit, StandaloneAudit from sqlmesh.core.dialect import schema_ -from sqlmesh.core.engine_adapter.shared import InsertOverwriteStrategy, DataObjectType, DataObject -from sqlmesh.core.model.meta import GrantsTargetLayer +from sqlmesh.core.engine_adapter.shared import (DataObject, DataObjectType, + InsertOverwriteStrategy) from sqlmesh.core.macros import RuntimeStage -from sqlmesh.core.model import ( - AuditResult, - IncrementalUnmanagedKind, - Model, - SeedModel, - SCDType2ByColumnKind, - SCDType2ByTimeKind, - ViewKind, - CustomKind, -) -from sqlmesh.core.model.kind import _Incremental, DbtCustomKind -from sqlmesh.utils import CompletionStatus, columns_to_types_all_known -from sqlmesh.core.schema_diff import ( - has_drop_alteration, - TableAlterOperation, - has_additive_alteration, -) -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Intervals, - Snapshot, - SnapshotId, - SnapshotIdBatch, - SnapshotInfoLike, - SnapshotTableCleanupTask, -) +from sqlmesh.core.model import (AuditResult, CustomKind, + IncrementalUnmanagedKind, Model, + SCDType2ByColumnKind, SCDType2ByTimeKind, + SeedModel, ViewKind) +from sqlmesh.core.model.kind import DbtCustomKind, _Incremental +from sqlmesh.core.model.meta import GrantsTargetLayer +from sqlmesh.core.schema_diff import (TableAlterOperation, + has_additive_alteration, + has_drop_alteration) +from sqlmesh.core.snapshot import (DeployabilityIndex, Intervals, Snapshot, + SnapshotId, SnapshotIdBatch, + SnapshotInfoLike, SnapshotTableCleanupTask) from sqlmesh.core.snapshot.execution_tracker import QueryExecutionTracker -from sqlmesh.utils import random_id, CorrelationId, AttributeDict -from sqlmesh.utils.concurrency import ( - concurrent_apply_to_snapshots, - concurrent_apply_to_values, - NodeExecutionFailedError, -) +from sqlmesh.utils import (AttributeDict, CompletionStatus, CorrelationId, + columns_to_types_all_known, random_id) +from sqlmesh.utils.concurrency import (NodeExecutionFailedError, + concurrent_apply_to_snapshots, + concurrent_apply_to_values) from sqlmesh.utils.date import TimeLike, now, time_like_to_str -from sqlmesh.utils.errors import ( - ConfigError, - DestructiveChangeError, - MigrationNotSupportedError, - SQLMeshError, - format_destructive_change_msg, - format_additive_change_msg, - AdditiveChangeError, -) +from sqlmesh.utils.errors import (AdditiveChangeError, ConfigError, + DestructiveChangeError, + MigrationNotSupportedError, SQLMeshError, + format_additive_change_msg, + format_destructive_change_msg) from sqlmesh.utils.jinja import MacroReturnVal if sys.version_info >= (3, 12): @@ -101,7 +83,9 @@ class SnapshotCreationFailedError(SQLMeshError): def __init__( - self, errors: t.List[NodeExecutionFailedError[SnapshotId]], skipped: t.List[SnapshotId] + self, + errors: t.List[NodeExecutionFailedError[SnapshotId]], + skipped: t.List[SnapshotId], ): messages = "\n\n".join(f"{error}\n {error.__cause__}" for error in errors) super().__init__(f"Physical table creation failed:\n\n{messages}") @@ -132,11 +116,15 @@ def __init__( selected_gateway: t.Optional[str] = None, ): self.adapters = ( - adapters if isinstance(adapters, t.Dict) else {selected_gateway or "": adapters} + adapters + if isinstance(adapters, t.Dict) + else {selected_gateway or "": adapters} ) self.execution_tracker = QueryExecutionTracker() self.adapters = { - gateway: adapter.with_settings(query_execution_tracker=self.execution_tracker) + gateway: adapter.with_settings( + query_execution_tracker=self.execution_tracker + ) for gateway, adapter in self.adapters.items() } self.adapter = ( @@ -259,7 +247,9 @@ def evaluate_and_fetch( existing_limit = query_or_df.args.get("limit") if existing_limit: - limit = min(limit, execute(exp.select(existing_limit.expression)).rows[0][0]) + limit = min( + limit, execute(exp.select(existing_limit.expression)).rows[0][0] + ) assert limit is not None return adapter._fetch_native_df(query_or_df.limit(limit)) @@ -286,11 +276,15 @@ def promote( on_complete: A callback to call on each successfully promoted snapshot. """ - tables_by_gateway: t.Dict[t.Union[str, None], t.List[exp.Table]] = defaultdict(list) + 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 + snapshot.model_gateway + if environment_naming_info.gateway_managed + else None ) adapter = self.get_adapter(gateway) table = snapshot.qualified_view_name.table_for_environment( @@ -304,7 +298,9 @@ def promote( self._create_catalogs(tables=tables, gateway=gateway) gateway_table_pairs = [ - (gateway, table) for gateway, tables in tables_by_gateway.items() for table in tables + (gateway, table) + for gateway, tables in tables_by_gateway.items() + for table in tables ] self._create_schemas(gateway_table_pairs=gateway_table_pairs) @@ -383,7 +379,9 @@ def create( """ deployability_index = deployability_index or DeployabilityIndex.all_deployable() - snapshots_to_create = self.get_snapshots_to_create(target_snapshots, deployability_index) + snapshots_to_create = self.get_snapshots_to_create( + target_snapshots, deployability_index + ) if not snapshots_to_create: return CompletionStatus.NOTHING_TO_DO if on_start: @@ -412,16 +410,22 @@ def create_physical_schemas( for snapshot in snapshots: if snapshot.is_model and not snapshot.is_symbolic: tables_by_gateway[snapshot.model_gateway].append( - snapshot.table_name(is_deployable=deployability_index.is_deployable(snapshot)) + snapshot.table_name( + is_deployable=deployability_index.is_deployable(snapshot) + ) ) gateway_table_pairs = [ - (gateway, table) for gateway, tables in tables_by_gateway.items() for table in tables + (gateway, table) + for gateway, tables in tables_by_gateway.items() + for table in tables ] self._create_schemas(gateway_table_pairs=gateway_table_pairs) def get_snapshots_to_create( - self, target_snapshots: t.Iterable[Snapshot], deployability_index: DeployabilityIndex + self, + target_snapshots: t.Iterable[Snapshot], + deployability_index: DeployabilityIndex, ) -> t.List[Snapshot]: """Returns a list of snapshots that need to have their physical tables created. @@ -488,7 +492,9 @@ def migrate( deployability_index: Determines snapshots that are deployable in the context of this evaluation. """ deployability_index = deployability_index or DeployabilityIndex.all_deployable() - target_data_objects = self._get_physical_data_objects(target_snapshots, deployability_index) + target_data_objects = self._get_physical_data_objects( + target_snapshots, deployability_index + ) if not target_data_objects: return @@ -526,7 +532,9 @@ def cleanup( on_complete: A callback to call on each successfully deleted database object. """ target_snapshots = [ - t for t in target_snapshots if t.snapshot.is_model and not t.snapshot.is_symbolic + t + for t in target_snapshots + if t.snapshot.is_model and not t.snapshot.is_symbolic ] available_gateways = set(self.adapters.keys()) skipped = [] @@ -560,7 +568,9 @@ def cleanup( raise_on_error=False, ) if errors: - errored_snapshots = "\n".join(f" {e.node.name}: {e.__cause__}" for e in errors) + errored_snapshots = "\n".join( + f" {e.node.name}: {e.__cause__}" for e in errors + ) raise SQLMeshError(f"\n{errored_snapshots}") def audit( @@ -596,7 +606,9 @@ def audit( ) if wap_id is not None: - deployability_index = deployability_index or DeployabilityIndex.all_deployable() + deployability_index = ( + deployability_index or DeployabilityIndex.all_deployable() + ) original_table_name = snapshot.table_name( is_deployable=deployability_index.is_deployable(snapshot) ) @@ -621,7 +633,10 @@ def audit( if audits_with_args: logger.info("Auditing snapshot %s", snapshot.snapshot_id) - if not deployability_index.is_deployable(snapshot) and not adapter.SUPPORTS_CLONING: + if ( + not deployability_index.is_deployable(snapshot) + and not adapter.SUPPORTS_CLONING + ): # For dev preview tables that aren't based on clones of the production table, only a subset of the data is typically available # However, users still expect audits to run anwyay. Some audits (such as row count) are practically guaranteed to fail # when run on only a subset of data, so we switch all audits to non blocking and the user can decide if they still want to proceed @@ -743,7 +758,9 @@ def _evaluate_snapshot( # Use the 'creating' stage if the table doesn't exist yet to preserve backwards compatibility with existing projects # that depend on a separate physical table creation stage. - runtime_stage = RuntimeStage.EVALUATING if target_table_exists else RuntimeStage.CREATING + runtime_stage = ( + RuntimeStage.EVALUATING if target_table_exists else RuntimeStage.CREATING + ) common_render_kwargs = dict( start=start, end=end, @@ -777,7 +794,9 @@ def _evaluate_snapshot( with ( adapter.transaction(), - adapter.session(snapshot.model.render_session_properties(**render_statements_kwargs)), + adapter.session( + snapshot.model.render_session_properties(**render_statements_kwargs) + ), ): evaluation_strategy.run_pre_statements( snapshot=snapshot, @@ -793,7 +812,9 @@ def _evaluate_snapshot( ) if not should_create_empty_table: # Or if the model is self-referential and its query is fully annotated with types - should_create_empty_table = model.depends_on_self and model.annotated + should_create_empty_table = ( + model.depends_on_self and model.annotated + ) if self._can_clone(snapshot, deployability_index): self._clone_snapshot_in_dev( snapshot=snapshot, @@ -806,7 +827,11 @@ def _evaluate_snapshot( ) runtime_stage = RuntimeStage.EVALUATING target_table_exists = True - elif should_create_empty_table or model.is_seed or model.kind.is_scd_type_2: + elif ( + should_create_empty_table + or model.is_seed + or model.kind.is_scd_type_2 + ): self._execute_create( snapshot=snapshot, table_name=target_table_name, @@ -834,7 +859,9 @@ def _evaluate_snapshot( and (model.wap_supported or adapter.wap_supported(target_table_name)) ): wap_id = random_id()[0:8] - logger.info("Using WAP ID '%s' for snapshot %s", wap_id, snapshot.snapshot_id) + logger.info( + "Using WAP ID '%s' for snapshot %s", wap_id, snapshot.snapshot_id + ) target_table_name = adapter.wap_prepare(target_table_name, wap_id) self._render_and_insert_snapshot( @@ -898,12 +925,15 @@ def create_snapshot( evaluation_strategy = _evaluation_strategy(snapshot, adapter) evaluation_strategy.run_pre_statements( - snapshot=snapshot, render_kwargs={**create_render_kwargs, "inside_transaction": False} + snapshot=snapshot, + render_kwargs={**create_render_kwargs, "inside_transaction": False}, ) with ( adapter.transaction(), - adapter.session(snapshot.model.render_session_properties(**create_render_kwargs)), + adapter.session( + snapshot.model.render_session_properties(**create_render_kwargs) + ), ): rendered_physical_properties = snapshot.model.render_physical_properties( **create_render_kwargs @@ -933,7 +963,8 @@ def create_snapshot( ) evaluation_strategy.run_post_statements( - snapshot=snapshot, render_kwargs={**create_render_kwargs, "inside_transaction": False} + snapshot=snapshot, + render_kwargs={**create_render_kwargs, "inside_transaction": False}, ) if on_complete is not None: @@ -946,7 +977,9 @@ def wap_publish_snapshot( deployability_index: t.Optional[DeployabilityIndex], ) -> None: deployability_index = deployability_index or DeployabilityIndex.all_deployable() - table_name = snapshot.table_name(is_deployable=deployability_index.is_deployable(snapshot)) + table_name = snapshot.table_name( + is_deployable=deployability_index.is_deployable(snapshot) + ) adapter = self.get_adapter(snapshot.model_gateway) adapter.wap_publish(table_name, wap_id) @@ -1099,7 +1132,9 @@ def _clone_snapshot_in_dev( source_table_name = snapshot.table_name() try: - logger.info(f"Cloning table '{source_table_name}' into '{target_table_name}'") + logger.info( + f"Cloning table '{source_table_name}' into '{target_table_name}'" + ) adapter.clone_table( target_table_name, snapshot.table_name(), @@ -1145,7 +1180,8 @@ def _migrate_snapshot( evaluation_strategy = _evaluation_strategy(snapshot, adapter) evaluation_strategy.run_pre_statements( - snapshot=snapshot, render_kwargs={**render_kwargs, "inside_transaction": False} + snapshot=snapshot, + render_kwargs={**render_kwargs, "inside_transaction": False}, ) with ( @@ -1186,7 +1222,8 @@ def _migrate_snapshot( ) evaluation_strategy.run_post_statements( - snapshot=snapshot, render_kwargs={**render_kwargs, "inside_transaction": False} + snapshot=snapshot, + render_kwargs={**render_kwargs, "inside_transaction": False}, ) # Retry in case when the table is migrated concurrently from another plan application @@ -1270,7 +1307,9 @@ def _promote_snapshot( if environment_naming_info.gateway_managed else self.adapter ) - table_name = snapshot.table_name(deployability_index.is_representative(snapshot)) + table_name = snapshot.table_name( + deployability_index.is_representative(snapshot) + ) view_name = snapshot.qualified_view_name.for_environment( environment_naming_info, dialect=adapter.dialect ) @@ -1373,7 +1412,8 @@ def _cleanup_snapshot( if adapter.get_data_object(table_name) is not None: raise logger.warning( - "Skipping cleanup of table '%s' because it does not exist", table_name + "Skipping cleanup of table '%s' because it does not exist", + table_name, ) if on_complete is not None: @@ -1422,7 +1462,9 @@ def _audit( elif isinstance(audit, StandaloneAudit): query = audit.render_audit_query(**kwargs) else: - raise SQLMeshError("Expected model or standalone audit. {snapshot}: {audit}") + raise SQLMeshError( + "Expected model or standalone audit. {snapshot}: {audit}" + ) count, *_ = adapter.fetchone( select("COUNT(*)").from_(query.subquery("audit")), @@ -1446,13 +1488,17 @@ def _create_catalogs( # attempt to create catalogs for the virtual layer if possible adapter = self.get_adapter(gateway) if adapter.SUPPORTS_CREATE_DROP_CATALOG: - unique_catalogs = {t.catalog for t in [exp.to_table(maybe_t) for maybe_t in tables]} + unique_catalogs = { + t.catalog for t in [exp.to_table(maybe_t) for maybe_t in tables] + } for catalog_name in unique_catalogs: adapter.create_catalog(catalog_name) def _create_schemas( self, - gateway_table_pairs: t.Iterable[t.Tuple[t.Optional[str], t.Union[exp.Table, str]]], + gateway_table_pairs: t.Iterable[ + t.Tuple[t.Optional[str], t.Union[exp.Table, str]] + ], ) -> None: table_exprs = [(gateway, exp.to_table(t)) for gateway, t in gateway_table_pairs] unique_schemas = { @@ -1481,7 +1527,9 @@ def get_adapter(self, gateway: t.Optional[str] = None) -> EngineAdapter: if gateway: if adapter := self.adapters.get(gateway): return adapter - raise SQLMeshError(f"Gateway '{gateway}' not found in the available engine adapters.") + raise SQLMeshError( + f"Gateway '{gateway}' not found in the available engine adapters." + ) return self.adapter def _execute_create( @@ -1531,7 +1579,9 @@ def _execute_create( render_kwargs={**create_render_kwargs, "inside_transaction": True}, ) - def _can_clone(self, snapshot: Snapshot, deployability_index: DeployabilityIndex) -> bool: + def _can_clone( + self, snapshot: Snapshot, deployability_index: DeployabilityIndex + ) -> bool: adapter = self.get_adapter(snapshot.model.gateway) return ( snapshot.is_forward_only @@ -1565,7 +1615,8 @@ def _get_physical_data_objects( return self._get_data_objects( target_snapshots, lambda s: exp.to_table( - s.table_name(deployability_index.is_deployable(s)), dialect=s.model.dialect + s.table_name(deployability_index.is_deployable(s)), + dialect=s.model.dialect, ), ) @@ -1615,16 +1666,20 @@ def _get_data_objects( A dictionary of snapshot IDs to existing data objects. If the data object for a snapshot is not found, it will not be included in the dictionary. """ - tables_by_gateway_and_schema: t.Dict[t.Union[str, None], t.Dict[exp.Table, set[str]]] = ( - defaultdict(lambda: defaultdict(set)) + tables_by_gateway_and_schema: t.Dict[ + t.Union[str, None], t.Dict[exp.Table, set[str]] + ] = defaultdict(lambda: defaultdict(set)) + snapshots_by_table_name: t.Dict[exp.Table, t.Dict[str, Snapshot]] = defaultdict( + dict ) - snapshots_by_table_name: t.Dict[exp.Table, t.Dict[str, Snapshot]] = defaultdict(dict) for snapshot in target_snapshots: if not snapshot.is_model or snapshot.is_symbolic: continue table = table_name_callable(snapshot) table_schema = d.schema_(table.db, catalog=table.catalog) - tables_by_gateway_and_schema[snapshot.model_gateway][table_schema].add(table.name) + tables_by_gateway_and_schema[snapshot.model_gateway][table_schema].add( + table.name + ) snapshots_by_table_name[table_schema][table.name] = snapshot def _get_data_objects_in_schema( @@ -1654,12 +1709,16 @@ def _get_data_objects_in_schema( snapshots_by_name = snapshots_by_table_name.get(schema, {}) for obj in objs: if obj.name in snapshots_by_name: - snapshot_id_to_obj[snapshots_by_name[obj.name].snapshot_id] = obj + snapshot_id_to_obj[ + snapshots_by_name[obj.name].snapshot_id + ] = obj return snapshot_id_to_obj -def _evaluation_strategy(snapshot: SnapshotInfoLike, adapter: EngineAdapter) -> EvaluationStrategy: +def _evaluation_strategy( + snapshot: SnapshotInfoLike, adapter: EngineAdapter +) -> EvaluationStrategy: klass: t.Type if snapshot.is_embedded: klass = EmbeddedStrategy @@ -1699,7 +1758,9 @@ def _evaluation_strategy(snapshot: SnapshotInfoLike, adapter: EngineAdapter) -> raise SQLMeshError( f"Missing the name of a custom evaluation strategy in model '{snapshot.name}'." ) - _, klass = get_custom_materialization_type_or_raise(snapshot.custom_materialization) + _, klass = get_custom_materialization_type_or_raise( + snapshot.custom_materialization + ) return klass(adapter) elif snapshot.is_managed: klass = EngineManagedStrategy @@ -1974,10 +2035,14 @@ def promote( def demote(self, view_name: str, **kwargs: t.Any) -> None: pass - def run_pre_statements(self, snapshot: Snapshot, render_kwargs: t.Dict[str, t.Any]) -> None: + def run_pre_statements( + self, snapshot: Snapshot, render_kwargs: t.Dict[str, t.Any] + ) -> None: pass - def run_post_statements(self, snapshot: Snapshot, render_kwargs: t.Dict[str, t.Any]) -> None: + def run_post_statements( + self, snapshot: Snapshot, render_kwargs: t.Dict[str, t.Any] + ) -> None: pass @@ -2032,7 +2097,9 @@ def promote( ) # Apply grants to the virtual layer (view) after promotion - self._apply_grants(model, view_name, GrantsTargetLayer.VIRTUAL, is_snapshot_deployable) + self._apply_grants( + model, view_name, GrantsTargetLayer.VIRTUAL, is_snapshot_deployable + ) def demote(self, view_name: str, **kwargs: t.Any) -> None: logger.info("Dropping view '%s'", view_name) @@ -2089,7 +2156,9 @@ def create( ) -> None: ctas_query = model.ctas_query(**render_kwargs) physical_properties = _adjust_physical_properties_for_engine( - self.adapter, model, kwargs.get("physical_properties", model.physical_properties) + self.adapter, + model, + kwargs.get("physical_properties", model.physical_properties), ) logger.info("Creating table '%s'", table_name) @@ -2104,7 +2173,9 @@ def create( clustered_by=model.clustered_by, table_properties=physical_properties, table_description=model.description if is_table_deployable else None, - column_descriptions=model.column_descriptions if is_table_deployable else None, + column_descriptions=( + model.column_descriptions if is_table_deployable else None + ), ) # If we create both temp and prod tables, we need to make sure that we dry run once. @@ -2128,7 +2199,9 @@ def create( clustered_by=model.clustered_by, table_properties=physical_properties, table_description=model.description if is_table_deployable else None, - column_descriptions=model.column_descriptions if is_table_deployable else None, + column_descriptions=( + model.column_descriptions if is_table_deployable else None + ), ) # Apply grants after table creation (unless explicitly skipped by caller) @@ -2166,10 +2239,15 @@ def migrate( # Apply grants after schema migration deployability_index = kwargs.get("deployability_index") is_snapshot_deployable = ( - deployability_index.is_deployable(snapshot) if deployability_index else False + deployability_index.is_deployable(snapshot) + if deployability_index + else False ) self._apply_grants( - snapshot.model, target_table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + snapshot.model, + target_table_name, + GrantsTargetLayer.PHYSICAL, + is_snapshot_deployable, ) def delete(self, name: str, **kwargs: t.Any) -> None: @@ -2206,7 +2284,9 @@ def _replace_query_for_model( columns_to_types, source_columns = None, None physical_properties = _adjust_physical_properties_for_engine( - self.adapter, model, kwargs.get("physical_properties", model.physical_properties) + self.adapter, + model, + kwargs.get("physical_properties", model.physical_properties), ) self.adapter.replace_query( name, @@ -2226,7 +2306,9 @@ def _replace_query_for_model( # Apply grants after table replacement (unless explicitly skipped by caller) if not skip_grants: is_snapshot_deployable = kwargs.get("is_snapshot_deployable", False) - self._apply_grants(model, name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable) + self._apply_grants( + model, name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + ) def _get_target_and_source_columns( self, @@ -2293,7 +2375,9 @@ def insert( **kwargs: t.Any, ) -> None: if is_first_insert: - self._replace_query_for_model(model, table_name, query_or_df, render_kwargs, **kwargs) + self._replace_query_for_model( + model, table_name, query_or_df, render_kwargs, **kwargs + ) else: columns_to_types, source_columns = self._get_target_and_source_columns( model, table_name, render_kwargs=render_kwargs @@ -2343,7 +2427,9 @@ def insert( **kwargs: t.Any, ) -> None: if is_first_insert: - self._replace_query_for_model(model, table_name, query_or_df, render_kwargs, **kwargs) + self._replace_query_for_model( + model, table_name, query_or_df, render_kwargs, **kwargs + ) else: columns_to_types, source_columns = self._get_target_and_source_columns( model, @@ -2351,7 +2437,9 @@ def insert( render_kwargs=render_kwargs, ) physical_properties = _adjust_physical_properties_for_engine( - self.adapter, model, kwargs.get("physical_properties", model.physical_properties) + self.adapter, + model, + kwargs.get("physical_properties", model.physical_properties), ) self.adapter.merge( table_name, @@ -2380,7 +2468,9 @@ def append( model, table_name, render_kwargs=render_kwargs ) physical_properties = _adjust_physical_properties_for_engine( - self.adapter, model, kwargs.get("physical_properties", model.physical_properties) + self.adapter, + model, + kwargs.get("physical_properties", model.physical_properties), ) self.adapter.merge( table_name, @@ -2430,7 +2520,10 @@ def insert( return self._replace_query_for_model( model, table_name, query_or_df, render_kwargs, **kwargs ) - if isinstance(model.kind, IncrementalUnmanagedKind) and model.kind.insert_overwrite: + if ( + isinstance(model.kind, IncrementalUnmanagedKind) + and model.kind.insert_overwrite + ): columns_to_types, source_columns = self._get_target_and_source_columns( model, table_name, @@ -2477,7 +2570,9 @@ def insert( render_kwargs: t.Dict[str, t.Any], **kwargs: t.Any, ) -> None: - self._replace_query_for_model(model, table_name, query_or_df, render_kwargs, **kwargs) + self._replace_query_for_model( + model, table_name, query_or_df, render_kwargs, **kwargs + ) class SeedStrategy(MaterializableStrategy): @@ -2530,7 +2625,10 @@ def create( # Apply grants after seed table creation and data insertion is_snapshot_deployable = kwargs.get("is_snapshot_deployable", False) self._apply_grants( - model, table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + model, + table_name, + GrantsTargetLayer.PHYSICAL, + is_snapshot_deployable, ) except Exception: self.adapter.drop_table(table_name) @@ -2587,9 +2685,13 @@ def create( logger.info("Creating table '%s'", table_name) columns_to_types = model.columns_to_types_or_raise if isinstance(model.kind, SCDType2ByTimeKind): - columns_to_types[model.kind.updated_at_name.name] = model.kind.time_data_type + columns_to_types[model.kind.updated_at_name.name] = ( + model.kind.time_data_type + ) physical_properties = _adjust_physical_properties_for_engine( - self.adapter, model, kwargs.get("physical_properties", model.physical_properties) + self.adapter, + model, + kwargs.get("physical_properties", model.physical_properties), ) self.adapter.create_table( table_name, @@ -2601,7 +2703,9 @@ def create( clustered_by=model.clustered_by, table_properties=physical_properties, table_description=model.description if is_table_deployable else None, - column_descriptions=model.column_descriptions if is_table_deployable else None, + column_descriptions=( + model.column_descriptions if is_table_deployable else None + ), ) else: # We assume that the data type for `updated_at_name` matches the data type that is defined for @@ -2660,7 +2764,9 @@ def insert( partitioned_by=model.partitioned_by, partition_interval_unit=model.partition_interval_unit, clustered_by=model.clustered_by, - table_properties=kwargs.get("physical_properties", model.physical_properties), + table_properties=kwargs.get( + "physical_properties", model.physical_properties + ), ) elif isinstance(model.kind, SCDType2ByColumnKind): self.adapter.scd_type_2_by_column( @@ -2683,7 +2789,9 @@ def insert( partitioned_by=model.partitioned_by, partition_interval_unit=model.partition_interval_unit, clustered_by=model.clustered_by, - table_properties=kwargs.get("physical_properties", model.physical_properties), + table_properties=kwargs.get( + "physical_properties", model.physical_properties + ), ) else: raise SQLMeshError( @@ -2692,7 +2800,9 @@ def insert( # Apply grants after SCD Type 2 table recreation is_snapshot_deployable = kwargs.get("is_snapshot_deployable", False) - self._apply_grants(model, table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable) + self._apply_grants( + model, table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + ) def append( self, @@ -2767,14 +2877,18 @@ def insert( replace=must_recreate_view, materialized=is_materialized_view, materialized_properties=materialized_properties, - view_properties=kwargs.get("physical_properties", model.physical_properties), + view_properties=kwargs.get( + "physical_properties", model.physical_properties + ), table_description=model.description, column_descriptions=model.column_descriptions, ) # Apply grants after view creation / replacement is_snapshot_deployable = kwargs.get("is_snapshot_deployable", False) - self._apply_grants(model, table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable) + self._apply_grants( + model, table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + ) def append( self, @@ -2805,7 +2919,10 @@ def create( if not skip_grants: # Always apply grants when present, even if view exists, to handle grants updates self._apply_grants( - model, table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + model, + table_name, + GrantsTargetLayer.PHYSICAL, + is_snapshot_deployable, ) return @@ -2826,9 +2943,13 @@ def create( replace=False, materialized=self._is_materialized_view(model), materialized_properties=materialized_properties, - view_properties=kwargs.get("physical_properties", model.physical_properties), + view_properties=kwargs.get( + "physical_properties", model.physical_properties + ), table_description=model.description if is_table_deployable else None, - column_descriptions=model.column_descriptions if is_table_deployable else None, + column_descriptions=( + model.column_descriptions if is_table_deployable else None + ), ) if not skip_grants: @@ -2850,7 +2971,9 @@ def migrate( logger.info("Migrating view '%s'", target_table_name) model = snapshot.model render_kwargs = dict( - execution_time=now(), snapshots=kwargs["snapshots"], engine_adapter=self.adapter + execution_time=now(), + snapshots=kwargs["snapshots"], + engine_adapter=self.adapter, ) is_materialized_view = self._is_materialized_view(model) @@ -2877,10 +3000,15 @@ def migrate( # Apply grants after view migration deployability_index = kwargs.get("deployability_index") is_snapshot_deployable = ( - deployability_index.is_deployable(snapshot) if deployability_index else False + deployability_index.is_deployable(snapshot) + if deployability_index + else False ) self._apply_grants( - snapshot.model, target_table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + snapshot.model, + target_table_name, + GrantsTargetLayer.PHYSICAL, + is_snapshot_deployable, ) def delete(self, name: str, **kwargs: t.Any) -> None: @@ -2942,7 +3070,9 @@ def insert( ] = None -def get_custom_materialization_kind_type(st: t.Type[CustomMaterialization]) -> t.Type[CustomKind]: +def get_custom_materialization_kind_type( + st: t.Type[CustomMaterialization], +) -> t.Type[CustomKind]: # try to read if there is a custom 'kind' type in use by inspecting the type signature # eg try to read 'MyCustomKind' from: # >>>> class MyCustomMaterialization(CustomMaterialization[MyCustomKind]) @@ -2993,7 +3123,9 @@ def get_custom_materialization_type( } if strategy_key not in _custom_materialization_type_cache: - raise ConfigError(f"Materialization strategy with name '{name}' was not found.") + raise ConfigError( + f"Materialization strategy with name '{name}' was not found." + ) except (SQLMeshError, ConfigError) as e: if raise_errors: raise e @@ -3005,7 +3137,10 @@ def get_custom_materialization_type( strategy_kind_type, strategy_type = _custom_materialization_type_cache[strategy_key] logger.debug( - "Resolved custom materialization '%s' to '%s' (%s)", name, strategy_type, strategy_kind_type + "Resolved custom materialization '%s' to '%s' (%s)", + name, + strategy_type, + strategy_kind_type, ) return strategy_kind_type, strategy_type @@ -3019,7 +3154,9 @@ def get_custom_materialization_type_or_raise( return types[0], types[1] # Shouldnt get here as get_custom_materialization_type() has raise_errors=True, but just in case... - raise SQLMeshError(f"Custom materialization '{name}' not present in the Python environment") + raise SQLMeshError( + f"Custom materialization '{name}' not present in the Python environment" + ) class DbtCustomMaterializationStrategy(MaterializableStrategy): @@ -3201,7 +3338,9 @@ def create( # We could deploy this to prod; create a proper managed table logger.info("Creating managed table: %s", table_name) physical_properties = _adjust_physical_properties_for_engine( - self.adapter, model, kwargs.get("physical_properties", model.physical_properties) + self.adapter, + model, + kwargs.get("physical_properties", model.physical_properties), ) self.adapter.create_managed_table( table_name=table_name, @@ -3218,7 +3357,10 @@ def create( # Apply grants after managed table creation if not skip_grants: self._apply_grants( - model, table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + model, + table_name, + GrantsTargetLayer.PHYSICAL, + is_snapshot_deployable, ) elif not is_table_deployable: @@ -3247,14 +3389,20 @@ def insert( deployability_index: DeployabilityIndex = kwargs["deployability_index"] snapshot: Snapshot = kwargs["snapshot"] is_snapshot_deployable = deployability_index.is_deployable(snapshot) - if is_first_insert and is_snapshot_deployable and not self.adapter.table_exists(table_name): + if ( + is_first_insert + and is_snapshot_deployable + and not self.adapter.table_exists(table_name) + ): self.adapter.create_managed_table( table_name=table_name, query=query_or_df, # type: ignore target_columns_to_types=model.columns_to_types, partitioned_by=model.partitioned_by, clustered_by=model.clustered_by, # type: ignore[arg-type] - table_properties=kwargs.get("physical_properties", model.physical_properties), + table_properties=kwargs.get( + "physical_properties", model.physical_properties + ), table_description=model.description, column_descriptions=model.column_descriptions, table_format=model.table_format, @@ -3313,10 +3461,15 @@ def migrate( # Apply grants after verifying no schema changes deployability_index = kwargs.get("deployability_index") is_snapshot_deployable = ( - deployability_index.is_deployable(snapshot) if deployability_index else False + deployability_index.is_deployable(snapshot) + if deployability_index + else False ) self._apply_grants( - snapshot.model, target_table_name, GrantsTargetLayer.PHYSICAL, is_snapshot_deployable + snapshot.model, + target_table_name, + GrantsTargetLayer.PHYSICAL, + is_snapshot_deployable, ) def delete(self, name: str, **kwargs: t.Any) -> None: @@ -3330,7 +3483,9 @@ def delete(self, name: str, **kwargs: t.Any) -> None: logger.info("Dropped dev preview for managed table '%s'", name) -def _intervals(snapshot: Snapshot, deployability_index: DeployabilityIndex) -> Intervals: +def _intervals( + snapshot: Snapshot, deployability_index: DeployabilityIndex +) -> Intervals: return ( snapshot.intervals if deployability_index.is_deployable(snapshot) @@ -3362,7 +3517,9 @@ def _check_destructive_schema_change( ) return raise DestructiveChangeError( - format_destructive_change_msg(snapshot_name, alter_operations, model_dialect) + format_destructive_change_msg( + snapshot_name, alter_operations, model_dialect + ) ) @@ -3375,9 +3532,9 @@ def _check_additive_schema_change( if not isinstance(snapshot.model.kind, _Incremental): return - if snapshot.needs_additive_check(allow_additive_snapshots) and has_additive_alteration( - alter_operations - ): + if snapshot.needs_additive_check( + allow_additive_snapshots + ) and has_additive_alteration(alter_operations): # Note: IGNORE filtering is applied before this function is called # so if we reach here, additive changes are not being ignored snapshot_name = snapshot.name @@ -3395,7 +3552,9 @@ def _check_additive_schema_change( return if snapshot.model.on_additive_change.is_error: raise AdditiveChangeError( - format_additive_change_msg(snapshot_name, alter_operations, model_dialect) + format_additive_change_msg( + snapshot_name, alter_operations, model_dialect + ) ) diff --git a/sqlmesh/core/snapshot/execution_tracker.py b/sqlmesh/core/snapshot/execution_tracker.py index bcafec8d28..22ec39c547 100644 --- a/sqlmesh/core/snapshot/execution_tracker.py +++ b/sqlmesh/core/snapshot/execution_tracker.py @@ -2,8 +2,9 @@ import typing as t from contextlib import contextmanager -from threading import local from dataclasses import dataclass, field +from threading import local + from sqlmesh.core.snapshot import SnapshotIdBatch diff --git a/sqlmesh/core/state_sync/__init__.py b/sqlmesh/core/state_sync/__init__.py index 12ea77ac8f..c5ec781bf3 100644 --- a/sqlmesh/core/state_sync/__init__.py +++ b/sqlmesh/core/state_sync/__init__.py @@ -14,10 +14,9 @@ adapter to read and write state to the underlying data store. """ -from sqlmesh.core.state_sync.base import ( - StateReader as StateReader, - StateSync as StateSync, - Versions as Versions, -) +from sqlmesh.core.state_sync.base import StateReader as StateReader +from sqlmesh.core.state_sync.base import StateSync as StateSync +from sqlmesh.core.state_sync.base import Versions as Versions from sqlmesh.core.state_sync.cache import CachingStateSync as CachingStateSync -from sqlmesh.core.state_sync.db import EngineAdapterStateSync as EngineAdapterStateSync +from sqlmesh.core.state_sync.db import \ + EngineAdapterStateSync as EngineAdapterStateSync diff --git a/sqlmesh/core/state_sync/base.py b/sqlmesh/core/state_sync/base.py index 5c35be5ccb..e38f317f82 100644 --- a/sqlmesh/core/state_sync/base.py +++ b/sqlmesh/core/state_sync/base.py @@ -9,31 +9,19 @@ from sqlglot import __version__ as SQLGLOT_VERSION from sqlmesh import migrations -from sqlmesh.core.environment import ( - Environment, - EnvironmentStatements, - EnvironmentSummary, -) -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotId, - SnapshotIdLike, - SnapshotIdAndVersionLike, - SnapshotInfoLike, - SnapshotNameVersion, - SnapshotIdAndVersion, -) +from sqlmesh.core.environment import (Environment, EnvironmentStatements, + EnvironmentSummary) +from sqlmesh.core.snapshot import (Snapshot, SnapshotId, SnapshotIdAndVersion, + SnapshotIdAndVersionLike, SnapshotIdLike, + SnapshotInfoLike, SnapshotNameVersion) from sqlmesh.core.snapshot.definition import Interval, SnapshotIntervals +from sqlmesh.core.state_sync.common import (ExpiredBatchRange, + ExpiredSnapshotBatch, + PromotionResult, StateStream) from sqlmesh.utils import major_minor from sqlmesh.utils.date import TimeLike from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.pydantic import PydanticModel, field_validator -from sqlmesh.core.state_sync.common import ( - StateStream, - ExpiredSnapshotBatch, - PromotionResult, - ExpiredBatchRange, -) logger = logging.getLogger(__name__) @@ -68,7 +56,9 @@ def _schema_version_validator(cls, v: t.Any) -> int: MIN_SQLMESH_VERSION = "0.134.0" MIGRATIONS = [ importlib.import_module(f"sqlmesh.migrations.{migration}") - for migration in sorted(info.name for info in pkgutil.iter_modules(migrations.__path__)) + for migration in sorted( + info.name for info in pkgutil.iter_modules(migrations.__path__) + ) ] # -1 to account for the baseline script SCHEMA_VERSION: int = MIN_SCHEMA_VERSION + len(MIGRATIONS) - 1 @@ -109,7 +99,9 @@ def get_snapshots_by_names( """ @abc.abstractmethod - def snapshots_exist(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> t.Set[SnapshotId]: + def snapshots_exist( + self, snapshot_ids: t.Iterable[SnapshotIdLike] + ) -> t.Set[SnapshotId]: """Checks if multiple snapshots exist in the state sync. Args: @@ -120,7 +112,9 @@ def snapshots_exist(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> t.Set[Sna """ @abc.abstractmethod - def refresh_snapshot_intervals(self, snapshots: t.Collection[Snapshot]) -> t.List[Snapshot]: + def refresh_snapshot_intervals( + self, snapshots: t.Collection[Snapshot] + ) -> t.List[Snapshot]: """Updates given snapshots with latest intervals from the state. Args: @@ -131,7 +125,9 @@ def refresh_snapshot_intervals(self, snapshots: t.Collection[Snapshot]) -> t.Lis """ @abc.abstractmethod - def nodes_exist(self, names: t.Iterable[str], exclude_external: bool = False) -> t.Set[str]: + def nodes_exist( + self, names: t.Iterable[str], exclude_external: bool = False + ) -> t.Set[str]: """Returns the node names that exist in the state sync. Args: @@ -212,7 +208,9 @@ def update_auto_restatements( """ @abc.abstractmethod - def get_environment_statements(self, environment: str) -> t.List[EnvironmentStatements]: + def get_environment_statements( + self, environment: str + ) -> t.List[EnvironmentStatements]: """Fetches environment statements from the environment_statements table. Returns: @@ -262,7 +260,8 @@ def raise_error( SQLMESH_VERSION, versions.sqlmesh_version, remote_package_version=versions.sqlmesh_version, - ahead=major_minor(SQLMESH_VERSION) > major_minor(versions.sqlmesh_version), + ahead=major_minor(SQLMESH_VERSION) + > major_minor(versions.sqlmesh_version), ) if SCHEMA_VERSION != versions.schema_version: @@ -280,7 +279,8 @@ def raise_error( SQLGLOT_VERSION, versions.sqlglot_version, remote_package_version=versions.sqlglot_version, - ahead=major_minor(SQLGLOT_VERSION) > major_minor(versions.sqlglot_version), + ahead=major_minor(SQLGLOT_VERSION) + > major_minor(versions.sqlglot_version), ) return versions @@ -493,7 +493,9 @@ def rollback(self) -> None: """Rollback to previous backed up state.""" @abc.abstractmethod - def add_snapshots_intervals(self, snapshots_intervals: t.Sequence[SnapshotIntervals]) -> None: + def add_snapshots_intervals( + self, snapshots_intervals: t.Sequence[SnapshotIntervals] + ) -> None: """Add snapshot intervals to state Args: @@ -517,7 +519,9 @@ def add_interval( is_dev: Indicates whether the given interval is being added while in development mode last_altered_ts: The timestamp of the last modification of the physical table """ - start_ts, end_ts = snapshot.inclusive_exclusive(start, end, strict=False, expand=False) + start_ts, end_ts = snapshot.inclusive_exclusive( + start, end, strict=False, expand=False + ) if not snapshot.version: raise SQLMeshError("Snapshot version must be set to add an interval.") intervals = [(start_ts, end_ts)] diff --git a/sqlmesh/core/state_sync/cache.py b/sqlmesh/core/state_sync/cache.py index 77f3fc6ba5..70779b111f 100644 --- a/sqlmesh/core/state_sync/cache.py +++ b/sqlmesh/core/state_sync/cache.py @@ -3,13 +3,9 @@ import typing as t from sqlmesh.core.model import SeedModel -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotId, - SnapshotIdLike, - SnapshotIdAndVersionLike, - SnapshotInfoLike, -) +from sqlmesh.core.snapshot import (Snapshot, SnapshotId, + SnapshotIdAndVersionLike, SnapshotIdLike, + SnapshotInfoLike) from sqlmesh.core.snapshot.definition import Interval, SnapshotIntervals from sqlmesh.core.state_sync.base import DelegatingStateSync, StateSync from sqlmesh.core.state_sync.common import ExpiredBatchRange @@ -69,13 +65,17 @@ def get_snapshots( for snapshot_id, snapshot in existing.items(): cached = self._from_cache(snapshot_id, now) - if cached and (not isinstance(cached.node, SeedModel) or cached.node.is_hydrated): + if cached and ( + not isinstance(cached.node, SeedModel) or cached.node.is_hydrated + ): continue self.snapshot_cache[snapshot_id] = (snapshot, expire_at) return existing - def snapshots_exist(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> t.Set[SnapshotId]: + def snapshots_exist( + self, snapshot_ids: t.Iterable[SnapshotIdLike] + ) -> t.Set[SnapshotId]: existing = set() missing = set() now = now_timestamp() @@ -121,7 +121,9 @@ def delete_expired_snapshots( current_ts=current_ts, ) - def add_snapshots_intervals(self, snapshots_intervals: t.Sequence[SnapshotIntervals]) -> None: + def add_snapshots_intervals( + self, snapshots_intervals: t.Sequence[SnapshotIntervals] + ) -> None: for snapshot_intervals in snapshots_intervals: if snapshot_intervals.snapshot_id: self.snapshot_cache.pop(snapshot_intervals.snapshot_id, None) diff --git a/sqlmesh/core/state_sync/common.py b/sqlmesh/core/state_sync/common.py index 6308c0c29d..a0f9b46a44 100644 --- a/sqlmesh/core/state_sync/common.py +++ b/sqlmesh/core/state_sync/common.py @@ -1,27 +1,24 @@ from __future__ import annotations +import abc +import itertools import logging import typing as t -from functools import wraps -import itertools -import abc - from dataclasses import dataclass +from functools import wraps from pydantic_core.core_schema import ValidationInfo from sqlglot import exp -from sqlmesh.utils.pydantic import PydanticModel, field_validator, validation_data -from sqlmesh.core.environment import Environment, EnvironmentStatements, EnvironmentNamingInfo -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotId, - SnapshotTableCleanupTask, - SnapshotTableInfo, -) +from sqlmesh.core.environment import (Environment, EnvironmentNamingInfo, + EnvironmentStatements) +from sqlmesh.core.snapshot import (Snapshot, SnapshotId, + SnapshotTableCleanupTask, SnapshotTableInfo) +from sqlmesh.utils.pydantic import (PydanticModel, field_validator, + validation_data) if t.TYPE_CHECKING: - from sqlmesh.core.state_sync.base import Versions, StateReader + from sqlmesh.core.state_sync.base import StateReader, Versions logger = logging.getLogger(__name__) @@ -46,7 +43,9 @@ def wrapper(self: t.Any, *args: t.Any, **kwargs: t.Any) -> t.Any: T = t.TypeVar("T") -def chunk_iterable(iterable: t.Iterable[T], size: int = 10) -> t.Iterable[t.Iterable[T]]: +def chunk_iterable( + iterable: t.Iterable[T], size: int = 10 +) -> t.Iterable[t.Iterable[T]]: iterator = iter(iterable) for first in iterator: yield itertools.chain([first], itertools.islice(iterator, size - 1)) @@ -171,7 +170,9 @@ def _expanded_tuple_comparison( An expanded OR expression representing the tuple comparison """ if operator not in (exp.GT, exp.GTE, exp.LT, exp.LTE): - raise ValueError(f"Unsupported operator: {operator}. Use GT, GTE, LT, or LTE.") + raise ValueError( + f"Unsupported operator: {operator}. Use GT, GTE, LT, or LTE." + ) # For <= and >=, we use the strict operator for all but the last column # e.g., (a, b) <= (x, y) becomes: a < x OR (a = x AND b <= y) @@ -193,7 +194,9 @@ def _expanded_tuple_comparison( conditions: t.List[exp.Expr] = [] for i in range(len(columns)): # Build equality conditions for all columns before current - equality_conditions = [exp.EQ(this=columns[j], expression=values[j]) for j in range(i)] + equality_conditions = [ + exp.EQ(this=columns[j], expression=values[j]) for j in range(i) + ] # Use the final operator for the last column, strict for others comparison_op = final_operator if i == len(columns) - 1 else strict_operator @@ -204,7 +207,11 @@ def _expanded_tuple_comparison( else: conditions.append(comparison_condition) - return exp.or_(*conditions) if len(conditions) > 1 else t.cast(exp.Condition, conditions[0]) + return ( + exp.or_(*conditions) + if len(conditions) > 1 + else t.cast(exp.Condition, conditions[0]) + ) @property def where_filter(self) -> exp.Condition: @@ -230,7 +237,9 @@ def where_filter(self) -> exp.Condition: exp.Literal.string(self.end.name), exp.Literal.string(self.end.identifier), ] - end_condition = self._expanded_tuple_comparison(columns, end_values, exp.LTE) + end_condition = self._expanded_tuple_comparison( + columns, end_values, exp.LTE + ) range_filter = exp.and_(start_condition, end_condition) else: range_filter = start_condition @@ -270,7 +279,9 @@ def _validate_removed_environment_naming_info( cls, v: t.Optional[EnvironmentNamingInfo], info: ValidationInfo ) -> t.Optional[EnvironmentNamingInfo]: if v and not validation_data(info).get("removed"): - raise ValueError("removed_environment_naming_info must be None if removed is empty") + raise ValueError( + "removed_environment_naming_info must be None if removed is empty" + ) return v @@ -298,7 +309,9 @@ def iter_expired_snapshot_batches( batch_size: Maximum number of snapshots to fetch per batch. """ - batch_size = batch_size if batch_size is not None else EXPIRED_SNAPSHOT_DEFAULT_BATCH_SIZE + batch_size = ( + batch_size if batch_size is not None else EXPIRED_SNAPSHOT_DEFAULT_BATCH_SIZE + ) batch_range = ExpiredBatchRange.init_batch_range(batch_size=batch_size) while True: @@ -313,9 +326,9 @@ def iter_expired_snapshot_batches( yield batch - assert isinstance(batch.batch_range.end, RowBoundary), ( - "Only RowBoundary is supported for pagination currently" - ) + assert isinstance( + batch.batch_range.end, RowBoundary + ), "Only RowBoundary is supported for pagination currently" batch_range = ExpiredBatchRange( start=batch.batch_range.end, end=LimitBoundary(batch_size=batch_size), diff --git a/sqlmesh/core/state_sync/db/environment.py b/sqlmesh/core/state_sync/db/environment.py index 0bb79afcc9..f3e94605c2 100644 --- a/sqlmesh/core/state_sync/db/environment.py +++ b/sqlmesh/core/state_sync/db/environment.py @@ -1,20 +1,19 @@ from __future__ import annotations -import typing as t import json import logging +import typing as t + from sqlglot import exp from sqlmesh.core import constants as c from sqlmesh.core.engine_adapter import EngineAdapter -from sqlmesh.core.state_sync.db.utils import ( - fetchall, - fetchone, -) -from sqlmesh.core.environment import Environment, EnvironmentStatements, EnvironmentSummary -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.core.environment import (Environment, EnvironmentStatements, + EnvironmentSummary) +from sqlmesh.core.state_sync.db.utils import fetchall, fetchone from sqlmesh.utils.date import now_timestamp, time_like_to_str from sqlmesh.utils.errors import SQLMeshError +from sqlmesh.utils.migration import blob_text_type, index_text_type if t.TYPE_CHECKING: import pandas as pd @@ -31,7 +30,9 @@ def __init__( ): self.engine_adapter = engine_adapter self.environments_table = exp.table_("_environments", db=schema) - self.environment_statements_table = exp.table_("_environment_statements", db=schema) + self.environment_statements_table = exp.table_( + "_environment_statements", db=schema + ) index_type = index_text_type(engine_adapter.dialect) blob_type = blob_text_type(engine_adapter.dialect) @@ -107,7 +108,9 @@ def update_environment_statements( if environment_statements: self.engine_adapter.insert_append( self.environment_statements_table, - _environment_statements_to_df(environment_name, plan_id, environment_statements), + _environment_statements_to_df( + environment_name, plan_id, environment_statements + ), target_columns_to_types=self._environment_statements_columns_to_types, track_rows_processed=False, ) @@ -151,7 +154,9 @@ def finalize(self, environment: Environment) -> None: stored_plan_id_row = fetchone(self.engine_adapter, stored_plan_id_query) if not stored_plan_id_row: - raise SQLMeshError(f"Missing environment '{environment.name}' can't be finalized") + raise SQLMeshError( + f"Missing environment '{environment.name}' can't be finalized" + ) stored_plan_id = stored_plan_id_row[0] if stored_plan_id != environment.plan_id: @@ -198,7 +203,9 @@ def delete_expired_environments( A list of deleted environments. """ current_ts = current_ts or now_timestamp() - expired_environments = self.get_expired_environments(current_ts=current_ts, name=name) + expired_environments = self.get_expired_environments( + current_ts=current_ts, name=name + ) self.engine_adapter.delete_from( self.environments_table, @@ -207,7 +214,10 @@ def delete_expired_environments( # Delete the expired environments' corresponding environment statements if expired_environments_exprs := [ - exp.EQ(this=exp.column("environment_name"), expression=exp.Literal.string(env.name)) + exp.EQ( + this=exp.column("environment_name"), + expression=exp.Literal.string(env.name), + ) for env in expired_environments ]: self.engine_adapter.delete_from( @@ -265,7 +275,9 @@ def get_environment( env = self._environment_from_row(row) return env - def get_environment_statements(self, environment: str) -> t.List[EnvironmentStatements]: + def get_environment_statements( + self, environment: str + ) -> t.List[EnvironmentStatements]: """Fetches the environment's statements from the environment_statements table. Args: environment: The environment name @@ -297,12 +309,20 @@ def get_environment_statements(self, environment: str) -> t.List[EnvironmentStat def _environment_from_row(self, row: t.Tuple[str, ...]) -> Environment: return Environment( - **{field: row[i] for i, field in enumerate(sorted(Environment.all_fields()))} + **{ + field: row[i] + for i, field in enumerate(sorted(Environment.all_fields())) + } ) - def _environment_summmary_from_row(self, row: t.Tuple[str, ...]) -> EnvironmentSummary: + def _environment_summmary_from_row( + self, row: t.Tuple[str, ...] + ) -> EnvironmentSummary: return EnvironmentSummary( - **{field: row[i] for i, field in enumerate(sorted(EnvironmentSummary.all_fields()))} + **{ + field: row[i] + for i, field in enumerate(sorted(EnvironmentSummary.all_fields())) + } ) def _environments_query( @@ -311,7 +331,9 @@ def _environments_query( lock_for_update: bool = False, required_fields: t.Optional[t.List[str]] = None, ) -> exp.Select: - query_fields = required_fields if required_fields else sorted(Environment.all_fields()) + query_fields = ( + required_fields if required_fields else sorted(Environment.all_fields()) + ) query = ( exp.select(*(exp.to_identifier(field) for field in query_fields)) .from_(self.environments_table) @@ -362,7 +384,9 @@ def _environment_to_df(environment: Environment) -> pd.DataFrame: "name": environment.name, "snapshots": json.dumps(environment.snapshot_dicts()), "start_at": time_like_to_str(environment.start_at), - "end_at": time_like_to_str(environment.end_at) if environment.end_at else None, + "end_at": ( + time_like_to_str(environment.end_at) if environment.end_at else None + ), "plan_id": environment.plan_id, "previous_plan_id": environment.previous_plan_id, "expiration_ts": environment.expiration_ts, @@ -388,7 +412,9 @@ def _environment_to_df(environment: Environment) -> pd.DataFrame: def _environment_statements_to_df( - environment_name: str, plan_id: str, environment_statements: t.List[EnvironmentStatements] + environment_name: str, + plan_id: str, + environment_statements: t.List[EnvironmentStatements], ) -> pd.DataFrame: import pandas as pd @@ -397,7 +423,9 @@ def _environment_statements_to_df( { "environment_name": environment_name, "plan_id": plan_id, - "environment_statements": json.dumps([e.dict() for e in environment_statements]), + "environment_statements": json.dumps( + [e.dict() for e in environment_statements] + ), } ] ) diff --git a/sqlmesh/core/state_sync/db/facade.py b/sqlmesh/core/state_sync/db/facade.py index 572e54b7f1..225390b0f3 100644 --- a/sqlmesh/core/state_sync/db/facade.py +++ b/sqlmesh/core/state_sync/db/facade.py @@ -19,50 +19,35 @@ import contextlib import logging import typing as t -from pathlib import Path from datetime import datetime - +from pathlib import Path from sqlmesh.core.console import Console, get_console from sqlmesh.core.engine_adapter import EngineAdapter -from sqlmesh.core.environment import Environment, EnvironmentStatements, EnvironmentSummary -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotIdAndVersion, - SnapshotId, - SnapshotIdLike, - SnapshotIdAndVersionLike, - SnapshotInfoLike, - SnapshotIntervals, - SnapshotNameVersion, - SnapshotTableInfo, - start_date, -) -from sqlmesh.core.snapshot.definition import ( - Interval, -) -from sqlmesh.core.state_sync.base import ( - StateSync, - Versions, -) -from sqlmesh.core.state_sync.common import ( - EnvironmentsChunk, - SnapshotsChunk, - VersionsChunk, - transactional, - StateStream, - chunk_iterable, - EnvironmentWithStatements, - ExpiredSnapshotBatch, - PromotionResult, - ExpiredBatchRange, -) -from sqlmesh.core.state_sync.db.interval import IntervalState +from sqlmesh.core.environment import (Environment, EnvironmentStatements, + EnvironmentSummary) +from sqlmesh.core.snapshot import (Snapshot, SnapshotId, SnapshotIdAndVersion, + SnapshotIdAndVersionLike, SnapshotIdLike, + SnapshotInfoLike, SnapshotIntervals, + SnapshotNameVersion, SnapshotTableInfo, + start_date) +from sqlmesh.core.snapshot.definition import Interval +from sqlmesh.core.state_sync.base import StateSync, Versions +from sqlmesh.core.state_sync.common import (EnvironmentsChunk, + EnvironmentWithStatements, + ExpiredBatchRange, + ExpiredSnapshotBatch, + PromotionResult, SnapshotsChunk, + StateStream, VersionsChunk, + chunk_iterable, transactional) from sqlmesh.core.state_sync.db.environment import EnvironmentState +from sqlmesh.core.state_sync.db.interval import IntervalState +from sqlmesh.core.state_sync.db.migrator import (StateMigrator, + _backup_table_name) from sqlmesh.core.state_sync.db.snapshot import SnapshotState from sqlmesh.core.state_sync.db.version import VersionState -from sqlmesh.core.state_sync.db.migrator import StateMigrator, _backup_table_name -from sqlmesh.utils.date import TimeLike, to_timestamp, time_like_to_str, now_timestamp +from sqlmesh.utils.date import (TimeLike, now_timestamp, time_like_to_str, + to_timestamp) from sqlmesh.utils.errors import ConflictingPlanError, SQLMeshError logger = logging.getLogger(__name__) @@ -93,7 +78,9 @@ def __init__( ): self.interval_state = IntervalState(engine_adapter, schema=schema) self.environment_state = EnvironmentState(engine_adapter, schema=schema) - self.snapshot_state = SnapshotState(engine_adapter, schema=schema, cache_dir=cache_dir) + self.snapshot_state = SnapshotState( + engine_adapter, schema=schema, cache_dir=cache_dir + ) self.version_state = VersionState(engine_adapter, schema=schema) self.migrator = StateMigrator( engine_adapter, @@ -179,11 +166,16 @@ def promote( ) existing_table_infos = ( - {table_info.name: table_info for table_info in existing_environment.promoted_snapshots} + { + table_info.name: table_info + for table_info in existing_environment.promoted_snapshots + } if existing_environment else {} ) - table_infos = {table_info.name: table_info for table_info in environment.promoted_snapshots} + table_infos = { + table_info.name: table_info for table_info in environment.promoted_snapshots + } views_that_changed_location: t.Set[SnapshotTableInfo] = set() if existing_environment: views_that_changed_location = { @@ -193,7 +185,9 @@ def promote( and existing_table_info.qualified_view_name.for_environment( existing_environment.naming_info ) - != table_infos[name].qualified_view_name.for_environment(environment.naming_info) + != table_infos[name].qualified_view_name.for_environment( + environment.naming_info + ) } if not existing_environment.expired: if environment.previous_plan_id != existing_environment.plan_id: @@ -208,7 +202,9 @@ def promote( existing_environment, no_gaps_snapshot_names, ) - demoted_snapshots = set(existing_environment.snapshots) - set(environment.snapshots) + demoted_snapshots = set(existing_environment.snapshots) - set( + environment.snapshots + ) # Update the updated_at attribute. self.snapshot_state.touch_snapshots(demoted_snapshots) @@ -217,7 +213,9 @@ def promote( } added_table_infos = set(table_infos.values()) - if existing_environment and environment.can_partially_promote(existing_environment): + if existing_environment and environment.can_partially_promote( + existing_environment + ): # Only promote new snapshots. added_table_infos -= set(existing_environment.promoted_snapshots) @@ -238,7 +236,9 @@ def promote( added=sorted(added_table_infos), removed=list(removed), removed_environment_naming_info=( - existing_environment.naming_info if removed and existing_environment else None + existing_environment.naming_info + if removed and existing_environment + else None ), ) @@ -279,7 +279,9 @@ def get_expired_snapshots( def get_expired_environments( self, current_ts: int, name: t.Optional[str] = None ) -> t.List[EnvironmentSummary]: - return self.environment_state.get_expired_environments(current_ts=current_ts, name=name) + return self.environment_state.get_expired_environments( + current_ts=current_ts, name=name + ) @transactional() def delete_expired_snapshots( @@ -295,22 +297,30 @@ def delete_expired_snapshots( ) if batch and batch.expired_snapshot_ids: self.snapshot_state.delete_snapshots(batch.expired_snapshot_ids) - self.interval_state.cleanup_intervals(batch.cleanup_tasks, batch.expired_snapshot_ids) + self.interval_state.cleanup_intervals( + batch.cleanup_tasks, batch.expired_snapshot_ids + ) @transactional() def delete_expired_environments( self, current_ts: t.Optional[int] = None, name: t.Optional[str] = None ) -> t.List[EnvironmentSummary]: current_ts = current_ts or now_timestamp() - return self.environment_state.delete_expired_environments(current_ts=current_ts, name=name) + return self.environment_state.delete_expired_environments( + current_ts=current_ts, name=name + ) def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: self.snapshot_state.delete_snapshots(snapshot_ids) - def snapshots_exist(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> t.Set[SnapshotId]: + def snapshots_exist( + self, snapshot_ids: t.Iterable[SnapshotIdLike] + ) -> t.Set[SnapshotId]: return self.snapshot_state.snapshots_exist(snapshot_ids) - def nodes_exist(self, names: t.Iterable[str], exclude_external: bool = False) -> t.Set[str]: + def nodes_exist( + self, names: t.Iterable[str], exclude_external: bool = False + ) -> t.Set[str]: return self.snapshot_state.nodes_exist(names, exclude_external) def remove_state(self, including_backup: bool = False) -> None: @@ -343,7 +353,9 @@ def update_auto_restatements( def get_environment(self, environment: str) -> t.Optional[Environment]: return self.environment_state.get_environment(environment) - def get_environment_statements(self, environment: str) -> t.List[EnvironmentStatements]: + def get_environment_statements( + self, environment: str + ) -> t.List[EnvironmentStatements]: return self.environment_state.get_environment_statements(environment) def get_environments(self) -> t.List[Environment]: @@ -386,7 +398,9 @@ def get_snapshots_by_names( exclude_expired: bool = True, ) -> t.Set[SnapshotIdAndVersion]: return self.snapshot_state.get_snapshots_by_names( - snapshot_names=snapshot_names, current_ts=current_ts, exclude_expired=exclude_expired + snapshot_names=snapshot_names, + current_ts=current_ts, + exclude_expired=exclude_expired, ) @transactional() @@ -401,13 +415,17 @@ def add_interval( super().add_interval(snapshot, start, end, is_dev, last_altered_ts) @transactional() - def add_snapshots_intervals(self, snapshots_intervals: t.Sequence[SnapshotIntervals]) -> None: + def add_snapshots_intervals( + self, snapshots_intervals: t.Sequence[SnapshotIntervals] + ) -> None: intervals_to_insert = [] for snapshot_intervals in snapshots_intervals: snapshot_intervals = snapshot_intervals.copy( update={ "intervals": _remove_partial_intervals( - snapshot_intervals.intervals, snapshot_intervals.snapshot_id, is_dev=False + snapshot_intervals.intervals, + snapshot_intervals.snapshot_id, + is_dev=False, ), "dev_intervals": _remove_partial_intervals( snapshot_intervals.dev_intervals, @@ -433,7 +451,9 @@ def remove_intervals( def compact_intervals(self) -> None: self.interval_state.compact_intervals() - def refresh_snapshot_intervals(self, snapshots: t.Collection[Snapshot]) -> t.List[Snapshot]: + def refresh_snapshot_intervals( + self, snapshots: t.Collection[Snapshot] + ) -> t.List[Snapshot]: return self.interval_state.refresh_snapshot_intervals(snapshots) def max_interval_end_per_model( @@ -447,7 +467,9 @@ def max_interval_end_per_model( return {} snapshots = ( - env.snapshots if not ensure_finalized_snapshots else env.finalized_or_current_snapshots + env.snapshots + if not ensure_finalized_snapshots + else env.finalized_or_current_snapshots ) if models is not None: snapshots = [s for s in snapshots if s.name in models] @@ -502,13 +524,16 @@ def export(self, environment_names: t.Optional[t.List[str]] = None) -> StateStre snapshot_ids_to_export |= set([s.snapshot_id for s in env.snapshots or []]) def _export_snapshots() -> t.Iterator[Snapshot]: - for chunk in chunk_iterable(snapshot_ids_to_export, SnapshotState.SNAPSHOT_BATCH_SIZE): + for chunk in chunk_iterable( + snapshot_ids_to_export, SnapshotState.SNAPSHOT_BATCH_SIZE + ): yield from self.get_snapshots(chunk).values() def _export_environments() -> t.Iterator[EnvironmentWithStatements]: for env in selected_environments: yield EnvironmentWithStatements( - environment=env, statements=self.get_environment_statements(env.name) + environment=env, + statements=self.get_environment_statements(env.name), ) return StateStream.from_iterators( @@ -552,7 +577,9 @@ def import_(self, stream: StateStream, clear: bool = True) -> None: self.snapshot_state.push_snapshots( snapshot_chunk, overwrite=overwrite_existing_snapshots ) - self.add_snapshots_intervals((s.snapshot_intervals for s in snapshot_chunk)) + self.add_snapshots_intervals( + (s.snapshot_intervals for s in snapshot_chunk) + ) auto_restatements.update( { @@ -609,7 +636,9 @@ def _ensure_no_gaps( and prev_snapshot.intervals ): start = to_timestamp( - start_date(target_snapshot, target_snapshots_by_name.values(), cache) + start_date( + target_snapshot, target_snapshots_by_name.values(), cache + ) ) end = prev_snapshot.intervals[-1][1] diff --git a/sqlmesh/core/state_sync/db/interval.py b/sqlmesh/core/state_sync/db/interval.py index 8ccdc58fa0..980534ad54 100644 --- a/sqlmesh/core/state_sync/db/interval.py +++ b/sqlmesh/core/state_sync/db/interval.py @@ -1,30 +1,23 @@ from __future__ import annotations -import typing as t import logging +import typing as t from sqlglot import exp from sqlmesh.core.engine_adapter import EngineAdapter -from sqlmesh.core.state_sync.db.utils import ( - snapshot_name_version_filter, - snapshot_id_filter, - create_batches, - fetchall, -) -from sqlmesh.core.snapshot import ( - SnapshotIntervals, - SnapshotIdLike, - SnapshotIdAndVersionLike, - SnapshotNameVersionLike, - SnapshotTableCleanupTask, - SnapshotNameVersion, - Snapshot, -) +from sqlmesh.core.snapshot import (Snapshot, SnapshotIdAndVersionLike, + SnapshotIdLike, SnapshotIntervals, + SnapshotNameVersion, + SnapshotNameVersionLike, + SnapshotTableCleanupTask) from sqlmesh.core.snapshot.definition import Interval -from sqlmesh.utils.migration import index_text_type +from sqlmesh.core.state_sync.db.utils import (create_batches, fetchall, + snapshot_id_filter, + snapshot_name_version_filter) from sqlmesh.utils import random_id from sqlmesh.utils.date import now_timestamp +from sqlmesh.utils.migration import index_text_type if t.TYPE_CHECKING: import pandas as pd @@ -63,7 +56,9 @@ def __init__( "last_altered_ts": exp.DataType.build("bigint"), } - def add_snapshots_intervals(self, snapshots_intervals: t.Sequence[SnapshotIntervals]) -> None: + def add_snapshots_intervals( + self, snapshots_intervals: t.Sequence[SnapshotIntervals] + ) -> None: if snapshots_intervals: self._push_snapshot_intervals(snapshots_intervals) @@ -76,7 +71,9 @@ def remove_intervals( t.Tuple[t.Union[SnapshotIdAndVersionLike, SnapshotIntervals], Interval] ] = snapshot_intervals if remove_shared_versions: - name_version_mapping = {s.name_version: interval for s, interval in snapshot_intervals} + name_version_mapping = { + s.name_version: interval for s, interval in snapshot_intervals + } all_snapshots = [] for where in snapshot_name_version_filter( self.engine_adapter, @@ -125,10 +122,14 @@ def get_snapshot_intervals( return self._get_snapshot_intervals(snapshots)[1] def compact_intervals(self) -> None: - interval_ids, snapshot_intervals = self._get_snapshot_intervals(uncompacted_only=True) + interval_ids, snapshot_intervals = self._get_snapshot_intervals( + uncompacted_only=True + ) logger.info( - "Compacting %s intervals for %s snapshots", len(interval_ids), len(snapshot_intervals) + "Compacting %s intervals for %s snapshots", + len(interval_ids), + len(snapshot_intervals), ) self._push_snapshot_intervals(snapshot_intervals, is_compacted=True) @@ -141,7 +142,9 @@ def compact_intervals(self) -> None: self.intervals_table, exp.column("id").isin(*interval_id_batch) ) - def refresh_snapshot_intervals(self, snapshots: t.Collection[Snapshot]) -> t.List[Snapshot]: + def refresh_snapshot_intervals( + self, snapshots: t.Collection[Snapshot] + ) -> t.List[Snapshot]: if not snapshots: return [] @@ -164,12 +167,17 @@ def max_interval_end_per_model( result: t.Dict[str, int] = {} for where in snapshot_name_version_filter( - self.engine_adapter, snapshots, alias=table_alias, batch_size=self.SNAPSHOT_BATCH_SIZE + self.engine_adapter, + snapshots, + alias=table_alias, + batch_size=self.SNAPSHOT_BATCH_SIZE, ): query = ( exp.select( name_col, - exp.func("MAX", exp.column("end_ts", table=table_alias)).as_("max_end_ts"), + exp.func("MAX", exp.column("end_ts", table=table_alias)).as_( + "max_end_ts" + ), ) .from_(exp.to_table(self.intervals_table).as_(table_alias)) .where(where, copy=False) @@ -312,11 +320,13 @@ def _get_snapshot_intervals( ) if pending_restatement_interval_merge_key not in intervals: - intervals[pending_restatement_interval_merge_key] = SnapshotIntervals( - name=name, - identifier=None, - version=version, - dev_version=None, + intervals[pending_restatement_interval_merge_key] = ( + SnapshotIntervals( + name=name, + identifier=None, + version=version, + dev_version=None, + ) ) if is_removed: @@ -395,7 +405,10 @@ def _update_intervals_for_deleted_snapshots( return for where in snapshot_id_filter( - self.engine_adapter, snapshot_ids, alias=None, batch_size=self.SNAPSHOT_BATCH_SIZE + self.engine_adapter, + snapshot_ids, + alias=None, + batch_size=self.SNAPSHOT_BATCH_SIZE, ): # Nullify the identifier for dev intervals # Set is_compacted to False so that it's compacted during the next compaction @@ -412,7 +425,9 @@ def _update_intervals_for_deleted_snapshots( where=where.and_(exp.column("is_dev").not_()), ) - def _delete_intervals_by_dev_version(self, targets: t.List[SnapshotTableCleanupTask]) -> None: + def _delete_intervals_by_dev_version( + self, targets: t.List[SnapshotTableCleanupTask] + ) -> None: """Deletes dev intervals for snapshot dev versions that are no longer used.""" dev_keys_to_delete = [ SnapshotNameVersion(name=t.snapshot.name, version=t.snapshot.dev_version) @@ -429,9 +444,13 @@ def _delete_intervals_by_dev_version(self, targets: t.List[SnapshotTableCleanupT alias=None, batch_size=self.SNAPSHOT_BATCH_SIZE, ): - self.engine_adapter.delete_from(self.intervals_table, where.and_(exp.column("is_dev"))) + self.engine_adapter.delete_from( + self.intervals_table, where.and_(exp.column("is_dev")) + ) - def _delete_intervals_by_version(self, targets: t.List[SnapshotTableCleanupTask]) -> None: + def _delete_intervals_by_version( + self, targets: t.List[SnapshotTableCleanupTask] + ) -> None: """Deletes intervals for snapshot versions that are no longer used.""" non_dev_keys_to_delete = [t.snapshot for t in targets if not t.dev_table_only] if not non_dev_keys_to_delete: diff --git a/sqlmesh/core/state_sync/db/migrator.py b/sqlmesh/core/state_sync/db/migrator.py index 8d73e1d395..c93c3fcee8 100644 --- a/sqlmesh/core/state_sync/db/migrator.py +++ b/sqlmesh/core/state_sync/db/migrator.py @@ -14,31 +14,18 @@ from sqlmesh.core.console import Console, get_console from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.environment import Environment -from sqlmesh.core.snapshot import ( - Node, - Snapshot, - SnapshotFingerprint, - SnapshotId, - SnapshotTableInfo, - fingerprint_from_node, -) -from sqlmesh.core.snapshot.definition import ( - _parents_from_node, -) -from sqlmesh.core.state_sync.base import ( - MIGRATIONS, - MIN_SCHEMA_VERSION, - MIN_SQLMESH_VERSION, -) +from sqlmesh.core.snapshot import (Node, Snapshot, SnapshotFingerprint, + SnapshotId, SnapshotTableInfo, + fingerprint_from_node) +from sqlmesh.core.snapshot.definition import _parents_from_node +from sqlmesh.core.state_sync.base import (MIGRATIONS, MIN_SCHEMA_VERSION, + MIN_SQLMESH_VERSION) from sqlmesh.core.state_sync.db.environment import EnvironmentState from sqlmesh.core.state_sync.db.interval import IntervalState from sqlmesh.core.state_sync.db.snapshot import SnapshotState +from sqlmesh.core.state_sync.db.utils import (SQLMESH_VERSION, fetchall, + snapshot_id_filter) from sqlmesh.core.state_sync.db.version import VersionState -from sqlmesh.core.state_sync.db.utils import ( - SQLMESH_VERSION, - snapshot_id_filter, - fetchall, -) from sqlmesh.utils import major_minor from sqlmesh.utils.dag import DAG from sqlmesh.utils.date import now_timestamp @@ -95,7 +82,10 @@ def migrate( try: migrate_rows = self._apply_migrations(schema, skip_backup) - if not migrate_rows and major_minor(SQLMESH_VERSION) == versions.minor_sqlmesh_version: + if ( + not migrate_rows + and major_minor(SQLMESH_VERSION) == versions.minor_sqlmesh_version + ): return if migrate_rows: @@ -146,7 +136,9 @@ def rollback(self) -> None: for optional_table in self._optional_state_tables: if self.engine_adapter.table_exists(_backup_table_name(optional_table)): - self._restore_table(optional_table, _backup_table_name(optional_table)) + self._restore_table( + optional_table, _backup_table_name(optional_table) + ) logger.info("Migration rollback successful.") @@ -177,20 +169,29 @@ def _apply_migrations( if not skip_backup and should_backup: self._backup_state() - snapshot_count_before = self.snapshot_state.count() if versions.schema_version else None + snapshot_count_before = ( + self.snapshot_state.count() if versions.schema_version else None + ) - state_table_exist = any(self.engine_adapter.table_exists(t) for t in self._state_tables) + state_table_exist = any( + self.engine_adapter.table_exists(t) for t in self._state_tables + ) for migration in migrations: logger.info(f"Applying migration {migration}") migration.migrate_schemas(engine_adapter=self.engine_adapter, schema=schema) if state_table_exist: # No need to run DML for the initial migration since all tables are empty - migration.migrate_rows(engine_adapter=self.engine_adapter, schema=schema) + migration.migrate_rows( + engine_adapter=self.engine_adapter, schema=schema + ) snapshot_count_after = self.snapshot_state.count() - if snapshot_count_before is not None and snapshot_count_before != snapshot_count_after: + if ( + snapshot_count_before is not None + and snapshot_count_before != snapshot_count_after + ): scripts = f"{versions.schema_version} - {versions.schema_version + len(migrations)}" raise SQLMeshError( f"Number of snapshots before ({snapshot_count_before}) and after " @@ -199,7 +200,8 @@ def _apply_migrations( ) migrate_snapshots_and_environments = ( - bool(migrations) or major_minor(SQLGLOT_VERSION) != versions.minor_sqlglot_version + bool(migrations) + or major_minor(SQLGLOT_VERSION) != versions.minor_sqlglot_version ) return migrate_snapshots_and_environments @@ -258,7 +260,9 @@ def _migrate_snapshot_rows( dag: DAG[SnapshotId] = DAG() for snapshot_id, raw_snapshot in raw_snapshots.items(): - parent_ids = [SnapshotId.parse_obj(p_id) for p_id in raw_snapshot.get("parents", [])] + parent_ids = [ + SnapshotId.parse_obj(p_id) for p_id in raw_snapshot.get("parents", []) + ] dag.add(snapshot_id, [p_id for p_id in parent_ids if p_id in raw_snapshots]) reversed_dag_raw = dag.reversed.graph @@ -281,7 +285,9 @@ def _push_new_snapshots() -> None: existing_new_snapshots = self.snapshot_state.snapshots_exist(new_snapshots) new_snapshots_to_push = [ - s for s in new_snapshots.values() if s.snapshot_id not in existing_new_snapshots + s + for s in new_snapshots.values() + if s.snapshot_id not in existing_new_snapshots ] if new_snapshots_to_push: logger.info("Pushing %s migrated snapshots", len(new_snapshots_to_push)) @@ -332,7 +338,9 @@ def _visit( for parent_node in _parents_from_node(node, nodes).values() ) except Exception: - logger.exception("Could not compute fingerprint for %s", snapshot.snapshot_id) + logger.exception( + "Could not compute fingerprint for %s", snapshot.snapshot_id + ) return # Reset the effective_from date for the new snapshot to avoid unexpected backfills. @@ -432,7 +440,9 @@ def _backup_state(self) -> None: self.engine_adapter.drop_table(backup_name) self.engine_adapter.create_table_like(backup_name, table) self.engine_adapter.insert_append( - backup_name, exp.select("*").from_(table), track_rows_processed=False + backup_name, + exp.select("*").from_(table), + track_rows_processed=False, ) def _restore_table( diff --git a/sqlmesh/core/state_sync/db/snapshot.py b/sqlmesh/core/state_sync/db/snapshot.py index 9b4337b504..6acf2b915b 100644 --- a/sqlmesh/core/state_sync/db/snapshot.py +++ b/sqlmesh/core/state_sync/db/snapshot.py @@ -1,43 +1,32 @@ from __future__ import annotations -import typing as t import json import logging -from pathlib import Path +import typing as t from collections import defaultdict +from pathlib import Path + from sqlglot import exp from sqlmesh.core.engine_adapter import EngineAdapter -from sqlmesh.core.state_sync.db.utils import ( - snapshot_name_filter, - snapshot_name_version_filter, - snapshot_id_filter, - fetchone, - fetchall, -) from sqlmesh.core.environment import Environment -from sqlmesh.core.model import SeedModel, ModelKindName +from sqlmesh.core.model import ModelKindName, SeedModel +from sqlmesh.core.snapshot import (Snapshot, SnapshotFingerprint, SnapshotId, + SnapshotIdAndVersion, SnapshotIdLike, + SnapshotInfoLike, SnapshotNameVersion, + SnapshotNameVersionLike, + SnapshotTableCleanupTask) from sqlmesh.core.snapshot.cache import SnapshotCache -from sqlmesh.core.snapshot import ( - SnapshotIdLike, - SnapshotNameVersionLike, - SnapshotTableCleanupTask, - SnapshotNameVersion, - SnapshotInfoLike, - Snapshot, - SnapshotIdAndVersion, - SnapshotId, - SnapshotFingerprint, -) -from sqlmesh.core.state_sync.common import ( - RowBoundary, - ExpiredSnapshotBatch, - ExpiredBatchRange, - LimitBoundary, -) -from sqlmesh.utils.migration import index_text_type, blob_text_type -from sqlmesh.utils.date import now_timestamp, TimeLike, to_timestamp +from sqlmesh.core.state_sync.common import (ExpiredBatchRange, + ExpiredSnapshotBatch, + LimitBoundary, RowBoundary) +from sqlmesh.core.state_sync.db.utils import (fetchall, fetchone, + snapshot_id_filter, + snapshot_name_filter, + snapshot_name_version_filter) from sqlmesh.utils import unique +from sqlmesh.utils.date import TimeLike, now_timestamp, to_timestamp +from sqlmesh.utils.migration import blob_text_type, index_text_type if t.TYPE_CHECKING: import pandas as pd @@ -84,7 +73,9 @@ def __init__( self._snapshot_cache = SnapshotCache(cache_dir) - def push_snapshots(self, snapshots: t.Iterable[Snapshot], overwrite: bool = False) -> None: + def push_snapshots( + self, snapshots: t.Iterable[Snapshot], overwrite: bool = False + ) -> None: """Pushes snapshots to the state store. Args: @@ -118,9 +109,9 @@ def unpause_snapshots( snapshots: t.Collection[SnapshotInfoLike], unpaused_dt: TimeLike, ) -> None: - unrestorable_snapshots_by_forward_only: t.Dict[bool, t.List[SnapshotNameVersion]] = ( - defaultdict(list) - ) + unrestorable_snapshots_by_forward_only: t.Dict[ + bool, t.List[SnapshotNameVersion] + ] = defaultdict(list) for snapshot in snapshots: # We need to mark all other snapshots that have forward-only opposite to the target snapshot as unrestorable @@ -150,7 +141,10 @@ def unpause_snapshots( ) # Mark unrestorable snapshots - for forward_only, snapshot_name_versions in unrestorable_snapshots_by_forward_only.items(): + for ( + forward_only, + snapshot_name_versions, + ) in unrestorable_snapshots_by_forward_only.items(): forward_only_exp = exp.column("forward_only").is_(exp.convert(forward_only)) for where in snapshot_name_version_filter( self.engine_adapter, @@ -189,7 +183,10 @@ def get_expired_snapshots( environment.snapshots if environment.finalized_ts is not None # If the environment is not finalized, check both the current snapshots and the previous finalized snapshots - else [*environment.snapshots, *(environment.previous_finalized_snapshots or [])] + else [ + *environment.snapshots, + *(environment.previous_finalized_snapshots or []), + ] ) } @@ -261,7 +258,9 @@ def _is_snapshot_used(snapshot: SnapshotIdAndVersion) -> bool: cleanup_targets: t.List[t.Tuple[SnapshotId, bool]] = [] for snapshot in expired_snapshots: - shared_version_snapshots = snapshots_by_version[(snapshot.name, snapshot.version)] + shared_version_snapshots = snapshots_by_version[ + (snapshot.name, snapshot.version) + ] shared_version_snapshots.discard(snapshot.snapshot_id) shared_dev_version_snapshots = snapshots_by_dev_version[ @@ -355,7 +354,9 @@ def get_snapshots_by_names( if exclude_expired: current_ts = current_ts or now_timestamp() - unexpired_expr = (exp.column("updated_ts") + exp.column("ttl_ms")) > current_ts + unexpired_expr = ( + exp.column("updated_ts") + exp.column("ttl_ms") + ) > current_ts else: unexpired_expr = None @@ -375,7 +376,12 @@ def get_snapshots_by_names( for name, identifier, version, kind_name, dev_version, fingerprint in fetchall( self.engine_adapter, exp.select( - "name", "identifier", "version", "kind_name", "dev_version", "fingerprint" + "name", + "identifier", + "version", + "kind_name", + "dev_version", + "fingerprint", ) .from_(self.snapshots_table) .where(where) @@ -383,7 +389,9 @@ def get_snapshots_by_names( ) } - def snapshots_exist(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> t.Set[SnapshotId]: + def snapshots_exist( + self, snapshot_ids: t.Iterable[SnapshotIdLike] + ) -> t.Set[SnapshotId]: """Checks if snapshots exist. Args: @@ -399,11 +407,15 @@ def snapshots_exist(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> t.Set[Sna ) for name, identifier in fetchall( self.engine_adapter, - exp.select("name", "identifier").from_(self.snapshots_table).where(where), + exp.select("name", "identifier") + .from_(self.snapshots_table) + .where(where), ) } - def nodes_exist(self, names: t.Iterable[str], exclude_external: bool = False) -> t.Set[str]: + def nodes_exist( + self, names: t.Iterable[str], exclude_external: bool = False + ) -> t.Set[str]: """Checks if nodes with given names exist. Args: @@ -425,7 +437,9 @@ def nodes_exist(self, names: t.Iterable[str], exclude_external: bool = False) -> .distinct() ) if exclude_external: - query = query.where(exp.column("kind_name").neq(ModelKindName.EXTERNAL.value)) + query = query.where( + exp.column("kind_name").neq(ModelKindName.EXTERNAL.value) + ) return {name for (name,) in fetchall(self.engine_adapter, query)} def update_auto_restatements( @@ -465,7 +479,9 @@ def update_auto_restatements( def count(self) -> int: """Counts the number of snapshots in the state.""" - result = fetchone(self.engine_adapter, exp.select("COUNT(*)").from_(self.snapshots_table)) + result = fetchone( + self.engine_adapter, exp.select("COUNT(*)").from_(self.snapshots_table) + ) return result[0] if result else 0 def clear_cache(self) -> None: @@ -523,7 +539,9 @@ def _get_snapshots( def _loader(snapshot_ids_to_load: t.Set[SnapshotId]) -> t.Collection[Snapshot]: fetched_snapshots: t.Dict[SnapshotId, Snapshot] = {} - for query in self._get_snapshots_expressions(snapshot_ids_to_load, lock_for_update): + for query in self._get_snapshots_expressions( + snapshot_ids_to_load, lock_for_update + ): for ( serialized_snapshot, _, @@ -545,9 +563,13 @@ def _loader(snapshot_ids_to_load: t.Set[SnapshotId]) -> t.Collection[Snapshot]: ) snapshot_id = snapshot.snapshot_id if snapshot_id in fetched_snapshots: - other = duplicates.get(snapshot_id, fetched_snapshots[snapshot_id]) + other = duplicates.get( + snapshot_id, fetched_snapshots[snapshot_id] + ) duplicates[snapshot_id] = ( - snapshot if snapshot.updated_ts > other.updated_ts else other + snapshot + if snapshot.updated_ts > other.updated_ts + else other ) fetched_snapshots[snapshot_id] = duplicates[snapshot_id] else: @@ -561,7 +583,9 @@ def _loader(snapshot_ids_to_load: t.Set[SnapshotId]) -> t.Collection[Snapshot]: if cached_snapshots: cached_snapshots_in_state: t.Set[SnapshotId] = set() for where in snapshot_id_filter( - self.engine_adapter, cached_snapshots, batch_size=self.SNAPSHOT_BATCH_SIZE + self.engine_adapter, + cached_snapshots, + batch_size=self.SNAPSHOT_BATCH_SIZE, ): query = ( exp.select( @@ -575,13 +599,17 @@ def _loader(snapshot_ids_to_load: t.Set[SnapshotId]) -> t.Collection[Snapshot]: ) .from_(exp.to_table(self.snapshots_table).as_("snapshots")) .join( - exp.to_table(self.auto_restatements_table).as_("auto_restatements"), + exp.to_table(self.auto_restatements_table).as_( + "auto_restatements" + ), on=exp.and_( exp.column("name", table="snapshots").eq( exp.column("snapshot_name", table="auto_restatements") ), exp.column("version", table="snapshots").eq( - exp.column("snapshot_version", table="auto_restatements") + exp.column( + "snapshot_version", table="auto_restatements" + ) ), ), join_type="left", @@ -761,7 +789,9 @@ def _snapshots_to_df(snapshots: t.Iterable[Snapshot]) -> pd.DataFrame: "identifier": snapshot.identifier, "version": snapshot.version, "snapshot": _snapshot_to_json(snapshot), - "kind_name": snapshot.model_kind_name.value if snapshot.model_kind_name else None, + "kind_name": ( + snapshot.model_kind_name.value if snapshot.model_kind_name else None + ), "updated_ts": snapshot.updated_ts, "unpaused_ts": snapshot.unpaused_ts, "ttl_ms": snapshot.ttl_ms, @@ -775,7 +805,9 @@ def _snapshots_to_df(snapshots: t.Iterable[Snapshot]) -> pd.DataFrame: ) -def _auto_restatements_to_df(auto_restatements: t.Dict[SnapshotNameVersion, int]) -> pd.DataFrame: +def _auto_restatements_to_df( + auto_restatements: t.Dict[SnapshotNameVersion, int], +) -> pd.DataFrame: import pandas as pd return pd.DataFrame( diff --git a/sqlmesh/core/state_sync/db/utils.py b/sqlmesh/core/state_sync/db/utils.py index b0f321e21f..465102eac3 100644 --- a/sqlmesh/core/state_sync/db/utils.py +++ b/sqlmesh/core/state_sync/db/utils.py @@ -1,13 +1,13 @@ from __future__ import annotations -import typing as t import logging +import typing as t from sqlglot import exp + from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.snapshot import SnapshotIdLike, SnapshotNameVersionLike - logger = logging.getLogger(__name__) try: @@ -123,9 +123,17 @@ def create_batches(l: t.List[T], batch_size: int) -> t.List[t.List[T]]: return [l[i : i + batch_size] for i in range(0, len(l), batch_size)] -def fetchone(engine_adapter: EngineAdapter, query: t.Union[exp.Expr, str]) -> t.Optional[t.Tuple]: - return engine_adapter.fetchone(query, ignore_unsupported_errors=True, quote_identifiers=True) +def fetchone( + engine_adapter: EngineAdapter, query: t.Union[exp.Expr, str] +) -> t.Optional[t.Tuple]: + return engine_adapter.fetchone( + query, ignore_unsupported_errors=True, quote_identifiers=True + ) -def fetchall(engine_adapter: EngineAdapter, query: t.Union[exp.Expr, str]) -> t.List[t.Tuple]: - return engine_adapter.fetchall(query, ignore_unsupported_errors=True, quote_identifiers=True) +def fetchall( + engine_adapter: EngineAdapter, query: t.Union[exp.Expr, str] +) -> t.List[t.Tuple]: + return engine_adapter.fetchall( + query, ignore_unsupported_errors=True, quote_identifiers=True + ) diff --git a/sqlmesh/core/state_sync/db/version.py b/sqlmesh/core/state_sync/db/version.py index c95592bc31..95a8db1684 100644 --- a/sqlmesh/core/state_sync/db/version.py +++ b/sqlmesh/core/state_sync/db/version.py @@ -8,14 +8,8 @@ from sqlglot.helper import seq_get from sqlmesh.core.engine_adapter import EngineAdapter -from sqlmesh.core.state_sync.db.utils import ( - fetchone, - SQLMESH_VERSION, -) -from sqlmesh.core.state_sync.base import ( - SCHEMA_VERSION, - Versions, -) +from sqlmesh.core.state_sync.base import SCHEMA_VERSION, Versions +from sqlmesh.core.state_sync.db.utils import SQLMESH_VERSION, fetchone from sqlmesh.utils.migration import index_text_type logger = logging.getLogger(__name__) @@ -70,5 +64,7 @@ def get_versions(self) -> Versions: return no_version return Versions( - schema_version=row[0], sqlglot_version=row[1], sqlmesh_version=seq_get(row, 2) + schema_version=row[0], + sqlglot_version=row[1], + sqlmesh_version=seq_get(row, 2), ) diff --git a/sqlmesh/core/state_sync/export_import.py b/sqlmesh/core/state_sync/export_import.py index 2461ee50fa..1aca55c992 100644 --- a/sqlmesh/core/state_sync/export_import.py +++ b/sqlmesh/core/state_sync/export_import.py @@ -1,30 +1,25 @@ import json import typing as t - -from sqlmesh.core.state_sync import StateSync -from sqlmesh.core.snapshot import Snapshot -from sqlmesh.utils.date import now, to_tstz -from sqlmesh.utils.pydantic import _expression_encoder -from sqlmesh.core.state_sync import Versions -from sqlmesh.core.state_sync.common import ( - EnvironmentsChunk, - SnapshotsChunk, - VersionsChunk, - EnvironmentWithStatements, - StateStream, -) -from sqlmesh.core.console import Console from pathlib import Path -from sqlmesh.core.console import NoopConsole import json_stream -from json_stream import streamable_dict, to_standard_types, streamable_list -from json_stream.writer import StreamableDict +from json_stream import streamable_dict, streamable_list, to_standard_types from json_stream.base import StreamingJSONObject from json_stream.dump import JSONStreamEncoder -from sqlmesh.utils.errors import SQLMeshError +from json_stream.writer import StreamableDict from sqlglot import exp -from sqlmesh.utils.pydantic import DEFAULT_ARGS as PYDANTIC_DEFAULT_ARGS, PydanticModel + +from sqlmesh.core.console import Console, NoopConsole +from sqlmesh.core.snapshot import Snapshot +from sqlmesh.core.state_sync import StateSync, Versions +from sqlmesh.core.state_sync.common import (EnvironmentsChunk, + EnvironmentWithStatements, + SnapshotsChunk, StateStream, + VersionsChunk) +from sqlmesh.utils.date import now, to_tstz +from sqlmesh.utils.errors import SQLMeshError +from sqlmesh.utils.pydantic import DEFAULT_ARGS as PYDANTIC_DEFAULT_ARGS +from sqlmesh.utils.pydantic import PydanticModel, _expression_encoder class SQLMeshJSONStreamEncoder(JSONStreamEncoder): @@ -40,7 +35,9 @@ def _dump_pydantic_model(model: PydanticModel) -> t.Dict[str, t.Any]: return model.model_dump(mode="json", **dump_args) -def _export(state_stream: StateStream, importable: bool, console: Console) -> StreamableDict: +def _export( + state_stream: StateStream, importable: bool, console: Console +) -> StreamableDict: """ Return the state in a format 'json_stream' can stream to a file @@ -69,7 +66,11 @@ def _dump_environments( @streamable_dict def _do_export() -> t.Iterator[t.Tuple[str, t.Any]]: - yield "metadata", {"timestamp": to_tstz(now()), "file_version": 1, "importable": importable} + yield "metadata", { + "timestamp": to_tstz(now()), + "file_version": 1, + "importable": importable, + } for state_chunk in state_stream: if isinstance(state_chunk, VersionsChunk): @@ -91,7 +92,10 @@ def _do_export() -> t.Iterator[t.Tuple[str, t.Any]]: def _import( - state_sync: StateSync, data: t.Callable[[], StreamingJSONObject], clear: bool, console: Console + state_sync: StateSync, + data: t.Callable[[], StreamingJSONObject], + clear: bool, + console: Console, ) -> None: """ Load the state defined by the :data into the supplied :state_sync. The data is in the same format as written by dump() @@ -133,7 +137,9 @@ def _load_environments() -> t.Iterator[EnvironmentWithStatements]: timestamp = metadata["timestamp"] if not isinstance(timestamp, str): - raise ValueError(f"'timestamp' contains an invalid value. Expecting str, got: {timestamp}") + raise ValueError( + f"'timestamp' contains an invalid value. Expecting str, got: {timestamp}" + ) console.update_state_import_progress( timestamp=timestamp, state_file_version=metadata["file_version"] ) @@ -141,7 +147,9 @@ def _load_environments() -> t.Iterator[EnvironmentWithStatements]: versions = Versions.model_validate(to_standard_types(data()["versions"])) stream = StateStream.from_iterators( - versions=versions, snapshots=_load_snapshots(), environments=_load_environments() + versions=versions, + snapshots=_load_snapshots(), + environments=_load_environments(), ) console.update_state_import_progress(versions=versions) @@ -170,7 +178,9 @@ def export_state( importable = False if local_snapshots else True - json_stream = _export(state_stream=state_stream, importable=importable, console=console) + json_stream = _export( + state_stream=state_stream, importable=importable, console=console + ) with output_file.open(mode="w", encoding="utf8") as fh: json.dump(json_stream, fh, indent=2, cls=SQLMeshJSONStreamEncoder) @@ -199,12 +209,16 @@ def import_state( file_version = metadata.get("file_version") if file_version is None: - raise SQLMeshError("Unable to determine state file format version from the input file") + raise SQLMeshError( + "Unable to determine state file format version from the input file" + ) try: int(file_version) except ValueError: - raise SQLMeshError(f"Unable to parse state file format version: {file_version}") + raise SQLMeshError( + f"Unable to parse state file format version: {file_version}" + ) if not metadata.get("importable", False): # this can happen if the state file was created from local unversioned snapshots that were not sourced from the project state database diff --git a/sqlmesh/core/table_diff.py b/sqlmesh/core/table_diff.py index 97cb0c19ba..343e05456d 100644 --- a/sqlmesh/core/table_diff.py +++ b/sqlmesh/core/table_diff.py @@ -4,18 +4,17 @@ import typing as t from functools import cached_property -from sqlmesh.core.dialect import to_schema -from sqlmesh.core.engine_adapter.mixins import RowDiffMixin -from sqlmesh.core.engine_adapter.athena import AthenaEngineAdapter from sqlglot import exp, parse_one from sqlglot.helper import ensure_list from sqlglot.optimizer.normalize_identifiers import normalize_identifiers from sqlglot.optimizer.qualify_columns import quote_identifiers from sqlglot.optimizer.scope import find_all_in_scope -from sqlmesh.utils.pydantic import PydanticModel +from sqlmesh.core.dialect import to_schema +from sqlmesh.core.engine_adapter.athena import AthenaEngineAdapter +from sqlmesh.core.engine_adapter.mixins import RowDiffMixin from sqlmesh.utils.errors import SQLMeshError - +from sqlmesh.utils.pydantic import PydanticModel if t.TYPE_CHECKING: import pandas as pd @@ -90,7 +89,10 @@ def removed(self) -> t.List[t.Tuple[str, exp.DataType]]: def modified(self) -> t.Dict[str, t.Tuple[exp.DataType, exp.DataType]]: """Columns with modified types.""" modified = {} - for column in self._comparable_source_schema.keys() & self._comparable_target_schema.keys(): + for column in ( + self._comparable_source_schema.keys() + & self._comparable_target_schema.keys() + ): source_type = self._comparable_source_schema[column] target_type = self._comparable_target_schema[column] @@ -99,7 +101,8 @@ def modified(self) -> t.Dict[str, t.Tuple[exp.DataType, exp.DataType]]: if self.ignore_case: modified = { - self._original_column_name(c, self.source_schema): dt for c, dt in modified.items() + self._original_column_name(c, self.source_schema): dt + for c, dt in modified.items() } return modified @@ -276,7 +279,9 @@ def target_schema(self) -> t.Dict[str, exp.DataType]: return self.adapter.columns(self.target_table) @cached_property - def key_columns(self) -> t.Tuple[t.List[exp.Column], t.List[exp.Column], t.List[str]]: + def key_columns( + self, + ) -> t.Tuple[t.List[exp.Column], t.List[exp.Column], t.List[str]]: dialect = self.model_dialect or self.dialect # If the columns to join on are explicitly specified, then just return them @@ -343,10 +348,14 @@ def row_diff( ) -> RowDiff: if self._row_diff is None: source_schema = { - c: t for c, t in self.source_schema.items() if c not in self.skip_columns + c: t + for c, t in self.source_schema.items() + if c not in self.skip_columns } target_schema = { - c: t for c, t in self.target_schema.items() if c not in self.skip_columns + c: t + for c, t in self.target_schema.items() + if c not in self.skip_columns } s_selects = {c: exp.column(c, "s").as_(f"s__{c}") for c in source_schema} @@ -358,7 +367,9 @@ def row_diff( f"t__{SQLMESH_JOIN_KEY_COL}" ) - matched_columns = {c: t for c, t in source_schema.items() if t == target_schema.get(c)} + matched_columns = { + c: t for c, t in source_schema.items() if t == target_schema.get(c) + } s_index, t_index, index_cols = self.key_columns s_index_names = [c.name for c in s_index] @@ -369,7 +380,9 @@ def _column_expr(name: str, table: str) -> exp.Expr: qualified_column = exp.column(name, table) if column_type.is_type(*exp.DataType.REAL_TYPES): - return self.adapter._normalize_decimal_value(qualified_column, self.decimals) + return self.adapter._normalize_decimal_value( + qualified_column, self.decimals + ) if column_type.is_type(*exp.DataType.NESTED_TYPES): return self.adapter._normalize_nested_value(qualified_column) @@ -377,13 +390,17 @@ def _column_expr(name: str, table: str) -> exp.Expr: comparisons = [ exp.Case() - .when(_column_expr(c, "s").eq(_column_expr(c, "t")), exp.Literal.number(1)) .when( - exp.column(c, "s").is_(exp.Null()) & exp.column(c, "t").is_(exp.Null()), + _column_expr(c, "s").eq(_column_expr(c, "t")), exp.Literal.number(1) + ) + .when( + exp.column(c, "s").is_(exp.Null()) + & exp.column(c, "t").is_(exp.Null()), exp.Literal.number(1), ) .when( - exp.column(c, "s").is_(exp.Null()) | exp.column(c, "t").is_(exp.Null()), + exp.column(c, "s").is_(exp.Null()) + | exp.column(c, "t").is_(exp.Null()), exp.Literal.number(0), ) .else_(exp.Literal.number(0)) @@ -423,10 +440,16 @@ def _column_expr(name: str, table: str) -> exp.Expr: *s_selects.values(), *t_selects.values(), exp.func( - "IF", exp.column(SQLMESH_JOIN_KEY_COL, "s").is_(exp.Null()).not_(), 1, 0 + "IF", + exp.column(SQLMESH_JOIN_KEY_COL, "s").is_(exp.Null()).not_(), + 1, + 0, ).as_("s_exists"), exp.func( - "IF", exp.column(SQLMESH_JOIN_KEY_COL, "t").is_(exp.Null()).not_(), 1, 0 + "IF", + exp.column(SQLMESH_JOIN_KEY_COL, "t").is_(exp.Null()).not_(), + 1, + 0, ).as_("t_exists"), exp.func( "IF", @@ -486,14 +509,18 @@ def _column_expr(name: str, table: str) -> exp.Expr: ) query = self.adapter.ensure_nulls_for_unmatched_after_join( - quote_identifiers(base_query.copy(), dialect=self.model_dialect or self.dialect) + quote_identifiers( + base_query.copy(), dialect=self.model_dialect or self.dialect + ) ) if not temp_schema: temp_schema = "sqlmesh_temp" schema = to_schema(temp_schema, dialect=self.dialect) - temp_table = exp.table_("diff", db=schema.db, catalog=schema.catalog, quoted=True) + temp_table = exp.table_( + "diff", db=schema.db, catalog=schema.catalog, quoted=True + ) temp_table_kwargs: t.Dict[str, t.Any] = {} if isinstance(self.adapter, AthenaEngineAdapter): @@ -513,7 +540,10 @@ def _column_expr(name: str, table: str) -> exp.Expr: ) with self.adapter.temp_table( - query, name=temp_table, target_columns_to_types=None, **temp_table_kwargs + query, + name=temp_table, + target_columns_to_types=None, + **temp_table_kwargs, ) as table: summary_sums = [ exp.func("SUM", "s_exists").as_("s_count"), @@ -527,18 +557,20 @@ def _column_expr(name: str, table: str) -> exp.Expr: if not skip_grain_check: summary_sums.extend( [ - parse_one(f"COUNT(DISTINCT(s__{SQLMESH_JOIN_KEY_COL}))").as_( - "distinct_count_s" - ), - parse_one(f"COUNT(DISTINCT(t__{SQLMESH_JOIN_KEY_COL}))").as_( - "distinct_count_t" - ), + parse_one( + f"COUNT(DISTINCT(s__{SQLMESH_JOIN_KEY_COL}))" + ).as_("distinct_count_s"), + parse_one( + f"COUNT(DISTINCT(t__{SQLMESH_JOIN_KEY_COL}))" + ).as_("distinct_count_t"), ] ) summary_query = exp.select(*summary_sums).from_(table) - stats_df = self.adapter.fetchdf(summary_query, quote_identifiers=True).fillna(0) + stats_df = self.adapter.fetchdf( + summary_query, quote_identifiers=True + ).fillna(0) stats_df["s_only_count"] = stats_df["s_count"] - stats_df["join_count"] stats_df["t_only_count"] = stats_df["t_count"] - stats_df["join_count"] stats = stats_df.iloc[0].to_dict() @@ -551,7 +583,8 @@ def _column_expr(name: str, table: str) -> exp.Expr: 100 * ( exp.cast( - exp.func("SUM", name(c)), exp.DataType.build("NUMERIC") + exp.func("SUM", name(c)), + exp.DataType.build("NUMERIC"), ) / exp.func("COUNT", name(c)) ), @@ -565,8 +598,9 @@ def _column_expr(name: str, table: str) -> exp.Expr: ) column_stats = ( - self.adapter.fetchdf(column_stats_query, quote_identifiers=True) - .T.rename( + self.adapter.fetchdf( + column_stats_query, quote_identifiers=True + ).T.rename( columns={0: "pct_match"}, index=lambda x: str(x).replace("_matches", "") if x else "", ) @@ -622,9 +656,9 @@ def _column_expr(name: str, table: str) -> exp.Expr: for c, n in joined_renamed_cols.items() } - joined_sample = sample[sample[SQLMESH_SAMPLE_TYPE_COL] == "common_rows"][ - joined_sample_cols - ] + joined_sample = sample[ + sample[SQLMESH_SAMPLE_TYPE_COL] == "common_rows" + ][joined_sample_cols] joined_sample.rename( columns=joined_renamed_cols, inplace=True, @@ -637,7 +671,8 @@ def _column_expr(name: str, table: str) -> exp.Expr: ] ] s_sample.rename( - columns={c: c.replace("s__", "") for c in s_sample.columns}, inplace=True + columns={c: c.replace("s__", "") for c in s_sample.columns}, + inplace=True, ) t_sample = sample[sample[SQLMESH_SAMPLE_TYPE_COL] == "target_only"][ @@ -647,7 +682,8 @@ def _column_expr(name: str, table: str) -> exp.Expr: ] ] t_sample.rename( - columns={c: c.replace("t__", "") for c in t_sample.columns}, inplace=True + columns={c: c.replace("t__", "") for c in t_sample.columns}, + inplace=True, ) sample.drop( @@ -692,30 +728,41 @@ def _fetch_sample( source_only_sample = ( exp.select( - exp.Literal.string("source_only").as_(sample_type), *rendered_data_column_names + exp.Literal.string("source_only").as_(sample_type), + *rendered_data_column_names, ) .from_(sample_table) - .where(exp.and_(exp.column("s_exists").eq(1), exp.column("row_joined").eq(0))) + .where( + exp.and_(exp.column("s_exists").eq(1), exp.column("row_joined").eq(0)) + ) .order_by(*(name(s_selects[c.name]) for c in s_index)) .limit(limit) ) target_only_sample = ( exp.select( - exp.Literal.string("target_only").as_(sample_type), *rendered_data_column_names + exp.Literal.string("target_only").as_(sample_type), + *rendered_data_column_names, ) .from_(sample_table) - .where(exp.and_(exp.column("t_exists").eq(1), exp.column("row_joined").eq(0))) + .where( + exp.and_(exp.column("t_exists").eq(1), exp.column("row_joined").eq(0)) + ) .order_by(*(name(t_selects[c.name]) for c in t_index)) .limit(limit) ) common_rows_sample = ( exp.select( - exp.Literal.string("common_rows").as_(sample_type), *rendered_data_column_names + exp.Literal.string("common_rows").as_(sample_type), + *rendered_data_column_names, ) .from_(sample_table) - .where(exp.and_(exp.column("row_joined").eq(1), exp.column("row_full_match").eq(0))) + .where( + exp.and_( + exp.column("row_joined").eq(1), exp.column("row_full_match").eq(0) + ) + ) .order_by( *(name(s_selects[c.name]) for c in s_index), *(name(t_selects[c.name]) for c in t_index), @@ -731,11 +778,15 @@ def _fetch_sample( .select(sample_type, *rendered_data_column_names) .from_("source_only") .union( - exp.select(sample_type, *rendered_data_column_names).from_("target_only"), + exp.select(sample_type, *rendered_data_column_names).from_( + "target_only" + ), distinct=False, ) .union( - exp.select(sample_type, *rendered_data_column_names).from_("common_rows"), + exp.select(sample_type, *rendered_data_column_names).from_( + "common_rows" + ), distinct=False, ) ) diff --git a/sqlmesh/core/test/__init__.py b/sqlmesh/core/test/__init__.py index 6353370f45..80c4c9520a 100644 --- a/sqlmesh/core/test/__init__.py +++ b/sqlmesh/core/test/__init__.py @@ -1,9 +1,9 @@ from __future__ import annotations -from sqlmesh.core.test.definition import ModelTest as ModelTest, generate_test as generate_test -from sqlmesh.core.test.discovery import ( - ModelTestMetadata as ModelTestMetadata, - filter_tests_by_patterns as filter_tests_by_patterns, -) +from sqlmesh.core.test.definition import ModelTest as ModelTest +from sqlmesh.core.test.definition import generate_test as generate_test +from sqlmesh.core.test.discovery import ModelTestMetadata as ModelTestMetadata +from sqlmesh.core.test.discovery import \ + filter_tests_by_patterns as filter_tests_by_patterns from sqlmesh.core.test.result import ModelTextTestResult as ModelTextTestResult from sqlmesh.core.test.runner import run_tests as run_tests diff --git a/sqlmesh/core/test/definition.py b/sqlmesh/core/test/definition.py index 136de947a3..8fb8ae4673 100644 --- a/sqlmesh/core/test/definition.py +++ b/sqlmesh/core/test/definition.py @@ -1,19 +1,17 @@ from __future__ import annotations -import sys - import datetime +import sys import threading import typing as t import unittest from collections import Counter -from contextlib import nullcontext, contextmanager, AbstractContextManager +from contextlib import AbstractContextManager, contextmanager, nullcontext +from io import StringIO from itertools import chain from pathlib import Path from unittest.mock import patch - -from io import StringIO from sqlglot import Dialect, exp from sqlglot.optimizer.annotate_types import annotate_types from sqlglot.optimizer.normalize_identifiers import normalize_identifiers @@ -23,16 +21,16 @@ from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.macros import RuntimeStage from sqlmesh.core.model import Model, PythonModel, SqlModel -from sqlmesh.utils import UniqueKeyDict, random_id, type_is_known, yaml -from sqlmesh.utils.date import date_dict, pandas_timestamp_to_pydatetime, to_datetime +from sqlmesh.utils import (UniqueKeyDict, Verbosity, random_id, type_is_known, + yaml) +from sqlmesh.utils.date import (date_dict, pandas_timestamp_to_pydatetime, + to_datetime) from sqlmesh.utils.errors import ConfigError, TestError -from sqlmesh.utils.yaml import load as yaml_load -from sqlmesh.utils import Verbosity from sqlmesh.utils.rich import df_to_table +from sqlmesh.utils.yaml import load as yaml_load if t.TYPE_CHECKING: import pandas as pd - from sqlglot.dialects.dialect import DialectType Row = t.Dict[str, t.Any] @@ -44,7 +42,9 @@ "execution_time", "latest", # all built-in datetime macro var names - *date_dict(execution_time="1970-01-01", start="1970-01-01", end="1970-01-01").keys(), + *date_dict( + execution_time="1970-01-01", start="1970-01-01", end="1970-01-01" + ).keys(), } @@ -102,7 +102,8 @@ def __init__( if self.engine_adapter.default_catalog: self._fixture_catalog: t.Optional[exp.Identifier] = normalize_identifiers( exp.parse_identifier( - self.engine_adapter.default_catalog, dialect=self._test_adapter_dialect + self.engine_adapter.default_catalog, + dialect=self._test_adapter_dialect, ), dialect=self._test_adapter_dialect, ) @@ -114,10 +115,14 @@ def __init__( self._fixture_schema = exp.parse_identifier( self.body.get("schema") or f"sqlmesh_test_{random_id(short=True)}" ) - self._qualified_fixture_schema = schema_(self._fixture_schema, self._fixture_catalog) + self._qualified_fixture_schema = schema_( + self._fixture_schema, self._fixture_catalog + ) self._transforms = self._test_adapter_dialect.generator_class.TRANSFORMS - self._execution_time = str(self.body.get("vars", {}).get("execution_time") or "") + self._execution_time = str( + self.body.get("vars", {}).get("execution_time") or "" + ) if self._execution_time: # Normalizes the execution time by converting it into UTC timezone @@ -147,15 +152,17 @@ def __init__( def defaultTestResult(self) -> unittest.TestResult: from sqlmesh.core.test.result import ModelTextTestResult - return ModelTextTestResult(stream=sys.stdout, descriptions=True, verbosity=self.verbosity) + return ModelTextTestResult( + stream=sys.stdout, descriptions=True, verbosity=self.verbosity + ) def shortDescription(self) -> t.Optional[str]: return self.body.get("description") def setUp(self) -> None: """Load all input tables""" - import pandas as pd import numpy as np + import pandas as pd self.engine_adapter.create_schema(self._qualified_fixture_schema) @@ -167,7 +174,9 @@ def setUp(self) -> None: if model: inferred_columns_to_types = model.columns_to_types or {} columns_to_known_types = { - c: t for c, t in inferred_columns_to_types.items() if type_is_known(t) + c: t + for c, t in inferred_columns_to_types.items() + if type_is_known(t) } all_types_are_known = bool(inferred_columns_to_types) and ( len(columns_to_known_types) == len(inferred_columns_to_types) @@ -180,9 +189,14 @@ def setUp(self) -> None: if not all_types_are_known and rows: for col, value in rows[0].items(): if col not in columns_to_known_types: - v_type = annotate_types(exp.convert(value)).type or type(value).__name__ + v_type = ( + annotate_types(exp.convert(value)).type + or type(value).__name__ + ) v_type = exp.maybe_parse( - v_type, into=exp.DataType, dialect=self._test_adapter_dialect + v_type, + into=exp.DataType, + dialect=self._test_adapter_dialect, ) if not type_is_known(v_type): @@ -202,7 +216,8 @@ def setUp(self) -> None: ) if columns_to_known_types: columns_to_known_types = { - col: columns_to_known_types[col] for col in query_or_df.named_selects + col: columns_to_known_types[col] + for col in query_or_df.named_selects } else: query_or_df = self._create_df(values, columns=columns_to_known_types) @@ -218,7 +233,9 @@ def setUp(self) -> None: def tearDown(self) -> None: """Drop all fixture tables.""" if not self.preserve_fixtures: - self.engine_adapter.drop_schema(self._qualified_fixture_schema, cascade=True) + self.engine_adapter.drop_schema( + self._qualified_fixture_schema, cascade=True + ) def assert_equal( self, @@ -294,14 +311,22 @@ def _to_hashable(x: t.Any) -> t.Any: return tuple(_to_hashable(v) for v in x) if isinstance(x, dict): return tuple((k, _to_hashable(v)) for k, v in x.items()) - return str(x) if isinstance(x, DATETIME_TYPES) or not isinstance(x, t.Hashable) else x + return ( + str(x) + if isinstance(x, DATETIME_TYPES) or not isinstance(x, t.Hashable) + else x + ) actual = actual.apply(lambda col: col.map(_to_hashable)) expected = expected.apply(lambda col: col.map(_to_hashable)) if sort: - actual = actual.sort_values(by=actual.columns.to_list()).reset_index(drop=True) - expected = expected.sort_values(by=expected.columns.to_list()).reset_index(drop=True) + actual = actual.sort_values(by=actual.columns.to_list()).reset_index( + drop=True + ) + expected = expected.sort_values(by=expected.columns.to_list()).reset_index( + drop=True + ) try: pd.testing.assert_frame_equal( @@ -334,16 +359,22 @@ def _to_hashable(x: t.Any) -> t.Any: missing_rows = _row_difference(expected, actual) if not missing_rows.empty: args[0] += f"\n\nMissing rows:\n\n{missing_rows}" - args.append(df_to_table(f"Missing rows{failed_subtest}", missing_rows)) + args.append( + df_to_table(f"Missing rows{failed_subtest}", missing_rows) + ) unexpected_rows = _row_difference(actual, expected) if not unexpected_rows.empty: args[0] += f"\n\nUnexpected rows:\n\n{unexpected_rows}" - args.append(df_to_table(f"Unexpected rows{failed_subtest}", unexpected_rows)) + args.append( + df_to_table(f"Unexpected rows{failed_subtest}", unexpected_rows) + ) else: - diff = expected.compare(actual).rename(columns={"self": "exp", "other": "act"}) + diff = expected.compare(actual).rename( + columns={"self": "exp", "other": "act"} + ) args.append(f"Data mismatch (exp: expected, act: actual)\n\n{diff}") @@ -406,7 +437,9 @@ def create_test( if name is None: _raise_error("Missing required 'model' field", path) - name = normalize_model_name(name, default_catalog=default_catalog, dialect=dialect) + name = normalize_model_name( + name, default_catalog=default_catalog, dialect=dialect + ) model = models.get(name) if not model: from sqlmesh.core.console import get_console @@ -421,7 +454,9 @@ def create_test( elif isinstance(model, PythonModel): test_type = PythonModelTest else: - _raise_error(f"Model '{name}' is an unsupported model type for testing", path) + _raise_error( + f"Model '{name}' is an unsupported model type for testing", path + ) try: return test_type( @@ -455,7 +490,9 @@ def _validate_and_normalize_test(self) -> None: partial = outputs.pop("partial", None) if ctes is None and query is None: - _raise_error("Incomplete test, outputs must contain 'query' or 'ctes'", self.path) + _raise_error( + "Incomplete test, outputs must contain 'query' or 'ctes'", self.path + ) def _normalize_rows( values: t.List[Row] | t.Dict, @@ -475,34 +512,47 @@ def _normalize_rows( path = values.get("path") if fmt == "csv": csv_settings = values.get("csv_settings") or {} - rows = pd.read_csv(path or StringIO(rows), **csv_settings).to_dict(orient="records") + rows = pd.read_csv(path or StringIO(rows), **csv_settings).to_dict( + orient="records" + ) elif fmt in (None, "yaml"): if path: input_rows = yaml_load(Path(path)) - rows = input_rows.get("rows") if isinstance(input_rows, dict) else input_rows + rows = ( + input_rows.get("rows") + if isinstance(input_rows, dict) + else input_rows + ) else: _raise_error(f"Unsupported data format '{fmt}' for '{name}'", self.path) if query is not None: if rows is not None: _raise_error( - f"Invalid test, cannot set both 'query' and 'rows' for '{name}'", self.path + f"Invalid test, cannot set both 'query' and 'rows' for '{name}'", + self.path, ) # We parse the user-supplied query using the testing adapter dialect, but we # normalize its identifiers according to the model's dialect, so that, e.g., # the projection names match those in its `columns_to_types` field values["query"] = normalize_identifiers( - exp.maybe_parse(query, dialect=self._test_adapter_dialect), dialect=dialect + exp.maybe_parse(query, dialect=self._test_adapter_dialect), + dialect=dialect, ) return values if rows is None: - _raise_error(f"Incomplete test, missing row data for '{name}'", self.path) + _raise_error( + f"Incomplete test, missing row data for '{name}'", self.path + ) assert isinstance(rows, list) values["rows"] = [ - {self._normalize_column_name(column): value for column, value in row.items()} + { + self._normalize_column_name(column): value + for column, value in row.items() + } for row in rows ] if partial: @@ -552,7 +602,10 @@ def _normalize_sources( for depends_on in self.model.depends_on: if depends_on not in inputs: - _raise_error(f"Incomplete test, missing input model '{depends_on}'", self.path) + _raise_error( + f"Incomplete test, missing input model '{depends_on}'", + self.path, + ) if self.model.depends_on_self and normalized_model_name not in inputs: inputs[normalized_model_name] = {"rows": []} @@ -560,7 +613,9 @@ def _normalize_sources( self.body["inputs"] = inputs if ctes: - outputs["ctes"] = _normalize_sources(ctes, partial=partial, with_default_catalog=False) + outputs["ctes"] = _normalize_sources( + ctes, partial=partial, with_default_catalog=False + ) if query or query == []: outputs["query"] = _normalize_rows( @@ -583,14 +638,20 @@ def _test_fixture_table(self, name: str) -> exp.Table: return table - def _normalize_model_name(self, name: str, with_default_catalog: bool = True) -> str: - normalized_name = self._normalized_model_name_cache.get((name, with_default_catalog)) + def _normalize_model_name( + self, name: str, with_default_catalog: bool = True + ) -> str: + normalized_name = self._normalized_model_name_cache.get( + (name, with_default_catalog) + ) if normalized_name is None: default_catalog = self.default_catalog if with_default_catalog else None normalized_name = normalize_model_name( name, default_catalog=default_catalog, dialect=self.dialect ) - self._normalized_model_name_cache[(name, with_default_catalog)] = normalized_name + self._normalized_model_name_cache[(name, with_default_catalog)] = ( + normalized_name + ) return normalized_name @@ -686,7 +747,8 @@ def test_ctes(self, ctes: t.Dict[str, exp.Expr], recursive: bool = False) -> Non with self.subTest(cte=cte_name): if cte_name not in ctes: _raise_error( - f"No CTE named {cte_name} found in model {self.model.name}", self.path + f"No CTE named {cte_name} found in model {self.model.name}", + self.path, ) cte_query = ctes[cte_name].this @@ -694,7 +756,9 @@ def test_ctes(self, ctes: t.Dict[str, exp.Expr], recursive: bool = False) -> Non sort = cte_query.args.get("order") is None partial = values.get("partial") - cte_query = exp.select(*_projection_identifiers(cte_query)).from_(cte_name) + cte_query = exp.select(*_projection_identifiers(cte_query)).from_( + cte_name + ) for alias, cte in ctes.items(): cte_query = cte_query.with_(alias, cte.this, recursive=recursive) @@ -702,11 +766,14 @@ def test_ctes(self, ctes: t.Dict[str, exp.Expr], recursive: bool = False) -> Non # Similar to the model's query, we render the CTE query under the locked context # so that the execution (fetchdf) can continue concurrently between the threads sql = cte_query.sql( - self._test_adapter_dialect, pretty=self.engine_adapter._pretty_sql + self._test_adapter_dialect, + pretty=self.engine_adapter._pretty_sql, ) actual = self._execute(sql) - expected = self._create_df(values, columns=cte_query.named_selects, partial=partial) + expected = self._create_df( + values, columns=cte_query.named_selects, partial=partial + ) self.assert_equal(expected, actual, sort=sort, partial=partial) @@ -715,14 +782,18 @@ def runTest(self) -> None: # Render the model's query and generate the SQL under the locked context so that # execution (fetchdf) can continue concurrently between the threads query = self._render_model_query() - sql = query.sql(self._test_adapter_dialect, pretty=self.engine_adapter._pretty_sql) + sql = query.sql( + self._test_adapter_dialect, pretty=self.engine_adapter._pretty_sql + ) with_clause = query.args.get("with_") if with_clause: self.test_ctes( { - self._normalize_model_name(cte.alias, with_default_catalog=False): cte + self._normalize_model_name( + cte.alias, with_default_catalog=False + ): cte for cte in query.ctes }, recursive=with_clause.recursive, @@ -734,20 +805,25 @@ def runTest(self) -> None: sort = query.args.get("order") is None actual = self._execute(sql) - expected = self._create_df(values, columns=self.model.columns_to_types, partial=partial) + expected = self._create_df( + values, columns=self.model.columns_to_types, partial=partial + ) self.assert_equal(expected, actual, sort=sort, partial=partial) def _render_model_query(self) -> exp.Query: variables = self.body.get("vars", {}).copy() - time_kwargs = {key: variables.pop(key) for key in TIME_KWARG_KEYS if key in variables} + time_kwargs = { + key: variables.pop(key) for key in TIME_KWARG_KEYS if key in variables + } query = self.model.render_query_or_raise( **time_kwargs, variables=variables, engine_adapter=self.engine_adapter, table_mapping={ - name: self._test_fixture_table(name).sql() for name in self.body.get("inputs", {}) + name: self._test_fixture_table(name).sql() + for name in self.body.get("inputs", {}) }, runtime_stage=RuntimeStage.TESTING, ) @@ -812,7 +888,9 @@ def runTest(self) -> None: actual_df = self._execute_model() actual_df.reset_index(drop=True, inplace=True) - expected = self._create_df(values, columns=self.model.columns_to_types, partial=partial) + expected = self._create_df( + values, columns=self.model.columns_to_types, partial=partial + ) self.assert_equal(expected, actual_df, sort=True, partial=partial) @@ -822,8 +900,14 @@ def _execute_model(self) -> pd.DataFrame: with self._concurrent_render_context(): variables = self.body.get("vars", {}).copy() - time_kwargs = {key: variables.pop(key) for key in TIME_KWARG_KEYS if key in variables} - df = next(self.model.render(context=self.context, variables=variables, **time_kwargs)) + time_kwargs = { + key: variables.pop(key) for key in TIME_KWARG_KEYS if key in variables + } + df = next( + self.model.render( + context=self.context, variables=variables, **time_kwargs + ) + ) assert not isinstance(df, exp.Expr) return df if isinstance(df, pd.DataFrame) else df.toPandas() @@ -880,7 +964,9 @@ def generate_test( # datetime or datetime.date objects based on column type inputs = { dep: pandas_timestamp_to_pydatetime( - engine_adapter.fetchdf(query).apply(lambda col: col.map(_normalize_df_value)), + engine_adapter.fetchdf(query).apply( + lambda col: col.map(_normalize_df_value) + ), models[dep].columns_to_types, ) .replace({np.nan: None}) @@ -922,7 +1008,9 @@ def generate_test( cte_query = cte.this cte_identifier = cte.args["alias"].this - cte_query = exp.select(*_projection_identifiers(cte_query)).from_(cte_identifier) + cte_query = exp.select(*_projection_identifiers(cte_query)).from_( + cte_identifier + ) for prev in chain(previous_ctes, [cte]): cte_query = cte_query.with_( @@ -949,7 +1037,8 @@ def generate_test( outputs["query"] = ( pandas_timestamp_to_pydatetime( - output.apply(lambda col: col.map(_normalize_df_value)), model.columns_to_types + output.apply(lambda col: col.map(_normalize_df_value)), + model.columns_to_types, ) .replace({np.nan: None}) .to_dict(orient="records") @@ -1024,12 +1113,16 @@ def _normalize_df_value(value: t.Any) -> t.Any: if "key" in value and "value" in value: # Maps returned by DuckDB look like: {'key': ['key1', 'key2'], 'value': [10, 20]} # so we convert to {'key1': 10, 'key2': 20} (TODO: handle more dialects here) - return {k: _normalize_df_value(v) for k, v in zip(value["key"], value["value"])} + return { + k: _normalize_df_value(v) for k, v in zip(value["key"], value["value"]) + } return {k: _normalize_df_value(v) for k, v in value.items()} return value -def _split_df_by_column_pairs(df: pd.DataFrame, pairs_per_chunk: int = 4) -> t.List[pd.DataFrame]: +def _split_df_by_column_pairs( + df: pd.DataFrame, pairs_per_chunk: int = 4 +) -> t.List[pd.DataFrame]: """Split a dataframe into chunks of column pairs. Args: @@ -1050,7 +1143,9 @@ def _split_df_by_column_pairs(df: pd.DataFrame, pairs_per_chunk: int = 4) -> t.L # Calculate columns per chunk to ensure equal distribution # We round down to nearest even number to ensure each chunk has even columns - columns_per_chunk = (total_columns // num_chunks) & ~1 # Round down to nearest even number + columns_per_chunk = ( + total_columns // num_chunks + ) & ~1 # Round down to nearest even number remainder = total_columns - (columns_per_chunk * num_chunks) chunks = [] diff --git a/sqlmesh/core/test/discovery.py b/sqlmesh/core/test/discovery.py index 9afe3dd7fc..bdcb0bcc62 100644 --- a/sqlmesh/core/test/discovery.py +++ b/sqlmesh/core/test/discovery.py @@ -43,6 +43,9 @@ def filter_tests_by_patterns( return unique( test for test, pattern in itertools.product(tests, patterns) - if ("*" in pattern and fnmatch.fnmatchcase(test.fully_qualified_test_name, pattern)) + if ( + "*" in pattern + and fnmatch.fnmatchcase(test.fully_qualified_test_name, pattern) + ) or pattern in test.fully_qualified_test_name ) diff --git a/sqlmesh/core/test/runner.py b/sqlmesh/core/test/runner.py index 284558e1c8..e9db5efe96 100644 --- a/sqlmesh/core/test/runner.py +++ b/sqlmesh/core/test/runner.py @@ -1,25 +1,22 @@ from __future__ import annotations -import time +import concurrent import threading +import time import typing as t import unittest -from io import StringIO - -import concurrent from concurrent.futures import ThreadPoolExecutor +from io import StringIO +from sqlmesh.core.config.connection import BaseDuckDBConnectionConfig from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.model import Model -from sqlmesh.core.test.definition import ModelTest as ModelTest, generate_test as generate_test -from sqlmesh.core.test.discovery import ( - ModelTestMetadata as ModelTestMetadata, -) -from sqlmesh.core.config.connection import BaseDuckDBConnectionConfig +from sqlmesh.core.test.definition import ModelTest as ModelTest +from sqlmesh.core.test.definition import generate_test as generate_test +from sqlmesh.core.test.discovery import ModelTestMetadata as ModelTestMetadata from sqlmesh.core.test.result import ModelTextTestResult as ModelTextTestResult from sqlmesh.utils import UniqueKeyDict, Verbosity - if t.TYPE_CHECKING: from sqlmesh.core.config.loader import C @@ -64,7 +61,9 @@ def create_testing_engine_adapters( # Ensure DuckDB connections are fully isolated from each other # by forcing the creation of a new adapter with SingletonConnectionPool test_connection.concurrent_tasks = 1 - adapter = test_connection.create_engine_adapter(register_comments_override=False) + adapter = test_connection.create_engine_adapter( + register_comments_override=False + ) test_connection.concurrent_tasks = concurrent_tasks elif gateway not in testing_adapter_by_gateway: # All other engines can share connections between threads @@ -123,7 +122,9 @@ def run_tests( ) # Ensure workers are not greater than the number of tests - num_workers = min(len(model_test_metadata) or 1, default_test_connection.concurrent_tasks) + num_workers = min( + len(model_test_metadata) or 1, default_test_connection.concurrent_tasks + ) def _run_single_test( metadata: ModelTestMetadata, engine_adapter: EngineAdapter @@ -160,7 +161,9 @@ def _run_single_test( try: with ThreadPoolExecutor(max_workers=num_workers) as pool: futures = [ - pool.submit(_run_single_test, metadata=metadata, engine_adapter=engine_adapter) + pool.submit( + _run_single_test, metadata=metadata, engine_adapter=engine_adapter + ) for metadata, engine_adapter in metadata_to_adapter.items() ] diff --git a/sqlmesh/core/user.py b/sqlmesh/core/user.py index f40188a471..2c3c8f143d 100644 --- a/sqlmesh/core/user.py +++ b/sqlmesh/core/user.py @@ -1,8 +1,10 @@ import typing as t from enum import Enum -from sqlmesh.core.notification_target import BasicSMTPNotificationTarget, NotificationTarget -from sqlmesh.utils.pydantic import PydanticModel, ValidationInfo, field_validator, validation_data +from sqlmesh.core.notification_target import (BasicSMTPNotificationTarget, + NotificationTarget) +from sqlmesh.utils.pydantic import (PydanticModel, ValidationInfo, + field_validator, validation_data) class UserRole(str, Enum): @@ -44,6 +46,8 @@ def validate_notification_targets( ) -> t.List[NotificationTarget]: email = validation_data(info).get("email") for target in v: - if isinstance(target, BasicSMTPNotificationTarget) and target.recipients != {email}: + if isinstance( + target, BasicSMTPNotificationTarget + ) and target.recipients != {email}: raise ValueError("Recipient emails do not match user email") return v diff --git a/sqlmesh/dbt/__init__.py b/sqlmesh/dbt/__init__.py index 690b1f5289..a39df23e2c 100644 --- a/sqlmesh/dbt/__init__.py +++ b/sqlmesh/dbt/__init__.py @@ -1,4 +1,4 @@ -from sqlmesh.dbt.builtin import ( - create_builtin_filters as create_builtin_filters, - create_builtin_globals as create_builtin_globals, -) +from sqlmesh.dbt.builtin import \ + create_builtin_filters as create_builtin_filters +from sqlmesh.dbt.builtin import \ + create_builtin_globals as create_builtin_globals diff --git a/sqlmesh/dbt/adapter.py b/sqlmesh/dbt/adapter.py index 7f7c7eb4fb..6f9157caf6 100644 --- a/sqlmesh/dbt/adapter.py +++ b/sqlmesh/dbt/adapter.py @@ -6,19 +6,22 @@ from sqlglot import exp, parse_one -from sqlmesh.core.dialect import normalize_and_quote, normalize_model_name, schema_ +from sqlmesh.core.dialect import (normalize_and_quote, normalize_model_name, + schema_) from sqlmesh.core.engine_adapter import EngineAdapter -from sqlmesh.core.snapshot import DeployabilityIndex, Snapshot, to_table_mapping +from sqlmesh.core.schema_diff import TableAlterOperation +from sqlmesh.core.snapshot import (DeployabilityIndex, Snapshot, + to_table_mapping) +from sqlmesh.utils import AttributeDict from sqlmesh.utils.errors import ConfigError, ParsetimeAdapterCallError from sqlmesh.utils.jinja import JinjaMacroRegistry -from sqlmesh.utils import AttributeDict -from sqlmesh.core.schema_diff import TableAlterOperation if t.TYPE_CHECKING: import agate from dbt.adapters.base import BaseRelation from dbt.adapters.base.column import Column from dbt.adapters.base.impl import AdapterResponse + from sqlmesh.core.engine_adapter.base import DataObject from sqlmesh.dbt.relation import Policy @@ -43,7 +46,9 @@ def __init__( self.quote_policy = quote_policy or Policy() @abc.abstractmethod - def get_relation(self, database: str, schema: str, identifier: str) -> t.Optional[BaseRelation]: + def get_relation( + self, database: str, schema: str, identifier: str + ) -> t.Optional[BaseRelation]: """Returns a single relation that matches the provided path.""" @abc.abstractmethod @@ -51,14 +56,18 @@ def load_relation(self, relation: BaseRelation) -> t.Optional[BaseRelation]: """Returns a single relation that matches the provided relation if present.""" @abc.abstractmethod - def list_relations(self, database: t.Optional[str], schema: str) -> t.List[BaseRelation]: + def list_relations( + self, database: t.Optional[str], schema: str + ) -> t.List[BaseRelation]: """Gets all relations in a given schema and optionally database. TODO: Add caching functionality to avoid repeat visits to DB """ @abc.abstractmethod - def list_relations_without_caching(self, schema_relation: BaseRelation) -> t.List[BaseRelation]: + def list_relations_without_caching( + self, schema_relation: BaseRelation + ) -> t.List[BaseRelation]: """Using the engine adapter, gets all the relations that match the given schema grain relation.""" @abc.abstractmethod @@ -90,7 +99,9 @@ def expand_target_column_types( """Expand to_relation's column types to match those of from_relation.""" @abc.abstractmethod - def rename_relation(self, from_relation: BaseRelation, to_relation: BaseRelation) -> None: + def rename_relation( + self, from_relation: BaseRelation, to_relation: BaseRelation + ) -> None: """Renames a relation (table) in the target database.""" @abc.abstractmethod @@ -109,11 +120,17 @@ def resolve_identifier(self, relation: BaseRelation) -> t.Optional[str]: def quote(self, identifier: str) -> str: """Returns a quoted identifier.""" - return exp.to_column(identifier).sql(dialect=self.project_dialect, identify=True) + return exp.to_column(identifier).sql( + dialect=self.project_dialect, identify=True + ) def quote_as_configured(self, value: str, component_type: str) -> str: """Returns the value quoted according to the quote policy.""" - return self.quote(value) if getattr(self.quote_policy, component_type, False) else value + return ( + self.quote(value) + if getattr(self.quote_policy, component_type, False) + else value + ) def dispatch( self, @@ -124,7 +141,9 @@ def dispatch( target_type = self.jinja_globals["target"]["type"] macro_suffix = f"__{macro_name}" - def _relevance(package_name_pair: t.Tuple[t.Optional[str], str]) -> t.Tuple[int, int]: + def _relevance( + package_name_pair: t.Tuple[t.Optional[str], str], + ) -> t.Tuple[int, int]: """Lower scores more relevant.""" macro_package, name = package_name_pair @@ -143,7 +162,10 @@ def _relevance(package_name_pair: t.Tuple[t.Optional[str], str]) -> t.Tuple[int, packages_to_check: t.List[t.Optional[str]] = [None] if macro_namespace is not None: if macro_namespace in jinja_env: - packages_to_check = [self.jinja_macros.root_package_name, macro_namespace] + packages_to_check = [ + self.jinja_macros.root_package_name, + macro_namespace, + ] # Add dbt packages as fallback packages_to_check.extend(k for k in jinja_env if k.startswith("dbt")) @@ -165,7 +187,9 @@ def _relevance(package_name_pair: t.Tuple[t.Optional[str], str]) -> t.Tuple[int, sorted_candidates = sorted(candidates, key=_relevance) return candidates[sorted_candidates[0]] - raise ConfigError(f"Macro '{macro_name}', package '{macro_namespace}' was not found.") + raise ConfigError( + f"Macro '{macro_name}', package '{macro_namespace}' was not found." + ) def type(self) -> str: return self.project_dialect or "" @@ -192,7 +216,9 @@ def graph(self) -> t.Any: class ParsetimeAdapter(BaseAdapter): - def get_relation(self, database: str, schema: str, identifier: str) -> t.Optional[BaseRelation]: + def get_relation( + self, database: str, schema: str, identifier: str + ) -> t.Optional[BaseRelation]: self._raise_parsetime_adapter_call_error("get relation") raise @@ -200,11 +226,15 @@ def load_relation(self, relation: BaseRelation) -> t.Optional[BaseRelation]: self._raise_parsetime_adapter_call_error("load relation") raise - def list_relations(self, database: t.Optional[str], schema: str) -> t.List[BaseRelation]: + def list_relations( + self, database: t.Optional[str], schema: str + ) -> t.List[BaseRelation]: self._raise_parsetime_adapter_call_error("list relation") raise - def list_relations_without_caching(self, schema_relation: BaseRelation) -> t.List[BaseRelation]: + def list_relations_without_caching( + self, schema_relation: BaseRelation + ) -> t.List[BaseRelation]: self._raise_parsetime_adapter_call_error("list relation") raise @@ -232,7 +262,9 @@ def expand_target_column_types( ) -> None: self._raise_parsetime_adapter_call_error("expand target column types") - def rename_relation(self, from_relation: BaseRelation, to_relation: BaseRelation) -> None: + def rename_relation( + self, from_relation: BaseRelation, to_relation: BaseRelation + ) -> None: self._raise_parsetime_adapter_call_error("rename relation") def execute( @@ -296,28 +328,43 @@ def get_relation( return self.load_relation(self._table_to_relation(target_table)) def load_relation(self, relation: BaseRelation) -> t.Optional[BaseRelation]: - mapped_table = self._map_table_name(self._normalize(self._relation_to_table(relation))) + mapped_table = self._map_table_name( + self._normalize(self._relation_to_table(relation)) + ) data_object = self.engine_adapter.get_data_object(mapped_table) - return self._data_object_to_relation(data_object) if data_object is not None else None + return ( + self._data_object_to_relation(data_object) + if data_object is not None + else None + ) - def list_relations(self, database: t.Optional[str], schema: str) -> t.List[BaseRelation]: + def list_relations( + self, database: t.Optional[str], schema: str + ) -> t.List[BaseRelation]: target_schema = schema_(schema, catalog=database) # Normalize before converting to a relation; otherwise, it will be too late, # as quotes will have already been applied. target_schema = self._normalize(target_schema) - return self.list_relations_without_caching(self._table_to_relation(target_schema)) + return self.list_relations_without_caching( + self._table_to_relation(target_schema) + ) - def list_relations_without_caching(self, schema_relation: BaseRelation) -> t.List[BaseRelation]: + def list_relations_without_caching( + self, schema_relation: BaseRelation + ) -> t.List[BaseRelation]: schema = self._normalize(self._schema(schema_relation)) relations = [ - self._data_object_to_relation(do) for do in self.engine_adapter.get_data_objects(schema) + self._data_object_to_relation(do) + for do in self.engine_adapter.get_data_objects(schema) ] return relations def get_columns_in_relation(self, relation: BaseRelation) -> t.List[Column]: - mapped_table = self._map_table_name(self._normalize(self._relation_to_table(relation))) + mapped_table = self._map_table_name( + self._normalize(self._relation_to_table(relation)) + ) if self.project_dialect == "bigquery": # dbt.adapters.bigquery.column.BigQueryColumn has a different constructor signature @@ -338,13 +385,17 @@ def get_columns_in_relation(self, relation: BaseRelation) -> t.List[Column]: Column.from_description( name=name, raw_data_type=dtype.sql(dialect=self.project_dialect) ) - for name, dtype in self.engine_adapter.columns(table_name=mapped_table).items() + for name, dtype in self.engine_adapter.columns( + table_name=mapped_table + ).items() ] return [ self.column_type.from_description( name=name, raw_data_type=dtype.sql(dialect=self.project_dialect) ) - for name, dtype in self.engine_adapter.columns(table_name=mapped_table).items() + for name, dtype in self.engine_adapter.columns( + table_name=mapped_table + ).items() ] def get_missing_columns( @@ -368,12 +419,16 @@ def drop_schema(self, relation: BaseRelation) -> None: def drop_relation(self, relation: BaseRelation) -> None: if relation.schema is not None and relation.identifier is not None: - self.engine_adapter.drop_table(self._normalize(self._relation_to_table(relation))) + self.engine_adapter.drop_table( + self._normalize(self._relation_to_table(relation)) + ) def expand_target_column_types( self, from_relation: BaseRelation, to_relation: BaseRelation ) -> None: - from_dbt_columns = {c.name: c for c in self.get_columns_in_relation(from_relation)} + from_dbt_columns = { + c.name: c for c in self.get_columns_in_relation(from_relation) + } to_dbt_columns = {c.name: c for c in self.get_columns_in_relation(to_relation)} from_table_name = self._normalize(self._relation_to_table(from_relation)) @@ -403,7 +458,9 @@ def expand_target_column_types( if alter_expressions: self.engine_adapter.alter_table(alter_expressions) - def rename_relation(self, from_relation: BaseRelation, to_relation: BaseRelation) -> None: + def rename_relation( + self, from_relation: BaseRelation, to_relation: BaseRelation + ) -> None: old_table_name = self._normalize(self._relation_to_table(from_relation)) new_table_name = self._normalize(self._relation_to_table(to_relation)) @@ -415,7 +472,7 @@ def execute( import pandas as pd from dbt.adapters.base.impl import AdapterResponse - from sqlmesh.dbt.util import pandas_to_agate, empty_table + from sqlmesh.dbt.util import empty_table, pandas_to_agate # mypy bug: https://github.com/python/mypy/issues/10740 exec_func: t.Callable[..., None | pd.DataFrame] = ( @@ -424,7 +481,9 @@ def execute( expression = parse_one(sql, read=self.project_dialect) with normalize_and_quote( - expression, t.cast(str, self.project_dialect), self.engine_adapter.default_catalog + expression, + t.cast(str, self.project_dialect), + self.engine_adapter.default_catalog, ) as expression: expression = exp.replace_tables( expression, self.table_mapping, dialect=self.project_dialect, copy=False @@ -444,11 +503,15 @@ def execute( return AdapterResponse("Success"), empty_table() def resolve_schema(self, relation: BaseRelation) -> t.Optional[str]: - schema = self._map_table_name(self._normalize(self._relation_to_table(relation))).db + schema = self._map_table_name( + self._normalize(self._relation_to_table(relation)) + ).db return schema if schema else None def resolve_identifier(self, relation: BaseRelation) -> t.Optional[str]: - identifier = self._map_table_name(self._normalize(self._relation_to_table(relation))).name + identifier = self._map_table_name( + self._normalize(self._relation_to_table(relation)) + ).name return identifier if identifier else None def _map_table_name(self, table: exp.Table) -> exp.Table: @@ -459,7 +522,9 @@ def _map_table_name(self, table: exp.Table) -> exp.Table: if not physical_table_name: return table - logger.debug("Resolved ref '%s' to snapshot table '%s'", name, physical_table_name) + logger.debug( + "Resolved ref '%s' to snapshot table '%s'", name, physical_table_name + ) return exp.to_table(physical_table_name, dialect=self.project_dialect) @@ -496,8 +561,12 @@ def _schema(self, schema_relation: BaseRelation) -> exp.Table: assert schema_relation.schema is not None return exp.Table( this=None, - db=exp.to_identifier(schema_relation.schema, quoted=self.quote_policy.schema), - catalog=exp.to_identifier(schema_relation.database, quoted=self.quote_policy.database), + db=exp.to_identifier( + schema_relation.schema, quoted=self.quote_policy.schema + ), + catalog=exp.to_identifier( + schema_relation.database, quoted=self.quote_policy.database + ), ) def _normalize(self, input_table: exp.Table) -> exp.Table: diff --git a/sqlmesh/dbt/basemodel.py b/sqlmesh/dbt/basemodel.py index 32a76aba13..db86635a3c 100644 --- a/sqlmesh/dbt/basemodel.py +++ b/sqlmesh/dbt/basemodel.py @@ -1,10 +1,10 @@ from __future__ import annotations +import logging import typing as t from abc import abstractmethod from enum import Enum from pathlib import Path -import logging from pydantic import Field from sqlglot.helper import ensure_list @@ -15,19 +15,10 @@ from sqlmesh.core.model import Model from sqlmesh.core.model.common import ParsableSql from sqlmesh.core.node import DbtNodeInfo -from sqlmesh.dbt.column import ( - ColumnConfig, - column_descriptions_to_sqlmesh, - column_types_to_sqlmesh, -) -from sqlmesh.dbt.common import ( - DbtConfig, - Dependencies, - GeneralConfig, - RAW_CODE_KEY, - SqlStr, - sql_str_validator, -) +from sqlmesh.dbt.column import (ColumnConfig, column_descriptions_to_sqlmesh, + column_types_to_sqlmesh) +from sqlmesh.dbt.common import (RAW_CODE_KEY, DbtConfig, Dependencies, + GeneralConfig, SqlStr, sql_str_validator) from sqlmesh.dbt.relation import Policy, RelationType from sqlmesh.dbt.test import TestConfig from sqlmesh.dbt.util import DBT_VERSION @@ -150,7 +141,9 @@ class BaseModelConfig(GeneralConfig): @field_validator("pre_hook", "post_hook", mode="before") @classmethod - def _validate_hooks(cls, v: t.Union[str, t.List[t.Union[SqlStr, str]]]) -> t.List[Hook]: + def _validate_hooks( + cls, v: t.Union[str, t.List[t.Union[SqlStr, str]]] + ) -> t.List[Hook]: hooks = [] for hook in ensure_list(v): if isinstance(hook, Hook): @@ -292,7 +285,9 @@ def sqlmesh_config_fields(self) -> t.Set[str]: @property def node_info(self) -> DbtNodeInfo: - return DbtNodeInfo(unique_id=self.unique_id, name=self.name, fqn=self.fqn, alias=self.alias) + return DbtNodeInfo( + unique_id=self.unique_id, name=self.name, fqn=self.fqn, alias=self.alias + ) def sqlmesh_model_kwargs( self, @@ -345,17 +340,23 @@ def sqlmesh_model_kwargs( "jinja_macros": jinja_macros, "path": self.path, "pre_statements": [ - ParsableSql(sql=d.jinja_statement(hook.sql).sql(), transaction=hook.transaction) + ParsableSql( + sql=d.jinja_statement(hook.sql).sql(), transaction=hook.transaction + ) for hook in self.pre_hook ], "post_statements": [ - ParsableSql(sql=d.jinja_statement(hook.sql).sql(), transaction=hook.transaction) + ParsableSql( + sql=d.jinja_statement(hook.sql).sql(), transaction=hook.transaction + ) for hook in self.post_hook ], "tags": self.tags, "physical_schema_mapping": context.sqlmesh_config.physical_schema_mapping, "default_catalog": context.target.database, - "grain": [d.parse_one(g, dialect=model_dialect) for g in ensure_list(self.grain)], + "grain": [ + d.parse_one(g, dialect=model_dialect) for g in ensure_list(self.grain) + ], **self.sqlmesh_config_kwargs, } @@ -371,7 +372,8 @@ def sqlmesh_model_kwargs( # - https://docs.getdbt.com/reference/resource-configs/column_types if column_types_override: model_kwargs["columns"] = ( - column_types_to_sqlmesh(column_types_override, self.dialect(context)) or None + column_types_to_sqlmesh(column_types_override, self.dialect(context)) + or None ) return model_kwargs @@ -394,7 +396,10 @@ def _model_jinja_context( model_node: AttributeDict[str, t.Any] = AttributeDict(attributes) else: model_node = AttributeDict( - filter(lambda kv: kv[0] in dependencies.model_attrs.attrs, attributes.items()) + filter( + lambda kv: kv[0] in dependencies.model_attrs.attrs, + attributes.items(), + ) ) # We exclude the raw SQL code to reduce the payload size. It's still accessible through diff --git a/sqlmesh/dbt/builtin.py b/sqlmesh/dbt/builtin.py index fa05e3d7f9..dc1087d652 100644 --- a/sqlmesh/dbt/builtin.py +++ b/sqlmesh/dbt/builtin.py @@ -26,7 +26,8 @@ from sqlmesh.utils import AttributeDict, debug_mode_enabled, yaml from sqlmesh.utils.date import now from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.jinja import JinjaMacroRegistry, MacroReference, MacroReturnVal +from sqlmesh.utils.jinja import (JinjaMacroRegistry, MacroReference, + MacroReturnVal) logger = logging.getLogger(__name__) @@ -117,8 +118,10 @@ def __init__(self, adapter: BaseAdapter): self.adapter = adapter self._results: t.Dict[str, AttributeDict] = {} - def store_result(self, name: str, response: t.Any, agate_table: t.Optional[agate.Table]) -> str: - from sqlmesh.dbt.util import empty_table, as_matrix + def store_result( + self, name: str, response: t.Any, agate_table: t.Optional[agate.Table] + ) -> str: + from sqlmesh.dbt.util import as_matrix, empty_table if agate_table is None: agate_table = empty_table() @@ -136,7 +139,9 @@ def load_result(self, name: str) -> t.Optional[AttributeDict]: return self._results.get(name) def run_query(self, sql: str) -> agate.Table: - self.statement("run_query_statement", fetch_result=True, auto_begin=False, caller=sql) + self.statement( + "run_query_statement", fetch_result=True, auto_begin=False, caller=sql + ) resp = self.load_result("run_query_statement") assert resp is not None return resp["table"] @@ -165,7 +170,9 @@ def statement( raise NotImplementedError( "SQLMesh's dbt integration only supports SQL statements at this time." ) - res, table = self.adapter.execute(sql, fetch=fetch_result, auto_begin=auto_begin) + res, table = self.adapter.execute( + sql, fetch=fetch_result, auto_begin=auto_begin + ) if name: self.store_result(name, res, table) return "" @@ -210,7 +217,9 @@ def set(self, name: str, value: t.Any) -> str: self._config.update({name: value}) return "" - def _validate(self, name: str, validator: t.Callable, value: t.Optional[t.Any] = None) -> None: + def _validate( + self, name: str, validator: t.Callable, value: t.Optional[t.Any] = None + ) -> None: try: validator(value) except Exception as e: @@ -287,7 +296,11 @@ def ref( relation_info = refs.get(ref_name) if not relation_info: versioned_infos = sorted( - [(r, info) for r, info in refs.items() if r.startswith(f"{ref_name}_v")], + [ + (r, info) + for r, info in refs.items() + if r.startswith(f"{ref_name}_v") + ], key=lambda i: i[0], ) if versioned_infos: @@ -306,7 +319,9 @@ def generate_source(sources: t.Dict[str, t.Any], api: Api) -> t.Callable: def source(package: str, name: str) -> t.Optional[BaseRelation]: relation_info = sources.get(f"{package}.{name}") if relation_info is None: - logger.debug("Could not resolve source package='%s' name='%s'", package, name) + logger.debug( + "Could not resolve source package='%s' name='%s'", package, name + ) return None # Clickhouse uses a 2-level schema.table naming scheme, where the second level is called @@ -402,6 +417,7 @@ def _try_literal_eval(value: str) -> t.Any: def debug() -> str: import sys + import ipdb # type: ignore frame = sys._getframe(3) @@ -456,7 +472,9 @@ def create_builtin_globals( jinja_globals = jinja_globals.copy() target: t.Optional[AttributeDict] = jinja_globals.get("target", None) - project_dialect = jinja_globals.pop("dialect", None) or (target.get("type") if target else None) + project_dialect = jinja_globals.pop("dialect", None) or ( + target.get("type") if target else None + ) api = Api(project_dialect) builtin_globals["api"] = api @@ -510,7 +528,9 @@ def create_builtin_globals( ) if (model := jinja_globals.pop("model", None)) is not None: - if isinstance(model_instance := jinja_globals.pop("model_instance", None), SqlModel): + if isinstance( + model_instance := jinja_globals.pop("model_instance", None), SqlModel + ): builtin_globals["model"] = AttributeDict( {**model, RAW_CODE_KEY: model_instance.query.name} ) @@ -585,7 +605,11 @@ def _relation_info_to_relation( quote_policy = Policy( **{ **asdict(target_quote_policy), - **{k: v for k, v in relation_info.pop("quote_policy", {}).items() if v is not None}, + **{ + k: v + for k, v in relation_info.pop("quote_policy", {}).items() + if v is not None + }, } ) return relation_type.create(**relation_info, quote_policy=quote_policy) diff --git a/sqlmesh/dbt/column.py b/sqlmesh/dbt/column.py index 80a6ad9325..e84f786d34 100644 --- a/sqlmesh/dbt/column.py +++ b/sqlmesh/dbt/column.py @@ -1,7 +1,7 @@ from __future__ import annotations -import typing as t import logging +import typing as t from sqlglot import exp, parse_one from sqlglot.helper import ensure_list @@ -50,7 +50,9 @@ def column_types_to_sqlmesh( return col_types_to_sqlmesh -def column_descriptions_to_sqlmesh(columns: t.Dict[str, ColumnConfig]) -> t.Dict[str, str]: +def column_descriptions_to_sqlmesh( + columns: t.Dict[str, ColumnConfig], +) -> t.Dict[str, str]: """ Get the sqlmesh column types diff --git a/sqlmesh/dbt/common.py b/sqlmesh/dbt/common.py index 67e1a788cf..e1deef9538 100644 --- a/sqlmesh/dbt/common.py +++ b/sqlmesh/dbt/common.py @@ -8,9 +8,9 @@ from ruamel.yaml.constructor import DuplicateKeyError from sqlglot.helper import ensure_list -from sqlmesh.dbt.util import DBT_VERSION from sqlmesh.core.config.base import BaseConfig, UpdateStrategy from sqlmesh.core.config.common import DBT_PROJECT_FILENAME +from sqlmesh.dbt.util import DBT_VERSION from sqlmesh.utils import AttributeDict from sqlmesh.utils.conversions import ensure_bool, try_str_to_bool from sqlmesh.utils.errors import ConfigError @@ -40,7 +40,10 @@ def load_yaml(source: str | Path) -> t.Dict: try: return load( - source, render_jinja=False, allow_duplicate_keys=True, keep_last_duplicate_key=True + source, + render_jinja=False, + allow_duplicate_keys=True, + keep_last_duplicate_key=True, ) except DuplicateKeyError as ex: raise ConfigError(f"{source}: {ex}" if isinstance(source, Path) else f"{ex}") @@ -117,7 +120,9 @@ def _validate_list(cls, v: t.Union[str, t.List[str]]) -> t.List[str]: @field_validator("meta", mode="before") @classmethod - def _validate_meta(cls, v: t.Optional[t.Dict[str, t.Union[str, t.Any]]]) -> t.Dict[str, t.Any]: + def _validate_meta( + cls, v: t.Optional[t.Dict[str, t.Union[str, t.Any]]] + ) -> t.Dict[str, t.Any]: return parse_meta(v) _FIELD_UPDATE_STRATEGY: t.ClassVar[t.Dict[str, UpdateStrategy]] = { @@ -211,7 +216,8 @@ def union(self, other: Dependencies) -> Dependencies: attrs=self.model_attrs.attrs | other.model_attrs.attrs, all_attrs=self.model_attrs.all_attrs or other.model_attrs.all_attrs, ), - has_dynamic_var_names=self.has_dynamic_var_names or other.has_dynamic_var_names, + has_dynamic_var_names=self.has_dynamic_var_names + or other.has_dynamic_var_names, ) @field_validator("macros", mode="after") @@ -253,7 +259,9 @@ def jinja_end(sql: str, start: int) -> int: if start == -1: continue extracted = no_config[start : jinja_end(no_config, start)] - only_config = SqlStr("\n".join([only_config, extracted]) if only_config else extracted) + only_config = SqlStr( + "\n".join([only_config, extracted]) if only_config else extracted + ) no_config = SqlStr(no_config.replace(extracted, "").strip()) return (no_config, only_config) diff --git a/sqlmesh/dbt/context.py b/sqlmesh/dbt/context.py index 29eb03700d..7ad50d9275 100644 --- a/sqlmesh/dbt/context.py +++ b/sqlmesh/dbt/context.py @@ -13,13 +13,10 @@ from sqlmesh.dbt.manifest import ManifestHelper from sqlmesh.dbt.target import TargetConfig from sqlmesh.utils import AttributeDict -from sqlmesh.utils.errors import ConfigError, SQLMeshError, MissingModelError, MissingSourceError -from sqlmesh.utils.jinja import ( - JinjaGlobalAttribute, - JinjaMacroRegistry, - MacroInfo, - MacroReference, -) +from sqlmesh.utils.errors import (ConfigError, MissingModelError, + MissingSourceError, SQLMeshError) +from sqlmesh.utils.jinja import (JinjaGlobalAttribute, JinjaMacroRegistry, + MacroInfo, MacroReference) if t.TYPE_CHECKING: from jinja2 import Environment @@ -106,9 +103,13 @@ def add_variables(self, variables: t.Dict[str, t.Any]) -> None: self._variables.update(variables) self._jinja_environment = None - def set_and_render_variables(self, variables: t.Dict[str, t.Any], package: str) -> None: + def set_and_render_variables( + self, variables: t.Dict[str, t.Any], package: str + ) -> None: package_macros = self.jinja_macros.copy( - update={"top_level_packages": [*self.jinja_macros.top_level_packages, package]} + update={ + "top_level_packages": [*self.jinja_macros.top_level_packages, package] + } ) jinja_environment = package_macros.build_environment(**self.jinja_globals) @@ -197,7 +198,8 @@ def refs(self) -> t.Dict[str, t.Union[ModelConfig, SeedConfig]]: if not self._refs: # Refs can be called with or without package name. for model in t.cast( - t.Dict[str, t.Union[ModelConfig, SeedConfig]], {**self._seeds, **self._models} + t.Dict[str, t.Union[ModelConfig, SeedConfig]], + {**self._seeds, **self._models}, ).values(): name = model.name config_name = model.config_name @@ -218,7 +220,9 @@ def target(self) -> TargetConfig: @target.setter def target(self, value: TargetConfig) -> None: if not self.project_name: - raise ConfigError("Project name must be set in the context in order to use a target.") + raise ConfigError( + "Project name must be set in the context in order to use a target." + ) self._target = value self._jinja_environment = None @@ -239,7 +243,9 @@ def copy(self) -> DbtContext: @property def jinja_environment(self) -> Environment: if self._jinja_environment is None: - self._jinja_environment = self.jinja_macros.build_environment(**self.jinja_globals) + self._jinja_environment = self.jinja_macros.build_environment( + **self.jinja_globals + ) return self._jinja_environment @property @@ -247,7 +253,9 @@ def jinja_globals(self) -> t.Dict[str, JinjaGlobalAttribute]: output: t.Dict[str, JinjaGlobalAttribute] = { "vars": AttributeDict(self.variables), "refs": AttributeDict({k: v.relation_info for k, v in self.refs.items()}), - "sources": AttributeDict({k: v.relation_info for k, v in self.sources.items()}), + "sources": AttributeDict( + {k: v.relation_info for k, v in self.sources.items()} + ), } if self.project_name is not None: output["project_name"] = self.project_name @@ -287,7 +295,9 @@ def context_for_dependencies(self, dependencies: Dependencies) -> DbtContext: else: raise MissingSourceError(source) - variables = {k: v for k, v in self.variables.items() if k in dependencies.variables} + variables = { + k: v for k, v in self.variables.items() if k in dependencies.variables + } dependency_context.sources = sources dependency_context.seeds = seeds @@ -298,12 +308,18 @@ def context_for_dependencies(self, dependencies: Dependencies) -> DbtContext: return dependency_context def create_relation( - self, relation_info: AttributeDict[str, t.Any], quote_policy: t.Optional[Policy] = None + self, + relation_info: AttributeDict[str, t.Any], + quote_policy: t.Optional[Policy] = None, ) -> BaseRelation: if not self.target: - raise SQLMeshError("Target must be configured before calling create_relation.") + raise SQLMeshError( + "Target must be configured before calling create_relation." + ) return _relation_info_to_relation( - relation_info, self.target.relation_class, quote_policy or self.target.quote_policy + relation_info, + self.target.relation_class, + quote_policy or self.target.quote_policy, ) diff --git a/sqlmesh/dbt/loader.py b/sqlmesh/dbt/loader.py index fb3ecb2c77..2680072c6e 100644 --- a/sqlmesh/dbt/loader.py +++ b/sqlmesh/dbt/loader.py @@ -3,16 +3,13 @@ import logging import sys import typing as t -import sqlmesh.core.dialect as d -from pathlib import Path from collections import defaultdict -from sqlmesh.core.config import ( - Config, - ConnectionConfig, - GatewayConfig, - ModelDefaultsConfig, - DbtConfig as RootDbtConfig, -) +from pathlib import Path + +import sqlmesh.core.dialect as d +from sqlmesh.core.config import Config, ConnectionConfig +from sqlmesh.core.config import DbtConfig as RootDbtConfig +from sqlmesh.core.config import GatewayConfig, ModelDefaultsConfig from sqlmesh.core.environment import EnvironmentStatements from sqlmesh.core.loader import CacheBase, LoadedProject, Loader from sqlmesh.core.macros import MacroRegistry, macro @@ -26,11 +23,9 @@ from sqlmesh.dbt.project import Project from sqlmesh.dbt.target import TargetConfig from sqlmesh.utils import UniqueKeyDict -from sqlmesh.utils.errors import ConfigError, MissingModelError, BaseMissingReferenceError -from sqlmesh.utils.jinja import ( - JinjaMacroRegistry, - make_jinja_registry, -) +from sqlmesh.utils.errors import (BaseMissingReferenceError, ConfigError, + MissingModelError) +from sqlmesh.utils.jinja import JinjaMacroRegistry, make_jinja_registry if sys.version_info >= (3, 12): from importlib import metadata @@ -58,7 +53,9 @@ def sqlmesh_config( ) -> Config: project_root = project_root or Path() context = DbtContext( - project_root=project_root, profiles_dir=profiles_dir, profile_name=dbt_profile_name + project_root=project_root, + profiles_dir=profiles_dir, + profile_name=dbt_profile_name, ) # note: Profile.load() is called twice with different DbtContext's: @@ -160,7 +157,9 @@ def _load_models( models: UniqueKeyDict[str, Model] = UniqueKeyDict("models") def _to_sqlmesh(config: BMC, context: DbtContext) -> Model: - logger.debug("Converting '%s' to sqlmesh format", config.canonical_name(context)) + logger.debug( + "Converting '%s' to sqlmesh format", config.canonical_name(context) + ) return config.to_sqlmesh( context, audit_definitions=audits, @@ -178,13 +177,22 @@ def _to_sqlmesh(config: BMC, context: DbtContext) -> Model: # Now that config is rendered, create the sqlmesh models for package in project.packages.values(): package_context = project.context.copy() - package_context.set_and_render_variables(package.variables, package.name) - package_models: t.Dict[str, BaseModelConfig] = {**package.models, **package.seeds} + package_context.set_and_render_variables( + package.variables, package.name + ) + package_models: t.Dict[str, BaseModelConfig] = { + **package.models, + **package.seeds, + } - package_models_by_path: t.Dict[Path, t.List[BaseModelConfig]] = defaultdict(list) + package_models_by_path: t.Dict[Path, t.List[BaseModelConfig]] = ( + defaultdict(list) + ) for model in package_models.values(): if isinstance(model, ModelConfig) and not model.sql.strip(): - logger.info(f"Skipping empty model '{model.name}' at path '{model.path}'.") + logger.info( + f"Skipping empty model '{model.name}' at path '{model.path}'." + ) continue package_models_by_path[model.path].append(model) @@ -211,14 +219,18 @@ def _load_audits( logger.debug("Converting audits to sqlmesh") for package in project.packages.values(): package_context = project.context.copy() - package_context.set_and_render_variables(package.variables, package.name) + package_context.set_and_render_variables( + package.variables, package.name + ) for test in package.tests.values(): logger.debug("Converting '%s' to sqlmesh format", test.name) try: audits[test.canonical_name] = test.to_sqlmesh(package_context) except BaseMissingReferenceError as e: - ref_type = "model" if isinstance(e, MissingModelError) else "source" + ref_type = ( + "model" if isinstance(e, MissingModelError) else "source" + ) logger.warning( "Skipping audit '%s' because %s '%s' is not a valid ref", test.name, @@ -285,18 +297,25 @@ def _load_requirements(self) -> t.Tuple[t.Dict[str, str], t.Set[str]]: target_packages.append(f"dbt-{project.context.target.type}") for target_package in target_packages: - if target_package in requirements or target_package in excluded_requirements: + if ( + target_package in requirements + or target_package in excluded_requirements + ): continue try: requirements[target_package] = metadata.version(target_package) except metadata.PackageNotFoundError: from sqlmesh.core.console import get_console - get_console().log_warning(f"dbt package {target_package} is not installed.") + get_console().log_warning( + f"dbt package {target_package} is not installed." + ) return requirements, excluded_requirements - def _load_environment_statements(self, macros: MacroRegistry) -> t.List[EnvironmentStatements]: + def _load_environment_statements( + self, macros: MacroRegistry + ) -> t.List[EnvironmentStatements]: """Loads dbt's on_run_start, on_run_end hooks into sqlmesh's before_all, after_all statements respectively.""" hooks_by_package_name: t.Dict[str, EnvironmentStatements] = {} @@ -305,24 +324,37 @@ def _load_environment_statements(self, macros: MacroRegistry) -> t.List[Environm for project in self._load_projects(): for package_name, package in project.packages.items(): package_context = project.context.copy() - package_context.set_and_render_variables(package.variables, package_name) + package_context.set_and_render_variables( + package.variables, package_name + ) on_run_start: t.List[str] = [ on_run_hook.sql - for on_run_hook in sorted(package.on_run_start.values(), key=lambda h: h.index) + for on_run_hook in sorted( + package.on_run_start.values(), key=lambda h: h.index + ) ] on_run_end: t.List[str] = [ on_run_hook.sql - for on_run_hook in sorted(package.on_run_end.values(), key=lambda h: h.index) + for on_run_hook in sorted( + package.on_run_end.values(), key=lambda h: h.index + ) ] if on_run_start or on_run_end: dependencies = Dependencies() - for hook in [*package.on_run_start.values(), *package.on_run_end.values()]: + for hook in [ + *package.on_run_start.values(), + *package.on_run_end.values(), + ]: dependencies = dependencies.union(hook.dependencies) - statements_context = package_context.context_for_dependencies(dependencies) + statements_context = package_context.context_for_dependencies( + dependencies + ) jinja_registry = make_jinja_registry( - statements_context.jinja_macros, package_name, set(dependencies.macros) + statements_context.jinja_macros, + package_name, + set(dependencies.macros), ) jinja_registry.add_globals(statements_context.jinja_globals) @@ -366,13 +398,17 @@ def _compute_yaml_max_mtime_per_subfolder( try: if nested.is_dir(): result.update( - self._compute_yaml_max_mtime_per_subfolder(nested, visited=visited) + self._compute_yaml_max_mtime_per_subfolder( + nested, visited=visited + ) ) elif nested.suffix.lower() in (".yaml", ".yml"): yaml_mtime = self._path_mtimes.get(nested) if yaml_mtime: max_mtime = ( - max(max_mtime, yaml_mtime) if max_mtime is not None else yaml_mtime + max(max_mtime, yaml_mtime) + if max_mtime is not None + else yaml_mtime ) except PermissionError: pass diff --git a/sqlmesh/dbt/manifest.py b/sqlmesh/dbt/manifest.py index aae4de871c..89e4c55126 100644 --- a/sqlmesh/dbt/manifest.py +++ b/sqlmesh/dbt/manifest.py @@ -35,18 +35,20 @@ from dbt.parser.manifest import ManifestLoader try: - from dbt.parser.sources import merge_freshness # type: ignore[attr-defined] + from dbt.parser.sources import \ + merge_freshness # type: ignore[attr-defined] except ImportError: # merge_freshness was renamed to merge_source_freshness in dbt 1.10 # ref: https://github.com/dbt-labs/dbt-core/commit/14fc39a76ff4830cdf2fcbe73f57ca27db500018#diff-1f09db95588f46879a83378c2a86d6b16b7cdfcaddbfe46afc5d919ee5e9a4d9R430 from dbt.parser.sources import merge_source_freshness as merge_freshness # type: ignore[no-redef,attr-defined] from dbt.tracking import do_not_track +from sqlglot.helper import ensure_list from sqlmesh.core import constants as c -from sqlmesh.utils.errors import SQLMeshError from sqlmesh.core.config import ModelDefaultsConfig -from sqlmesh.dbt.builtin import BUILTIN_FILTERS, BUILTIN_GLOBALS, OVERRIDDEN_MACROS +from sqlmesh.dbt.builtin import (BUILTIN_FILTERS, BUILTIN_GLOBALS, + OVERRIDDEN_MACROS) from sqlmesh.dbt.common import Dependencies from sqlmesh.dbt.model import ModelConfig from sqlmesh.dbt.package import HookConfig, MacroConfig, MaterializationConfig @@ -56,18 +58,14 @@ from sqlmesh.dbt.test import TestConfig from sqlmesh.dbt.util import DBT_VERSION from sqlmesh.utils.cache import FileCache -from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.jinja import ( - MacroInfo, - MacroReference, - extract_call_names, - jinja_call_arg_name, -) -from sqlglot.helper import ensure_list +from sqlmesh.utils.errors import ConfigError, SQLMeshError +from sqlmesh.utils.jinja import (MacroInfo, MacroReference, extract_call_names, + jinja_call_arg_name) if t.TYPE_CHECKING: from dbt.contracts.graph.manifest import Macro, Manifest from dbt.contracts.graph.nodes import ManifestNode, SourceDefinition + from sqlmesh.utils.jinja import CallNames logger = logging.getLogger(__name__) @@ -87,7 +85,8 @@ # Patch Semantic Manifest to skip validation and avoid Pydantic v1 errors on DBT 1.6 # We patch for 1.7+ since we don't care about semantic models if DBT_VERSION >= (1, 6, 0): - from dbt.contracts.graph.semantic_manifest import SemanticManifest # type: ignore + from dbt.contracts.graph.semantic_manifest import \ + SemanticManifest # type: ignore SemanticManifest.validate = lambda _: True # type: ignore @@ -120,7 +119,9 @@ def __init__( self._sources_per_package: t.Dict[str, SourceConfigs] = defaultdict(dict) self._macros_per_package: t.Dict[str, MacroConfigs] = defaultdict(dict) - self._macro_flatten_dependencies: t.Dict[str, t.Dict[str, Dependencies]] = defaultdict(dict) + self._macro_flatten_dependencies: t.Dict[str, t.Dict[str, Dependencies]] = ( + defaultdict(dict) + ) self._tests_by_owner: t.Dict[str, t.List[TestConfig]] = defaultdict(list) self._disabled_refs: t.Optional[t.Set[str]] = None @@ -219,7 +220,9 @@ def _load_all(self) -> None: if self._is_loaded: return - self._calls = {k: (v, False) for k, v in (self._call_cache.get("") or {}).items()} + self._calls = { + k: (v, False) for k, v in (self._call_cache.get("") or {}).items() + } self._load_macros() self._load_materializations() @@ -229,7 +232,9 @@ def _load_all(self) -> None: self._load_on_run_start_end() self._is_loaded = True - self._call_cache.put("", value={k: v for k, (v, used) in self._calls.items() if used}) + self._call_cache.put( + "", value={k: v for k, (v, used) in self._calls.items() if used} + ) def _load_sources(self) -> None: for source in self._manifest.sources.values(): @@ -254,9 +259,9 @@ def _load_sources(self) -> None: "freshness": freshness.to_dict() if freshness else None, } ) - self._sources_per_package[source.package_name][source_config.config_name] = ( - source_config - ) + self._sources_per_package[source.package_name][ + source_config.config_name + ] = source_config def _load_macros(self) -> None: for macro in self._manifest.macros.values(): @@ -303,7 +308,9 @@ def _load_materializations(self) -> None: mat_name = "_".join(name_parts[1:-1]) adapter = name_parts[-1] - dependencies = Dependencies(macros=_macro_references(self._manifest, macro)) + dependencies = Dependencies( + macros=_macro_references(self._manifest, macro) + ) macro.macro_sql = _strip_jinja_materialization_tags(macro.macro_sql) dependencies = dependencies.union( self._extra_dependencies(macro.macro_sql, macro.package_name) @@ -346,13 +353,21 @@ def _load_tests(self) -> None: sources=_sources(node), ) # Implicit dependencies for model test arg - dependencies.macros.append(MacroReference(package="dbt", name="get_where_subquery")) - dependencies.macros.append(MacroReference(package="dbt", name="should_store_failures")) + dependencies.macros.append( + MacroReference(package="dbt", name="get_where_subquery") + ) + dependencies.macros.append( + MacroReference(package="dbt", name="should_store_failures") + ) sql = node.raw_code if DBT_VERSION >= (1, 3, 0) else node.raw_sql # type: ignore - dependencies = dependencies.union(self._extra_dependencies(sql, node.package_name)) dependencies = dependencies.union( - self._flatten_dependencies_from_macros(dependencies.macros, node.package_name) + self._extra_dependencies(sql, node.package_name) + ) + dependencies = dependencies.union( + self._flatten_dependencies_from_macros( + dependencies.macros, node.package_name + ) ) test_model = _test_model(node) @@ -362,7 +377,9 @@ def _load_tests(self) -> None: test = TestConfig( sql=sql, model_name=test_model, - test_kwargs=node.test_metadata.kwargs if hasattr(node, "test_metadata") else {}, + test_kwargs=( + node.test_metadata.kwargs if hasattr(node, "test_metadata") else {} + ), dependencies=dependencies, **node_config, ) @@ -398,16 +415,23 @@ def _load_models_and_seeds(self) -> None: macros=macro_references, refs=_refs(node), sources=_sources(node) ) dependencies = dependencies.union( - self._extra_dependencies(sql, node.package_name, track_all_model_attrs=True) + self._extra_dependencies( + sql, node.package_name, track_all_model_attrs=True + ) ) - for hook in [*node_config.get("pre-hook", []), *node_config.get("post-hook", [])]: + for hook in [ + *node_config.get("pre-hook", []), + *node_config.get("post-hook", []), + ]: dependencies = dependencies.union( self._extra_dependencies( hook["sql"], node.package_name, track_all_model_attrs=True ) ) dependencies = dependencies.union( - self._flatten_dependencies_from_macros(dependencies.macros, node.package_name) + self._flatten_dependencies_from_macros( + dependencies.macros, node.package_name + ) ) self._models_per_package[node.package_name][node_name] = ModelConfig( @@ -441,24 +465,32 @@ def _load_on_run_start_end(self) -> None: refs=_refs(node), sources=_sources(node), ) - dependencies = dependencies.union(self._extra_dependencies(sql, node.package_name)) dependencies = dependencies.union( - self._flatten_dependencies_from_macros(dependencies.macros, node.package_name) + self._extra_dependencies(sql, node.package_name) + ) + dependencies = dependencies.union( + self._flatten_dependencies_from_macros( + dependencies.macros, node.package_name + ) ) if "on-run-start" in node.tags: - self._on_run_start_per_package[node.package_name][node_name] = HookConfig( - sql=sql, - index=getattr(node, "index", None) or 0, - path=node_path, - dependencies=dependencies, + self._on_run_start_per_package[node.package_name][node_name] = ( + HookConfig( + sql=sql, + index=getattr(node, "index", None) or 0, + path=node_path, + dependencies=dependencies, + ) ) else: - self._on_run_end_per_package[node.package_name][node_name] = HookConfig( - sql=sql, - index=getattr(node, "index", None) or 0, - path=node_path, - dependencies=dependencies, + self._on_run_end_per_package[node.package_name][node_name] = ( + HookConfig( + sql=sql, + index=getattr(node, "index", None) or 0, + path=node_path, + dependencies=dependencies, + ) ) @property @@ -491,7 +523,8 @@ def _load_manifest(self) -> Manifest: flags.set_from_args(args, None) if DBT_VERSION >= (1, 8, 0): - from dbt_common.context import set_invocation_context # type: ignore + from dbt_common.context import \ + set_invocation_context # type: ignore set_invocation_context(os.environ) @@ -524,7 +557,9 @@ def _load_manifest(self) -> Manifest: return manifest def _load_project(self, profile: Profile) -> Project: - project_renderer = DbtProjectYamlRenderer(profile, cli_vars=self.variable_overrides) + project_renderer = DbtProjectYamlRenderer( + profile, cli_vars=self.variable_overrides + ) return Project.from_project_root(str(self.project_path), project_renderer) def _load_profile(self) -> Profile: @@ -583,9 +618,9 @@ def _flatten_dependencies_from_macros( continue visited.add((macro_package, macro.name)) - macro_dependencies = self._macro_flatten_dependencies.get(macro_package, {}).get( - macro.name - ) + macro_dependencies = self._macro_flatten_dependencies.get( + macro_package, {} + ).get(macro.name) if not macro_dependencies: macro_config = self._macros_per_package[macro_package].get(macro.name) if not macro_config: @@ -599,7 +634,9 @@ def _flatten_dependencies_from_macros( # We don't need flatten macro dependencies. The jinja macro registry takes care of recursive # dependencies for us. macro_dependencies.macros = [] - self._macro_flatten_dependencies[macro_package][macro.name] = macro_dependencies + self._macro_flatten_dependencies[macro_package][ + macro.name + ] = macro_dependencies dependencies = dependencies.union(macro_dependencies) return dependencies @@ -627,7 +664,10 @@ def _extra_dependencies( track_all_model_attrs and not all_model_attrs and isinstance(node, jinja2.nodes.Call) - and any(isinstance(a, jinja2.nodes.Name) and a.name == "model" for a in node.args) + and any( + isinstance(a, jinja2.nodes.Name) and a.name == "model" + for a in node.args + ) ): all_model_attrs = True @@ -689,7 +729,9 @@ def _extra_dependencies( def _macro_reference_if_not_overridden( - package: t.Optional[str], name: str, if_not_overridden: t.Callable[[MacroReference], None] + package: t.Optional[str], + name: str, + if_not_overridden: t.Callable[[MacroReference], None], ) -> None: reference = MacroReference(package=package, name=name) if reference not in OVERRIDDEN_MACROS: @@ -714,7 +756,9 @@ def _macro_references( macro_node = manifest.macros[macro_node_id] macro_name = macro_node.name macro_package = ( - macro_node.package_name if macro_node.package_name != node.package_name else None + macro_node.package_name + if macro_node.package_name != node.package_name + else None ) _macro_reference_if_not_overridden(macro_package, macro_name, result.add) return result @@ -783,7 +827,9 @@ def _convert_jinja_test_to_macro(test_jinja: str) -> str: macro_tag = re.sub(r"({%-?\s*)test\s+", r"\1macro test_", test_tag) macro = macro_tag + test_jinja[match.span()[-1] :] - return re.sub(ENDTEST_REGEX, lambda m: m.group(0).replace("endtest", "endmacro"), macro) + return re.sub( + ENDTEST_REGEX, lambda m: m.group(0).replace("endtest", "endmacro"), macro + ) def _strip_jinja_materialization_tags(materialization_jinja: str) -> str: @@ -829,7 +875,9 @@ def _build_test_name(node: ManifestNode, dependencies: Dependencies) -> str: if not model_name and dependencies.sources: # extract source and table names source_parts = list(dependencies.sources)[0].split(".") - source_name = "_".join(source_parts) if len(source_parts) == 2 else source_parts[-1] + source_name = ( + "_".join(source_parts) if len(source_parts) == 2 else source_parts[-1] + ) entity_name = model_name or source_name or "" entity_name = f"_{entity_name}" if entity_name else "" diff --git a/sqlmesh/dbt/model.py b/sqlmesh/dbt/model.py index 55994abf85..f5b7f3e66c 100644 --- a/sqlmesh/dbt/model.py +++ b/sqlmesh/dbt/model.py @@ -1,8 +1,8 @@ from __future__ import annotations import datetime -import typing as t import logging +import typing as t from sqlglot import exp from sqlglot.errors import SqlglotError @@ -12,28 +12,18 @@ from sqlmesh.core.config.base import UpdateStrategy from sqlmesh.core.config.common import VirtualEnvironmentMode from sqlmesh.core.console import get_console -from sqlmesh.core.model import ( - EmbeddedKind, - FullKind, - IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind, - IncrementalUnmanagedKind, - Model, - ModelKind, - SCDType2ByColumnKind, - ViewKind, - ManagedKind, - create_sql_model, -) -from sqlmesh.core.model.kind import ( - SCDType2ByTimeKind, - OnDestructiveChange, - OnAdditiveChange, - on_destructive_change_validator, - on_additive_change_validator, - DbtCustomKind, -) -from sqlmesh.dbt.basemodel import BaseModelConfig, Materialization, SnapshotStrategy +from sqlmesh.core.model import (EmbeddedKind, FullKind, + IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, + IncrementalUnmanagedKind, ManagedKind, Model, + ModelKind, SCDType2ByColumnKind, ViewKind, + create_sql_model) +from sqlmesh.core.model.kind import (DbtCustomKind, OnAdditiveChange, + OnDestructiveChange, SCDType2ByTimeKind, + on_additive_change_validator, + on_destructive_change_validator) +from sqlmesh.dbt.basemodel import (BaseModelConfig, Materialization, + SnapshotStrategy) from sqlmesh.dbt.common import SqlStr, sql_str_validator from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.pydantic import field_validator @@ -167,7 +157,9 @@ def _validate_list(cls, v: t.Union[str, t.List[str]]) -> t.List[str]: @field_validator("check_cols", mode="before") @classmethod - def _validate_check_cols(cls, v: t.Union[str, t.List[str]]) -> t.Union[str, t.List[str]]: + def _validate_check_cols( + cls, v: t.Union[str, t.List[str]] + ) -> t.Union[str, t.List[str]]: if isinstance(v, str) and v.lower() == "all": return "*" return ensure_list(v) @@ -214,7 +206,9 @@ def _validate_partition_by( "hour", ): granularity = v["granularity"] - raise ConfigError(f"Unexpected granularity '{granularity}' in partition_by '{v}'.") + raise ConfigError( + f"Unexpected granularity '{granularity}' in partition_by '{v}'." + ) if "data_type" in v and v["data_type"].lower() not in ( "timestamp", "date", @@ -222,7 +216,9 @@ def _validate_partition_by( "int64", ): data_type = v["data_type"] - raise ConfigError(f"Unexpected data_type '{data_type}' in partition_by '{v}'.") + raise ConfigError( + f"Unexpected data_type '{data_type}' in partition_by '{v}'." + ) return {"data_type": "date", "granularity": "day", **v} raise ConfigError(f"Invalid format for partition_by '{v}'") @@ -311,21 +307,26 @@ def model_kind(self, context: DbtContext) -> ModelKind: ) auto_restatement_cron_value = self._get_field_value("auto_restatement_cron") if auto_restatement_cron_value is not None: - incremental_kind_kwargs["auto_restatement_cron"] = auto_restatement_cron_value + incremental_kind_kwargs["auto_restatement_cron"] = ( + auto_restatement_cron_value + ) if materialization == Materialization.TABLE: return FullKind() if materialization == Materialization.VIEW: return ViewKind() if materialization == Materialization.INCREMENTAL: - incremental_by_kind_kwargs: t.Dict[str, t.Any] = {"dialect": self.dialect(context)} + incremental_by_kind_kwargs: t.Dict[str, t.Any] = { + "dialect": self.dialect(context) + } forward_only_value = self._get_field_value("forward_only") if forward_only_value is not None: incremental_kind_kwargs["forward_only"] = forward_only_value is_incremental_by_time_range = self.time_column or ( self.incremental_strategy - and self.incremental_strategy in {"microbatch", "incremental_by_time_range"} + and self.incremental_strategy + in {"microbatch", "incremental_by_time_range"} ) # Get shared incremental by kwargs for field in ("batch_size", "batch_concurrency", "lookback"): @@ -342,13 +343,16 @@ def model_kind(self, context: DbtContext) -> ModelKind: disable_restatement = False else: disable_restatement = ( - not self.full_refresh if self.full_refresh is not None else False + not self.full_refresh + if self.full_refresh is not None + else False ) incremental_by_kind_kwargs["disable_restatement"] = disable_restatement if is_incremental_by_time_range: - strategy = self.incremental_strategy or target.default_incremental_strategy( - IncrementalByTimeRangeKind + strategy = ( + self.incremental_strategy + or target.default_incremental_strategy(IncrementalByTimeRangeKind) ) if strategy not in INCREMENTAL_BY_TIME_RANGE_STRATEGIES: @@ -401,8 +405,9 @@ def model_kind(self, context: DbtContext) -> ModelKind: ) if self.unique_key: - strategy = self.incremental_strategy or target.default_incremental_strategy( - IncrementalByUniqueKeyKind + strategy = ( + self.incremental_strategy + or target.default_incremental_strategy(IncrementalByUniqueKeyKind) ) if ( self.incremental_strategy @@ -461,10 +466,14 @@ def model_kind(self, context: DbtContext) -> ModelKind: } if self.snapshot_strategy.is_check: return SCDType2ByColumnKind( - columns=self.check_cols, execution_time_as_valid_from=True, **shared_kwargs + columns=self.check_cols, + execution_time_as_valid_from=True, + **shared_kwargs, ) return SCDType2ByTimeKind( - updated_at_name=self.updated_at, updated_at_as_valid_from=True, **shared_kwargs + updated_at_name=self.updated_at, + updated_at_as_valid_from=True, + **shared_kwargs, ) if materialization == Materialization.DYNAMIC_TABLE: @@ -522,7 +531,9 @@ def _big_query_partition_by_expr(self, context: DbtContext) -> exp.Expr: dialect="bigquery", ) - def _get_custom_materialization(self, context: DbtContext) -> t.Optional[MaterializationConfig]: + def _get_custom_materialization( + self, context: DbtContext + ) -> t.Optional[MaterializationConfig]: materializations = context.manifest.materializations() name, target_adapter = self.materialized, context.target.dialect @@ -585,7 +596,9 @@ def to_sqlmesh( ) from e elif isinstance(self.partition_by, dict): if context.target.dialect == "bigquery": - partitioned_by.append(self._big_query_partition_by_expr(context)) + partitioned_by.append( + self._big_query_partition_by_expr(context) + ) else: logger.warning( "Ignoring partition_by config for model '%s' targeting %s. The format of the config field is only supported for BigQuery.", @@ -608,7 +621,10 @@ def to_sqlmesh( for c in self.cluster_by: try: cluster_expr = exp.maybe_parse( - c, into=exp.Cluster, prefix="CLUSTER BY", dialect=model_dialect + c, + into=exp.Cluster, + prefix="CLUSTER BY", + dialect=model_dialect, ) for expr in cluster_expr.expressions: clustered_by.append( @@ -627,23 +643,33 @@ def to_sqlmesh( if context.target.dialect == "bigquery": dbt_max_partition_blob = self._dbt_max_partition_blob() if dbt_max_partition_blob: - model_kwargs["pre_statements"].append(d.jinja_statement(dbt_max_partition_blob)) + model_kwargs["pre_statements"].append( + d.jinja_statement(dbt_max_partition_blob) + ) if self.partition_expiration_days is not None: - physical_properties["partition_expiration_days"] = self.partition_expiration_days + physical_properties["partition_expiration_days"] = ( + self.partition_expiration_days + ) if self.require_partition_filter is not None: - physical_properties["require_partition_filter"] = self.require_partition_filter + physical_properties["require_partition_filter"] = ( + self.require_partition_filter + ) if physical_properties: model_kwargs["physical_properties"] = physical_properties if context.target.dialect == "snowflake": if self.snowflake_warehouse is not None: - model_kwargs["session_properties"] = {"warehouse": self.snowflake_warehouse} + model_kwargs["session_properties"] = { + "warehouse": self.snowflake_warehouse + } if self.model_materialization == Materialization.DYNAMIC_TABLE: if not self.snowflake_warehouse: - raise ConfigError("`snowflake_warehouse` must be set for dynamic tables") + raise ConfigError( + "`snowflake_warehouse` must be set for dynamic tables" + ) if not self.target_lag: raise ConfigError("`target_lag` must be set for dynamic tables") @@ -679,7 +705,11 @@ def to_sqlmesh( if self.order_by: order_by = [] - for o in self.order_by if isinstance(self.order_by, list) else [self.order_by]: + for o in ( + self.order_by + if isinstance(self.order_by, list) + else [self.order_by] + ): try: order_by.append(d.parse_one(o, dialect=model_dialect)) except SqlglotError as e: @@ -710,7 +740,9 @@ def to_sqlmesh( ) if self.settings: - physical_properties.update({k: exp.var(v) for k, v in self.settings.items()}) + physical_properties.update( + {k: exp.var(v) for k, v in self.settings.items()} + ) if physical_properties: model_kwargs["physical_properties"] = physical_properties diff --git a/sqlmesh/dbt/package.py b/sqlmesh/dbt/package.py index dbaa832c22..b74a37456f 100644 --- a/sqlmesh/dbt/package.py +++ b/sqlmesh/dbt/package.py @@ -98,16 +98,24 @@ def load(self, package_root: Path) -> Package: all_variables.update(all_variables.pop(package_name, None) or {}) package_variables = { - var: value for var, value in all_variables.items() if not isinstance(value, dict) + var: value + for var, value in all_variables.items() + if not isinstance(value, dict) } tests = _fix_paths(self._context.manifest.tests(package_name), package_root) models = _fix_paths(self._context.manifest.models(package_name), package_root) seeds = _fix_paths(self._context.manifest.seeds(package_name), package_root) macros = _fix_paths(self._context.manifest.macros(package_name), package_root) - materializations = _fix_paths(self._context.manifest.materializations(), package_root) - on_run_start = _fix_paths(self._context.manifest.on_run_start(package_name), package_root) - on_run_end = _fix_paths(self._context.manifest.on_run_end(package_name), package_root) + materializations = _fix_paths( + self._context.manifest.materializations(), package_root + ) + on_run_start = _fix_paths( + self._context.manifest.on_run_start(package_name), package_root + ) + on_run_end = _fix_paths( + self._context.manifest.on_run_end(package_name), package_root + ) sources = self._context.manifest.sources(package_name) config_paths = { @@ -134,7 +142,13 @@ def load(self, package_root: Path) -> Package: T = t.TypeVar( - "T", TestConfig, ModelConfig, MacroConfig, MaterializationConfig, SeedConfig, HookConfig + "T", + TestConfig, + ModelConfig, + MacroConfig, + MaterializationConfig, + SeedConfig, + HookConfig, ) diff --git a/sqlmesh/dbt/profile.py b/sqlmesh/dbt/profile.py index a95c81501c..bcac64f247 100644 --- a/sqlmesh/dbt/profile.py +++ b/sqlmesh/dbt/profile.py @@ -51,7 +51,9 @@ def load(cls, context: DbtContext, target_name: t.Optional[str] = None) -> Profi if not context.profile_name: project_file = Path(context.project_root, PROJECT_FILENAME) if not project_file.exists(): - raise ConfigError(f"Could not find {PROJECT_FILENAME} in {context.project_root}") + raise ConfigError( + f"Could not find {PROJECT_FILENAME} in {context.project_root}" + ) project_yaml = load_yaml(project_file) context.profile_name = context.render( @@ -68,7 +70,9 @@ def load(cls, context: DbtContext, target_name: t.Optional[str] = None) -> Profi return Profile(profile_filepath, target_name, target) @classmethod - def _find_profile(cls, project_root: Path, profiles_dir: t.Optional[Path]) -> t.Optional[Path]: + def _find_profile( + cls, project_root: Path, profiles_dir: t.Optional[Path] + ) -> t.Optional[Path]: dir = os.environ.get("DBT_PROFILES_DIR", profiles_dir or "") path = Path(project_root, dir, cls.PROFILE_FILE) if path.exists(): @@ -89,11 +93,15 @@ def _read_profile( logger.debug("Processing profile '%s'.", path) project_data = load_yaml(path).get(context.profile_name) if not project_data: - raise ConfigError(f"Profile '{context.profile_name}' not found in profiles.") + raise ConfigError( + f"Profile '{context.profile_name}' not found in profiles." + ) outputs = project_data.get("outputs") if not outputs: - raise ConfigError(f"No outputs exist in profiles for '{context.profile_name}'.") + raise ConfigError( + f"No outputs exist in profiles for '{context.profile_name}'." + ) if not target_name: if "target" not in project_data: diff --git a/sqlmesh/dbt/project.py b/sqlmesh/dbt/project.py index 2b0a2e0c3f..c27ecb2aed 100644 --- a/sqlmesh/dbt/project.py +++ b/sqlmesh/dbt/project.py @@ -1,7 +1,7 @@ from __future__ import annotations -import typing as t import logging +import typing as t from pathlib import Path from sqlmesh.core.console import get_console @@ -36,7 +36,9 @@ def __init__( self.packages = packages @classmethod - def load(cls, context: DbtContext, variables: t.Optional[t.Dict[str, t.Any]] = None) -> Project: + def load( + cls, context: DbtContext, variables: t.Optional[t.Dict[str, t.Any]] = None + ) -> Project: """ Loads the configuration for the specified DBT project @@ -52,7 +54,9 @@ def load(cls, context: DbtContext, variables: t.Optional[t.Dict[str, t.Any]] = N project_file_path = Path(context.project_root, PROJECT_FILENAME) logger.debug("Processing project file '%s'.", project_file_path) if not project_file_path.exists(): - raise ConfigError(f"Could not find {PROJECT_FILENAME} in {context.project_root}") + raise ConfigError( + f"Could not find {PROJECT_FILENAME} in {context.project_root}" + ) project_yaml = load_yaml(project_file_path) project_name = context.render(project_yaml.get("name", "")) @@ -60,7 +64,9 @@ def load(cls, context: DbtContext, variables: t.Optional[t.Dict[str, t.Any]] = N if not context.project_name: raise ConfigError(f"{project_file_path.stem} must include project name.") - profile_name = context.render(project_yaml.get("profile", "")) or context.project_name + profile_name = ( + context.render(project_yaml.get("profile", "")) or context.project_name + ) context.profile_name = profile_name profile = Profile.load(context, context.target_name) @@ -104,7 +110,10 @@ def load(cls, context: DbtContext, variables: t.Optional[t.Dict[str, t.Any]] = N # 2. Package-scoped variables in the root project's dbt_project.yml # 3. Global project variables in the root project's dbt_project.yml # 4. Variables in the package's dbt_project.yml - all_project_variables = {**(project_yaml.get("vars") or {}), **(variable_overrides or {})} + all_project_variables = { + **(project_yaml.get("vars") or {}), + **(variable_overrides or {}), + } for name, package in packages.items(): if isinstance(all_project_variables.get(name), dict): project_vars_copy = all_project_variables.copy() diff --git a/sqlmesh/dbt/relation.py b/sqlmesh/dbt/relation.py index fff9f75593..390501f04e 100644 --- a/sqlmesh/dbt/relation.py +++ b/sqlmesh/dbt/relation.py @@ -1,6 +1,5 @@ from sqlmesh.dbt.util import DBT_VERSION - if DBT_VERSION >= (1, 8, 0): from dbt.adapters.contracts.relation import * # type: ignore # noqa: F403 else: diff --git a/sqlmesh/dbt/seed.py b/sqlmesh/dbt/seed.py index c0c8186f29..efa84d64ab 100644 --- a/sqlmesh/dbt/seed.py +++ b/sqlmesh/dbt/seed.py @@ -118,7 +118,9 @@ def cast(self, d: t.Any) -> t.Optional[int]: try: return int(d) except ValueError: - raise agate.exceptions.CastError('Can not parse value "%s" as Integer.' % d) + raise agate.exceptions.CastError( + 'Can not parse value "%s" as Integer.' % d + ) return super().cast(d) def jsonify(self, d: t.Any) -> str: diff --git a/sqlmesh/dbt/source.py b/sqlmesh/dbt/source.py index f08fa5744c..a284003dd5 100644 --- a/sqlmesh/dbt/source.py +++ b/sqlmesh/dbt/source.py @@ -98,7 +98,9 @@ def canonical_name(self, context: DbtContext) -> str: def relation_info(self) -> AttributeDict: extras = {} external_location = ( - self.source_meta.get("external_location", None) if self.source_meta else None + self.source_meta.get("external_location", None) + if self.source_meta + else None ) if external_location: extras["external"] = external_location.replace("{name}", self.table_name) diff --git a/sqlmesh/dbt/target.py b/sqlmesh/dbt/target.py index 62683ecfac..d89b2b2bbf 100644 --- a/sqlmesh/dbt/target.py +++ b/sqlmesh/dbt/target.py @@ -5,30 +5,26 @@ from pathlib import Path from dbt.adapters.base import BaseRelation, Column -from pydantic import Field, AliasChoices - +from pydantic import AliasChoices, Field + +from sqlmesh.core.config.connection import (AthenaConnectionConfig, + BigQueryConnectionConfig, + BigQueryConnectionMethod, + BigQueryPriority, + ClickhouseConnectionConfig, + ConnectionConfig, + DatabricksConnectionConfig, + DuckDBConnectionConfig, + MSSQLConnectionConfig, + PostgresConnectionConfig, + RedshiftConnectionConfig, + SnowflakeConnectionConfig, + TrinoAuthenticationMethod, + TrinoConnectionConfig) from sqlmesh.core.console import get_console -from sqlmesh.core.config.connection import ( - AthenaConnectionConfig, - BigQueryConnectionConfig, - BigQueryConnectionMethod, - BigQueryPriority, - ClickhouseConnectionConfig, - ConnectionConfig, - DatabricksConnectionConfig, - DuckDBConnectionConfig, - MSSQLConnectionConfig, - PostgresConnectionConfig, - RedshiftConnectionConfig, - SnowflakeConnectionConfig, - TrinoAuthenticationMethod, - TrinoConnectionConfig, -) -from sqlmesh.core.model import ( - IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind, - IncrementalUnmanagedKind, -) +from sqlmesh.core.model import (IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, + IncrementalUnmanagedKind) from sqlmesh.core.schema_diff import NestedSupport from sqlmesh.dbt.common import DbtConfig from sqlmesh.dbt.relation import Policy @@ -213,7 +209,9 @@ def validate_authentication(cls, data: t.Any) -> t.Any: ) if "threads" in data and t.cast(int, data["threads"]) > 1: - get_console().log_warning("DuckDB does not support concurrency - setting threads to 1.") + get_console().log_warning( + "DuckDB does not support concurrency - setting threads to 1." + ) return data @@ -307,7 +305,9 @@ def validate_authentication(cls, data: t.Any) -> t.Any: ): return data - raise ConfigError("No supported Snowflake authentication method found in target profile.") + raise ConfigError( + "No supported Snowflake authentication method found in target profile." + ) def default_incremental_strategy(self, kind: IncrementalKind) -> str: return "merge" @@ -531,7 +531,9 @@ def to_sqlmesh(self, **kwargs: t.Any) -> ConnectionConfig: access_token=self.token, concurrent_tasks=self.threads, catalog=self.database, - auth_type="databricks-oauth" if self.auth_type == "oauth" else self.auth_type, + auth_type=( + "databricks-oauth" if self.auth_type == "oauth" else self.auth_type + ), oauth_client_id=self.client_id, oauth_client_secret=self.client_secret, **kwargs, @@ -715,7 +717,9 @@ class MSSQLConfig(TargetConfig): retries: t.Optional[int] = None # Unused authentication parameters (not supported by pymssql) - windows_login: t.Optional[bool] = None # pymssql doesn't require this flag for Windows Auth + windows_login: t.Optional[bool] = ( + None # pymssql doesn't require this flag for Windows Auth + ) tenant_id: t.Optional[str] = None # Azure Active Directory auth client_id: t.Optional[str] = None # Azure Active Directory auth client_secret: t.Optional[str] = None # Azure Active Directory auth @@ -744,7 +748,9 @@ def validate_alias_fields(cls, data: t.Any) -> t.Any: @classmethod def _validate_authentication(cls, v: str) -> str: if v != "sql": - raise ConfigError("Only SQL and Windows Authentication are supported for SQL Server") + raise ConfigError( + "Only SQL and Windows Authentication are supported for SQL Server" + ) return v @field_validator("port") @@ -763,7 +769,8 @@ def column_class(cls) -> t.Type[Column]: from dbt.adapters.sqlserver.sqlserver_column import SQLServerColumn except ImportError: # <1.8.0 - from dbt.adapters.sqlserver.sql_server_column import SQLServerColumn # type: ignore + from dbt.adapters.sqlserver.sql_server_column import \ + SQLServerColumn # type: ignore return SQLServerColumn diff --git a/sqlmesh/dbt/test.py b/sqlmesh/dbt/test.py index c4a32b2189..1c3c50d878 100644 --- a/sqlmesh/dbt/test.py +++ b/sqlmesh/dbt/test.py @@ -6,15 +6,12 @@ from pathlib import Path from pydantic import Field + import sqlmesh.core.dialect as d from sqlmesh.core.audit import Audit, ModelAudit, StandaloneAudit from sqlmesh.core.node import DbtNodeInfo -from sqlmesh.dbt.common import ( - Dependencies, - GeneralConfig, - SqlStr, - sql_str_validator, -) +from sqlmesh.dbt.common import (Dependencies, GeneralConfig, SqlStr, + sql_str_validator) from sqlmesh.utils import AttributeDict from sqlmesh.utils.pydantic import field_validator @@ -61,9 +58,7 @@ class TestConfig(GeneralConfig): error_if: Conditional expression (default "!=0") to detect if error condition met (Not supported). """ - __test__ = ( - False # prevent pytest trying to collect this as a test class when it's imported in a test - ) + __test__ = False # prevent pytest trying to collect this as a test class when it's imported in a test # SQLMesh fields path: Path = Path() @@ -111,7 +106,11 @@ def _lowercase_name(cls, v: str) -> str: @property def canonical_name(self) -> str: - return f"{self.package_name}.{self.name}".lower() if self.package_name else self.name + return ( + f"{self.package_name}.{self.name}".lower() + if self.package_name + else self.name + ) @property def is_standalone(self) -> bool: @@ -159,7 +158,9 @@ def to_sqlmesh(self, context: DbtContext) -> Audit: } ) - query = d.jinja_query(self.sql.replace("**_dbt_generic_test_kwargs", self._kwargs())) + query = d.jinja_query( + self.sql.replace("**_dbt_generic_test_kwargs", self._kwargs()) + ) skip = not self.enabled blocking = self.severity == Severity.ERROR @@ -175,9 +176,13 @@ def to_sqlmesh(self, context: DbtContext) -> Audit: query=query, jinja_macros=jinja_macros, depends_on={ - model.canonical_name(context) for model in test_context.refs.values() + model.canonical_name(context) + for model in test_context.refs.values() }.union( - {source.canonical_name(context) for source in test_context.sources.values()} + { + source.canonical_name(context) + for source in test_context.sources.values() + } ), tags=self.tags, default_catalog=context.target.database, @@ -233,7 +238,10 @@ def relation_info(self) -> AttributeDict: @property def node_info(self) -> DbtNodeInfo: return DbtNodeInfo( - unique_id=self.unique_id, name=self.name, fqn=".".join(self.fqn), alias=self.alias + unique_id=self.unique_id, + name=self.name, + fqn=".".join(self.fqn), + alias=self.alias, ) diff --git a/sqlmesh/dbt/util.py b/sqlmesh/dbt/util.py index 0de16e3b3e..4f6c08d268 100644 --- a/sqlmesh/dbt/util.py +++ b/sqlmesh/dbt/util.py @@ -21,9 +21,11 @@ def _get_dbt_version() -> t.Tuple[int, int, int]: DBT_VERSION = _get_dbt_version() if DBT_VERSION >= (1, 8, 0): - from dbt_common.clients.agate_helper import table_from_data_flat, empty_table, as_matrix # type: ignore # noqa: F401 + from dbt_common.clients.agate_helper import ( # type: ignore # noqa: F401 + as_matrix, empty_table, table_from_data_flat) else: - from dbt.clients.agate_helper import table_from_data_flat, empty_table, as_matrix # type: ignore # noqa: F401 + from dbt.clients.agate_helper import ( # type: ignore # noqa: F401 + as_matrix, empty_table, table_from_data_flat) def pandas_to_agate(df: pd.DataFrame) -> agate.Table: diff --git a/sqlmesh/engines/spark/db_api/spark_session.py b/sqlmesh/engines/spark/db_api/spark_session.py index 04229f2a44..af7b439b78 100644 --- a/sqlmesh/engines/spark/db_api/spark_session.py +++ b/sqlmesh/engines/spark/db_api/spark_session.py @@ -4,7 +4,8 @@ import typing as t from threading import get_ident -from sqlmesh.engines.spark.db_api.errors import NotSupportedError, ProgrammingError +from sqlmesh.engines.spark.db_api.errors import (NotSupportedError, + ProgrammingError) if t.TYPE_CHECKING: from pyspark.sql import DataFrame, SparkSession @@ -60,7 +61,9 @@ def _fetch(self, size: t.Optional[int] = None) -> t.List[t.Tuple]: if size is None: size = len(self._last_output) - self._last_output_cursor - output = self._last_output[self._last_output_cursor : self._last_output_cursor + size] + output = self._last_output[ + self._last_output_cursor : self._last_output_cursor + size + ] self._last_output_cursor += size return output @@ -93,7 +96,9 @@ def cursor(self) -> SparkSessionCursor: from pyspark.errors import PySparkAttributeError try: - self.spark.sparkContext.setLocalProperty("spark.scheduler.pool", f"pool_{get_ident()}") + self.spark.sparkContext.setLocalProperty( + "spark.scheduler.pool", f"pool_{get_ident()}" + ) self.spark.conf.set("spark.sql.sources.partitionOverwriteMode", "dynamic") self.spark.conf.set("hive.exec.dynamic.partition", "true") self.spark.conf.set("hive.exec.dynamic.partition.mode", "nonstrict") @@ -123,7 +128,9 @@ def close(self) -> None: pass -def connection(spark: SparkSession, catalog: t.Optional[str] = None) -> SparkSessionConnection: +def connection( + spark: SparkSession, catalog: t.Optional[str] = None +) -> SparkSessionConnection: return SparkSessionConnection(spark, catalog) diff --git a/sqlmesh/integrations/dlt.py b/sqlmesh/integrations/dlt.py index d9cced8deb..9ea1d187b5 100644 --- a/sqlmesh/integrations/dlt.py +++ b/sqlmesh/integrations/dlt.py @@ -1,8 +1,10 @@ import typing as t -import click from datetime import datetime, timedelta, timezone + +import click from pydantic import ValidationError from sqlglot import exp, parse_one + from sqlmesh.core.config.connection import parse_connection_config from sqlmesh.core.context import Context from sqlmesh.utils.date import yesterday_ds @@ -37,6 +39,7 @@ def generate_dlt_models_and_settings( pipeline = dlt.attach(pipeline_name=pipeline_name, pipelines_dir=dlt_path or "") except CannotRestorePipelineException as e: from pathlib import Path + from dlt.common.pipeline import get_dlt_pipelines_dir searched_dir = dlt_path or get_dlt_pipelines_dir() @@ -79,7 +82,10 @@ def generate_dlt_models_and_settings( name: table for name, table in schema.tables.items() if ( - (has_table_seen_data(table) and not name.startswith(schema._dlt_tables_prefix)) + ( + has_table_seen_data(table) + and not name.startswith(schema._dlt_tables_prefix) + ) or name == schema.loads_table_name ) and (name in tables if tables else True) @@ -92,7 +98,9 @@ def generate_dlt_models_and_settings( # is_complete_column returns true if column contains a name and a data type for col in filter(is_complete_column, table["columns"].values()): - dlt_columns[col["name"]] = exp.DataType.build(str(col["data_type"]), dialect=dialect) + dlt_columns[col["name"]] = exp.DataType.build( + str(col["data_type"]), dialect=dialect + ) if col.get("primary_key"): primary_key.append(str(col["name"])) @@ -123,7 +131,9 @@ def generate_dlt_models_and_settings( if isinstance(column, str) ] select_columns = ( - ",\n".join(f" {column_name}" for column_name in column_types) if column_types else "" + ",\n".join(f" {column_name}" for column_name in column_types) + if column_types + else "" ) grain = f"\n grain ({', '.join(primary_key)})," if primary_key else "" @@ -160,7 +170,9 @@ def generate_dlt_models( if not tables and not force: existing_models = [m.name for m in context.models.values()] - sqlmesh_models = {model for model in sqlmesh_models if model[0] not in existing_models} + sqlmesh_models = { + model for model in sqlmesh_models if model[0] not in existing_models + } if sqlmesh_models: _create_object_files( @@ -183,7 +195,9 @@ def generate_incremental_model( ) -> str: """Generate the SQL definition for an incremental model.""" - time_column = parse_one(f"to_timestamp(CAST({load_id} AS DOUBLE))").sql(dialect=dialect) + time_column = parse_one(f"to_timestamp(CAST({load_id} AS DOUBLE))").sql( + dialect=dialect + ) from_clause = f"{from_table} as c" if parent_table: @@ -232,7 +246,11 @@ def format_config(configs: t.Dict[str, str], db_type: str) -> str: invalid_fields.append(error.get("loc", [])[0]) return "\n".join( - [f" {key}: {value}" for key, value in config.items() if key not in invalid_fields] + [ + f" {key}: {value}" + for key, value in config.items() + if key not in invalid_fields + ] ) diff --git a/sqlmesh/integrations/github/cicd/command.py b/sqlmesh/integrations/github/cicd/command.py index 5506d4917b..bf622544bb 100644 --- a/sqlmesh/integrations/github/cicd/command.py +++ b/sqlmesh/integrations/github/cicd/command.py @@ -6,14 +6,13 @@ import click from sqlmesh.core.analytics import cli_analytics -from sqlmesh.core.console import set_console, MarkdownConsole -from sqlmesh.integrations.github.cicd.controller import ( - GithubCheckConclusion, - GithubCheckStatus, - GithubController, - TestFailure, -) -from sqlmesh.utils.errors import CICDBotError, ConflictingPlanError, PlanError, LinterError +from sqlmesh.core.console import MarkdownConsole, set_console +from sqlmesh.integrations.github.cicd.controller import (GithubCheckConclusion, + GithubCheckStatus, + GithubController, + TestFailure) +from sqlmesh.utils.errors import (CICDBotError, ConflictingPlanError, + LinterError, PlanError) logger = logging.getLogger(__name__) @@ -37,7 +36,9 @@ def github(ctx: click.Context, token: str, full_logs: bool = False) -> None: # which can result in surprise newlines when outputting dates to backfill set_console( MarkdownConsole( - width=1000, warning_capture_only=not full_logs, error_capture_only=not full_logs + width=1000, + warning_capture_only=not full_logs, + error_capture_only=not full_logs, ) ) ctx.obj["github"] = GithubController( @@ -116,14 +117,18 @@ def _run_linter(controller: GithubController) -> bool: def run_tests(ctx: click.Context) -> None: """Runs the unit tests""" if not _run_tests(ctx.obj["github"]): - raise CICDBotError("Failed to run tests. See Pull Requests Checks for more information.") + raise CICDBotError( + "Failed to run tests. See Pull Requests Checks for more information." + ) def _update_pr_environment(controller: GithubController) -> bool: controller.update_pr_environment_check(status=GithubCheckStatus.IN_PROGRESS) try: controller.update_pr_environment() - conclusion = controller.update_pr_environment_check(status=GithubCheckStatus.COMPLETED) + conclusion = controller.update_pr_environment_check( + status=GithubCheckStatus.COMPLETED + ) return conclusion is not None and conclusion.is_success except Exception as e: logger.exception("Error occurred when updating PR environment") @@ -253,7 +258,8 @@ def _run_all(controller: GithubController) -> None: if controller.do_required_approval_check: if has_required_approval: controller.update_required_approval_check( - status=GithubCheckStatus.COMPLETED, conclusion=GithubCheckConclusion.SKIPPED + status=GithubCheckStatus.COMPLETED, + conclusion=GithubCheckConclusion.SKIPPED, ) else: controller.update_required_approval_check(status=GithubCheckStatus.QUEUED) @@ -288,23 +294,21 @@ def _run_all(controller: GithubController) -> None: status=GithubCheckStatus.COMPLETED, conclusion=GithubCheckConclusion.SKIPPED ) deployed_to_prod = False - if has_required_approval and prod_plan_generated and controller.pr_targets_prod_branch: + if ( + has_required_approval + and prod_plan_generated + and controller.pr_targets_prod_branch + ): deployed_to_prod = _deploy_production(controller) elif is_auto_deploying_prod: if controller.deploy_command_enabled and not has_required_approval: skip_reason = "Skipped Deploying to Production because a `/deploy` command has not been detected yet" elif controller.do_required_approval_check and not has_required_approval: - skip_reason = ( - "Skipped Deploying to Production because a required approver has not approved" - ) + skip_reason = "Skipped Deploying to Production because a required approver has not approved" elif not pr_environment_updated: - skip_reason = ( - "Skipped Deploying to Production because the PR environment was not updated" - ) + skip_reason = "Skipped Deploying to Production because the PR environment was not updated" elif not prod_plan_generated: - skip_reason = ( - "Skipped Deploying to Production because the production plan could not be generated" - ) + skip_reason = "Skipped Deploying to Production because the production plan could not be generated" else: skip_reason = "Skipped Deploying to Production for an unknown reason" controller.update_prod_environment_check( @@ -315,7 +319,11 @@ def _run_all(controller: GithubController) -> None: if ( not pr_environment_updated or not prod_plan_generated - or (has_required_approval and controller.pr_targets_prod_branch and not deployed_to_prod) + or ( + has_required_approval + and controller.pr_targets_prod_branch + and not deployed_to_prod + ) ): raise CICDBotError( "A step of the run-all check failed. See Pull Requests Checks for more information." diff --git a/sqlmesh/integrations/github/cicd/config.py b/sqlmesh/integrations/github/cicd/config.py index 2a2e2efa49..725f0f306d 100644 --- a/sqlmesh/integrations/github/cicd/config.py +++ b/sqlmesh/integrations/github/cicd/config.py @@ -5,9 +5,9 @@ from sqlmesh.core.config import CategorizerConfig from sqlmesh.core.config.base import BaseConfig +from sqlmesh.core.console import get_console from sqlmesh.utils.date import TimeLike from sqlmesh.utils.pydantic import model_validator -from sqlmesh.core.console import get_console class MergeMethod(str, Enum): @@ -29,7 +29,9 @@ class GithubCICDBotConfig(BaseConfig): default_pr_start: t.Optional[TimeLike] = None default_pr_preview_start: TimeLike = "yesterday" skip_pr_backfill_: t.Optional[bool] = Field(default=None, alias="skip_pr_backfill") - pr_include_unmodified_: t.Optional[bool] = Field(default=None, alias="pr_include_unmodified") + pr_include_unmodified_: t.Optional[bool] = Field( + default=None, alias="pr_include_unmodified" + ) run_on_deploy_to_prod: bool = False pr_environment_name: t.Optional[str] = None pr_min_intervals: t.Optional[int] = None @@ -47,9 +49,13 @@ def _validate(cls, data: t.Any) -> t.Any: return data if data.get("enable_deploy_command") and not data.get("merge_method"): - raise ValueError("merge_method must be set if enable_deploy_command is True") + raise ValueError( + "merge_method must be set if enable_deploy_command is True" + ) if data.get("command_namespace") and not data.get("enable_deploy_command"): - raise ValueError("enable_deploy_command must be set if command_namespace is set") + raise ValueError( + "enable_deploy_command must be set if command_namespace is set" + ) return data diff --git a/sqlmesh/integrations/github/cicd/controller.py b/sqlmesh/integrations/github/cicd/controller.py index 58b0664c2d..ef27744643 100644 --- a/sqlmesh/integrations/github/cicd/controller.py +++ b/sqlmesh/integrations/github/cicd/controller.py @@ -8,41 +8,33 @@ import re import traceback import typing as t -from enum import Enum -from pathlib import Path from dataclasses import dataclass +from enum import Enum from functools import cached_property +from pathlib import Path import requests +from sqlglot.errors import SqlglotError from sqlglot.helper import seq_get from sqlmesh.core import constants as c -from sqlmesh.core.console import SNAPSHOT_CHANGE_CATEGORY_STR, get_console, MarkdownConsole +from sqlmesh.core.config import Config +from sqlmesh.core.console import (SNAPSHOT_CHANGE_CATEGORY_STR, + MarkdownConsole, get_console) from sqlmesh.core.context import Context -from sqlmesh.core.test.result import ModelTextTestResult from sqlmesh.core.environment import Environment from sqlmesh.core.plan import Plan, PlanBuilder, SnapshotIntervals from sqlmesh.core.plan.definition import UserProvidedFlags -from sqlmesh.core.snapshot.definition import ( - Snapshot, - SnapshotChangeCategory, - SnapshotId, - SnapshotTableInfo, -) -from sqlglot.errors import SqlglotError +from sqlmesh.core.snapshot.definition import (Snapshot, SnapshotChangeCategory, + SnapshotId, SnapshotTableInfo) +from sqlmesh.core.test.result import ModelTextTestResult from sqlmesh.core.user import User -from sqlmesh.core.config import Config from sqlmesh.integrations.github.cicd.config import GithubCICDBotConfig -from sqlmesh.utils import word_characters_only, Verbosity +from sqlmesh.utils import Verbosity, word_characters_only from sqlmesh.utils.date import now -from sqlmesh.utils.errors import ( - CICDBotError, - NoChangesPlanError, - PlanError, - UncategorizedPlanError, - LinterError, - SQLMeshError, -) +from sqlmesh.utils.errors import (CICDBotError, LinterError, + NoChangesPlanError, PlanError, SQLMeshError, + UncategorizedPlanError) from sqlmesh.utils.pydantic import PydanticModel if t.TYPE_CHECKING: @@ -190,7 +182,9 @@ class BotCommand(Enum): DEPLOY_PROD = 2 @classmethod - def from_comment_body(cls, body: str, namespace: t.Optional[str] = None) -> BotCommand: + def from_comment_body( + cls, body: str, namespace: t.Optional[str] = None + ) -> BotCommand: body = body.strip() namespace = namespace.strip() if namespace else "" input_to_command = { @@ -294,7 +288,9 @@ def __init__( ) -> None: from github import Github - logger.debug(f"Initializing GithubController with paths: {paths} and config: {config}") + logger.debug( + f"Initializing GithubController with paths: {paths} and config: {config}" + ) self.config = config self._paths = paths @@ -310,11 +306,12 @@ def __init__( raise CICDBotError("Console must be a markdown console.") self._console = t.cast(MarkdownConsole, get_console()) - from github.Consts import DEFAULT_BASE_URL from github.Auth import Token + from github.Consts import DEFAULT_BASE_URL self._client: Github = client or Github( - base_url=os.environ.get("GITHUB_API_URL", DEFAULT_BASE_URL), auth=Token(self._token) + base_url=os.environ.get("GITHUB_API_URL", DEFAULT_BASE_URL), + auth=Token(self._token), ) self._repo: Repository = self._client.get_repo( @@ -323,7 +320,9 @@ def __init__( self._pull_request: PullRequest = self._repo.get_pull( self._event.pull_request_info.pr_number ) - self._issue: Issue = self._repo.get_issue(self._event.pull_request_info.pr_number) + self._issue: Issue = self._repo.get_issue( + self._event.pull_request_info.pr_number + ) self._reviews: t.Iterable[PullRequestReview] = self._pull_request.get_reviews() # TODO: The python module says that user names can be None and this is not currently handled self._approvers: t.Set[str] = { @@ -332,7 +331,9 @@ def __init__( if review.state.lower() == "approved" } logger.debug(f"Approvers: {', '.join(self._approvers)}") - self._context: Context = context or Context(paths=self._paths, config=self.config) + self._context: Context = context or Context( + paths=self._paths, config=self.config + ) # Bot config needs the context to be initialized logger.debug(f"Bot config: {self.bot_config.json(indent=2)}") @@ -360,7 +361,9 @@ def _required_approvers(self) -> t.List[User]: @property def _required_approvers_with_approval(self) -> t.List[User]: return [ - user for user in self._required_approvers if user.github_username in self._approvers + user + for user in self._required_approvers + if user.github_username in self._approvers ] @property @@ -368,7 +371,8 @@ def pr_environment_name(self) -> str: return Environment.sanitize_name( "_".join( [ - self.bot_config.pr_environment_name or self._event.pull_request_info.repo, + self.bot_config.pr_environment_name + or self._event.pull_request_info.repo, str(self._event.pull_request_info.pr_number), ] ) @@ -467,7 +471,9 @@ def bot_config(self) -> GithubCICDBotConfig: return bot_config @property - def modified_snapshots(self) -> t.Dict[SnapshotId, t.Union[Snapshot, SnapshotTableInfo]]: + def modified_snapshots( + self, + ) -> t.Dict[SnapshotId, t.Union[Snapshot, SnapshotTableInfo]]: return self.prod_plan_with_gaps.modified_snapshots @property @@ -483,7 +489,9 @@ def forward_only_plan(self) -> bool: default = self._context.config.plan.forward_only head_ref = self._pull_request.head.ref if isinstance(head_ref, str): - return head_ref.endswith(self.bot_config.forward_only_branch_suffix) or default + return ( + head_ref.endswith(self.bot_config.forward_only_branch_suffix) or default + ) return default @classmethod @@ -577,9 +585,7 @@ def get_pr_environment_summary( heading = f":warning: Action Required to create or update PR Environment `{self.pr_environment_name}` :warning:" summary = self._get_pr_environment_summary_action_required(exception) elif conclusion.is_failure: - heading = ( - f":x: Failed to create or update PR Environment `{self.pr_environment_name}` :x:" - ) + heading = f":x: Failed to create or update PR Environment `{self.pr_environment_name}` :x:" summary = self._get_pr_environment_summary_failure(exception) elif conclusion.is_skipped: heading = f":next_track_button: Skipped creating or updating PR Environment `{self.pr_environment_name}` :next_track_button:" @@ -607,7 +613,9 @@ def _get_pr_environment_summary_success(self) -> str: return summary - def _get_pr_environment_summary_skipped(self, exception: t.Optional[Exception] = None) -> str: + def _get_pr_environment_summary_skipped( + self, exception: t.Optional[Exception] = None + ) -> str: if isinstance(exception, NoChangesPlanError): skip_reason = "No changes were detected compared to the prod environment." elif isinstance(exception, TestFailure): @@ -622,7 +630,9 @@ def _get_pr_environment_summary_action_required( ) -> str: plan = self.pr_plan_or_none if isinstance(exception, UncategorizedPlanError) and plan: - failure_msg = f"The following models could not be categorized automatically:\n" + failure_msg = ( + f"The following models could not be categorized automatically:\n" + ) for snapshot in plan.uncategorized: failure_msg += f"- {snapshot.name}\n" failure_msg += ( @@ -634,7 +644,9 @@ def _get_pr_environment_summary_action_required( return failure_msg - def _get_pr_environment_summary_failure(self, exception: t.Optional[Exception] = None) -> str: + def _get_pr_environment_summary_failure( + self, exception: t.Optional[Exception] = None + ) -> str: console_output = self._console.consume_captured_output() failure_msg = "" @@ -679,7 +691,11 @@ def run_linter(self) -> None: def _get_or_create_comment(self, header: str = BOT_HEADER_MSG) -> IssueComment: comment = seq_get( - [comment for comment in self._issue.get_comments() if header in comment.body], + [ + comment + for comment in self._issue.get_comments() + if header in comment.body + ], 0, ) if not comment: @@ -713,7 +729,9 @@ def _get_merge_state_status(self) -> MergeStateStatus: ) if request.status_code == 200: merge_status = MergeStateStatus( - request.json()["data"]["repository"]["pullRequest"]["mergeStateStatus"].lower() + request.json()["data"]["repository"]["pullRequest"][ + "mergeStateStatus" + ].lower() ) logger.debug(f"Merge state status: {merge_status.value}") return merge_status @@ -750,7 +768,9 @@ def update_pr_environment(self) -> None: self._context.apply(self.pr_plan) # will raise if PR environment creation fails # update PR info comment - vde_title = "- :eyes: To **review** this PR's changes, use virtual data environment:" + vde_title = ( + "- :eyes: To **review** this PR's changes, use virtual data environment:" + ) comment_value = f"{vde_title}\n - `{self.pr_environment_name}`" if self.bot_config.enable_deploy_command: full_command = f"{self.bot_config.command_namespace or ''}/deploy" @@ -774,7 +794,10 @@ def deploy_to_prod(self) -> None: "PR is already merged and this event was triggered prior to the merge." ) merge_status = self._get_merge_state_status() - if self.bot_config.check_if_blocked_on_deploy_to_prod and merge_status.is_blocked: + if ( + self.bot_config.check_if_blocked_on_deploy_to_prod + and merge_status.is_blocked + ): raise CICDBotError( "Branch protection or ruleset requirement is likely not satisfied, e.g. missing CODEOWNERS approval. " "Please check PR and resolve any issues. To disable this check, set `check_if_blocked_on_deploy_to_prod` to false in the bot configuration." @@ -792,9 +815,7 @@ def deploy_to_prod(self) -> None: """ if self.forward_only_plan: - plan_summary = ( - f"{self.get_forward_only_plan_post_deployment_tip(self.prod_plan)}\n{plan_summary}" - ) + plan_summary = f"{self.get_forward_only_plan_post_deployment_tip(self.prod_plan)}\n{plan_summary}" self.update_sqlmesh_comment_info( value=plan_summary, @@ -837,7 +858,9 @@ def _update_check( full_summary = full_summary or title summary, text, *truncated = self._chunk_up_api_message(full_summary) + [None] if truncated and truncated[0] is not None: - logger.warning(f"Summary was too long so we truncated it. Full text: {full_summary}") + logger.warning( + f"Summary was too long so we truncated it. Full text: {full_summary}" + ) kwargs["output"] = {"title": title, "summary": summary} if text: kwargs["output"]["text"] = text @@ -858,7 +881,9 @@ def _update_check( } ) else: - logger.debug(f"Did not find check run in mapping so creating it. Name: {name}") + logger.debug( + f"Did not find check run in mapping so creating it. Name: {name}" + ) self._check_run_mapping[name] = self._repo.create_check_run(**kwargs) else: # Output the summary using print() so the newlines are resolved and the result can easily @@ -869,7 +894,8 @@ def _update_check( if conclusion: self._append_output( - word_characters_only(name.replace("SQLMesh - ", "").lower()), conclusion.value + word_characters_only(name.replace("SQLMesh - ", "").lower()), + conclusion.value, ) def _update_check_handler( @@ -879,7 +905,8 @@ def _update_check_handler( conclusion: t.Optional[GithubCheckConclusion], status_handler: t.Callable[[GithubCheckStatus], t.Tuple[str, t.Optional[str]]], conclusion_handler: t.Callable[ - [GithubCheckConclusion], t.Tuple[GithubCheckConclusion, str, t.Optional[str]] + [GithubCheckConclusion], + t.Tuple[GithubCheckConclusion, str, t.Optional[str]], ], ) -> None: if conclusion: @@ -948,7 +975,9 @@ def conclusion_handler( self._context.test_connection_config._engine_adapter.DIALECT, ) test_summary = self._console.consume_captured_output() - test_title = "Tests Passed" if result.wasSuccessful() else "Tests Failed" + test_title = ( + "Tests Passed" if result.wasSuccessful() else "Tests Failed" + ) test_conclusion = ( GithubCheckConclusion.SUCCESS if result.wasSuccessful() @@ -976,7 +1005,9 @@ def conclusion_handler( ) def update_required_approval_check( - self, status: GithubCheckStatus, conclusion: t.Optional[GithubCheckConclusion] = None + self, + status: GithubCheckStatus, + conclusion: t.Optional[GithubCheckConclusion] = None, ) -> None: """ Updates the status of the merge commit for the required approval. @@ -1051,7 +1082,9 @@ def conclusion_handler( GithubCheckStatus.IN_PROGRESS: f":rocket: Creating or Updating PR Environment `{self.pr_environment_name}`", }[status], ), - conclusion_handler=functools.partial(conclusion_handler, exception=exception), + conclusion_handler=functools.partial( + conclusion_handler, exception=exception + ), ) return conclusion @@ -1129,7 +1162,8 @@ def conclusion_handler( elif conclusion.is_failure: captured_errors = self._console.consume_captured_errors() summary = ( - captured_errors or f"{title}\n\n**Error:**\n```\n{traceback.format_exc()}\n```" + captured_errors + or f"{title}\n\n**Error:**\n```\n{traceback.format_exc()}\n```" ) elif conclusion.is_action_required: if plan_error: @@ -1137,7 +1171,9 @@ def conclusion_handler( else: summary = "Got an action required conclusion but no plan error was provided. This is unexpected." else: - summary = "**Generated Prod Plan**\n" + self.get_plan_summary(self.prod_plan) + summary = "**Generated Prod Plan**\n" + self.get_plan_summary( + self.prod_plan + ) return conclusion, title, summary @@ -1152,7 +1188,9 @@ def conclusion_handler( }[status], None, ), - conclusion_handler=functools.partial(conclusion_handler, skip_reason=skip_reason), + conclusion_handler=functools.partial( + conclusion_handler, skip_reason=skip_reason + ), ) def try_merge_pr(self) -> None: @@ -1161,7 +1199,9 @@ def try_merge_pr(self) -> None: performed """ if self.bot_config.merge_method: - logger.debug(f"Merging PR with merge method: {self.bot_config.merge_method.value}") + logger.debug( + f"Merging PR with merge method: {self.bot_config.merge_method.value}" + ) self._pull_request.merge(merge_method=self.bot_config.merge_method.value) else: logger.debug("No merge method defined so skipping merge") @@ -1175,7 +1215,9 @@ def get_command_from_comment(self) -> BotCommand: return BotCommand.INVALID if self._event.pull_request_comment_body is None: raise CICDBotError("Unable to get comment body") - logger.debug(f"Getting command from comment body: {self._event.pull_request_comment_body}") + logger.debug( + f"Getting command from comment body: {self._event.pull_request_comment_body}" + ) return BotCommand.from_comment_body( self._event.pull_request_comment_body, self.bot_config.command_namespace ) @@ -1272,11 +1314,21 @@ def _generate_pr_environment_summary_list(self, plan: Plan) -> str: sections = [ ("### Added", [r for r in table_records if r.is_added]), ("### Removed", [r for r in table_records if r.is_removed]), - ("### Directly Modified", [r for r in table_records if r.is_directly_modified]), - ("### Indirectly Modified", [r for r in table_records if r.is_indirectly_modified]), + ( + "### Directly Modified", + [r for r in table_records if r.is_directly_modified], + ), + ( + "### Indirectly Modified", + [r for r in table_records if r.is_indirectly_modified], + ), ( "### Metadata Updated", - [r for r in table_records if r.is_metadata_updated and not r.is_modified], + [ + r + for r in table_records + if r.is_metadata_updated and not r.is_modified + ], ), ] @@ -1392,7 +1444,11 @@ def loaded_intervals_rendered(self) -> str: @property def missing_intervals(self) -> t.Optional[SnapshotIntervals]: return next( - (si for si in self.plan.missing_intervals if si.snapshot_id == self.snapshot_id), + ( + si + for si in self.plan.missing_intervals + if si.snapshot_id == self.snapshot_id + ), None, ) @@ -1415,7 +1471,9 @@ def as_markdown_list_item(self) -> str: # note: this is to re-use the '[recreate view]' and '[full refresh]' text and keep it in sync with updates to the CLI # it doesnt actually use the passed intervals, those are handled differently - how_applied = _format_missing_intervals(self.snapshot, self.loaded_intervals) + how_applied = _format_missing_intervals( + self.snapshot, self.loaded_intervals + ) how_applied_str = f" [{how_applied}]" if how_applied else "" diff --git a/sqlmesh/integrations/slack.py b/sqlmesh/integrations/slack.py index 495978b984..6c099c814b 100644 --- a/sqlmesh/integrations/slack.py +++ b/sqlmesh/integrations/slack.py @@ -5,7 +5,6 @@ from enum import Enum from textwrap import dedent - SLACK_MAX_TEXT_LENGTH = 3000 SLACK_MAX_ALERT_PREVIEW_BLOCKS = 5 SLACK_MAX_ATTACHMENTS_BLOCKS = 50 @@ -47,7 +46,10 @@ def add_secondary_blocks(self, *blocks: TSlackBlock) -> "SlackMessageComposer": are always displayed. NOTICE: attachments blocks are deprecated by Slack """ self.slack_message["attachments"][0]["blocks"].extend(blocks) - if len(self.slack_message["attachments"][0]["blocks"]) >= SLACK_MAX_ATTACHMENTS_BLOCKS: + if ( + len(self.slack_message["attachments"][0]["blocks"]) + >= SLACK_MAX_ATTACHMENTS_BLOCKS + ): raise ValueError("Too many attachments") return self @@ -67,7 +69,9 @@ def _introspect(self) -> "SlackMessageComposer": return self -def normalize_message(message: t.Union[str, t.List[str], t.Tuple[str], t.Set[str]]) -> str: +def normalize_message( + message: t.Union[str, t.List[str], t.Tuple[str], t.Set[str]], +) -> str: """Normalize message to fit Slack's max text length""" if isinstance(message, (list, tuple, set)): message = "\n".join(message) diff --git a/sqlmesh/lsp/api.py b/sqlmesh/lsp/api.py index 882ca9825b..ac6bd88a27 100644 --- a/sqlmesh/lsp/api.py +++ b/sqlmesh/lsp/api.py @@ -8,11 +8,11 @@ """ import typing as t + from pydantic import field_validator -from sqlmesh.lsp.custom import ( - CustomMethodRequestBaseClass, - CustomMethodResponseBaseClass, -) + +from sqlmesh.lsp.custom import (CustomMethodRequestBaseClass, + CustomMethodResponseBaseClass) from web.server.models import LineageColumn, Model, TableDiff API_FEATURE = "sqlmesh/api" diff --git a/sqlmesh/lsp/completions.py b/sqlmesh/lsp/completions.py index 93162b15a4..82e49539a8 100644 --- a/sqlmesh/lsp/completions.py +++ b/sqlmesh/lsp/completions.py @@ -1,13 +1,12 @@ +import typing as t from functools import lru_cache + from sqlglot import Dialect, Tokenizer -from sqlmesh.lsp.custom import ( - AllModelsResponse, - MacroCompletion, - ModelCompletion, -) + from sqlmesh import macro -import typing as t from sqlmesh.lsp.context import AuditTarget, LSPContext, ModelTarget +from sqlmesh.lsp.custom import (AllModelsResponse, MacroCompletion, + ModelCompletion) from sqlmesh.lsp.uri import URI from sqlmesh.utils.lineage import generate_markdown_description @@ -26,7 +25,9 @@ def get_sql_completions( # Get keywords from file content if provided file_keywords = set() if content: - file_keywords = extract_keywords_from_content(content, get_dialect(context, file_uri)) + file_keywords = extract_keywords_from_content( + content, get_dialect(context, file_uri) + ) # Combine keywords - SQL keywords first, then file keywords all_keywords = list(sql_keywords) + list(file_keywords - sql_keywords) @@ -88,7 +89,9 @@ def get_macros( return [MacroCompletion(name=name, description=doc) for name, doc in macros.items()] -def get_keywords(context: t.Optional[LSPContext], file_uri: t.Optional[URI]) -> t.Set[str]: +def get_keywords( + context: t.Optional[LSPContext], file_uri: t.Optional[URI] +) -> t.Set[str]: """ Return a list of sql keywords for a given file. If no context is provided, return ANSI SQL keywords. @@ -99,7 +102,11 @@ def get_keywords(context: t.Optional[LSPContext], file_uri: t.Optional[URI]) -> If both a context and a file_uri are provided, returns the keywords for the dialect of the model that the file belongs to. """ - if file_uri is not None and context is not None and file_uri.to_path() in context.map: + if ( + file_uri is not None + and context is not None + and file_uri.to_path() in context.map + ): file_info = context.map[file_uri.to_path()] # Handle ModelInfo objects @@ -142,11 +149,17 @@ def get_keywords_from_tokenizer(dialect: t.Optional[str] = None) -> t.Set[str]: return expanded_keywords -def get_dialect(context: t.Optional[LSPContext], file_uri: t.Optional[URI]) -> t.Optional[str]: +def get_dialect( + context: t.Optional[LSPContext], file_uri: t.Optional[URI] +) -> t.Optional[str]: """ Get the dialect for a given file. """ - if file_uri is not None and context is not None and file_uri.to_path() in context.map: + if ( + file_uri is not None + and context is not None + and file_uri.to_path() in context.map + ): file_info = context.map[file_uri.to_path()] # Handle ModelInfo objects @@ -167,7 +180,9 @@ def get_dialect(context: t.Optional[LSPContext], file_uri: t.Optional[URI]) -> t return None -def extract_keywords_from_content(content: str, dialect: t.Optional[str] = None) -> t.Set[str]: +def extract_keywords_from_content( + content: str, dialect: t.Optional[str] = None +) -> t.Set[str]: """ Extract identifiers from SQL content using the tokenizer. diff --git a/sqlmesh/lsp/context.py b/sqlmesh/lsp/context.py index a94db7c421..f03ff55c23 100644 --- a/sqlmesh/lsp/context.py +++ b/sqlmesh/lsp/context.py @@ -1,19 +1,21 @@ +import typing as t from dataclasses import dataclass from pathlib import Path + +from lsprotocol import types from pygls.server import LanguageServer + from sqlmesh.core.context import Context -import typing as t -from sqlmesh.core.linter.rule import Range -from sqlmesh.core.model.definition import SqlModel, ExternalModel from sqlmesh.core.linter.definition import AnnotatedRuleViolation +from sqlmesh.core.linter.rule import Range +from sqlmesh.core.model.definition import ExternalModel, SqlModel from sqlmesh.core.schema_loader import get_columns from sqlmesh.lsp.commands import EXTERNAL_MODEL_UPDATE_COLUMNS -from sqlmesh.lsp.custom import ModelForRendering, TestEntry, RunTestResponse -from sqlmesh.lsp.custom import AllModelsResponse, RenderModelEntry -from sqlmesh.lsp.tests_ranges import get_test_ranges +from sqlmesh.lsp.custom import (AllModelsResponse, ModelForRendering, + RenderModelEntry, RunTestResponse, TestEntry) from sqlmesh.lsp.helpers import to_lsp_range +from sqlmesh.lsp.tests_ranges import get_test_ranges from sqlmesh.lsp.uri import URI -from lsprotocol import types from sqlmesh.utils import yaml from sqlmesh.utils.lineage import get_yaml_model_name_ranges @@ -88,7 +90,9 @@ def list_workspace_tests(self) -> t.List[TestEntry]: TestEntry( name=test.test_name, uri=URI.from_path(test.path).value, - range=test_uris.get(URI.from_path(test.path).value, {}).get(test.test_name), + range=test_uris.get(URI.from_path(test.path).value, {}).get( + test.test_name + ), ) for test in tests ] @@ -129,7 +133,9 @@ def run_test(self, uri: URI, test_name: str) -> RunTestResponse: match_patterns=[test_name], ) if results.testsRun != 1: - raise ValueError(f"Expected to run 1 test, but ran {results.testsRun} tests.") + raise ValueError( + f"Expected to run 1 test, but ran {results.testsRun} tests." + ) if len(results.successes) == 1: return RunTestResponse(success=True) return RunTestResponse( @@ -464,15 +470,19 @@ def diagnostic_to_lsp_diagnostic( return types.Diagnostic( range=diagnostic_range, message=diagnostic.violation_msg, - severity=types.DiagnosticSeverity.Error - if diagnostic.violation_type == "error" - else types.DiagnosticSeverity.Warning, + severity=( + types.DiagnosticSeverity.Error + if diagnostic.violation_type == "error" + else types.DiagnosticSeverity.Warning + ), source="sqlmesh", code=diagnostic.rule.name, code_description=types.CodeDescription(href=rule_uri), ) - def update_external_model_columns(self, ls: LanguageServer, uri: URI, model_name: str) -> bool: + def update_external_model_columns( + self, ls: LanguageServer, uri: URI, model_name: str + ) -> bool: """ Update the columns for an external model in the YAML file. Returns True if changed, False if didn't because of the columns already being up to date. @@ -487,7 +497,9 @@ def update_external_model_columns(self, ls: LanguageServer, uri: URI, model_name f"Expected a list of models in {uri.to_path()}, but got {type(models).__name__}" ) - existing_model = next((model for model in models if model.get("name") == model_name), None) + existing_model = next( + (model for model in models if model.get("name") == model_name), None + ) if existing_model is None: raise ValueError(f"Could not find model {model_name} in {uri.to_path()}") @@ -509,7 +521,8 @@ def update_external_model_columns(self, ls: LanguageServer, uri: URI, model_name # Model index to update model_index = next( - (i for i, model in enumerate(models) if model.get("name") == model_name), None + (i for i, model in enumerate(models) if model.get("name") == model_name), + None, ) if model_index is None: raise ValueError(f"Could not find model {model_name} in {uri.to_path()}") diff --git a/sqlmesh/lsp/custom.py b/sqlmesh/lsp/custom.py index 84be43ee0e..f2ae733da6 100644 --- a/sqlmesh/lsp/custom.py +++ b/sqlmesh/lsp/custom.py @@ -1,6 +1,7 @@ -from lsprotocol import types import typing as t +from lsprotocol import types + from sqlmesh.core.linter.rule import Range from sqlmesh.utils.pydantic import PydanticModel diff --git a/sqlmesh/lsp/errors.py b/sqlmesh/lsp/errors.py index a9e778a555..941ff84d96 100644 --- a/sqlmesh/lsp/errors.py +++ b/sqlmesh/lsp/errors.py @@ -1,10 +1,9 @@ -from lsprotocol.types import Diagnostic, DiagnosticSeverity, Range, Position +import typing as t + +from lsprotocol.types import Diagnostic, DiagnosticSeverity, Position, Range from sqlmesh.lsp.uri import URI -from sqlmesh.utils.errors import ( - ConfigError, -) -import typing as t +from sqlmesh.utils.errors import ConfigError ContextFailedError = t.Union[str, ConfigError, Exception] diff --git a/sqlmesh/lsp/helpers.py b/sqlmesh/lsp/helpers.py index 920a93f5c7..958ff771b2 100644 --- a/sqlmesh/lsp/helpers.py +++ b/sqlmesh/lsp/helpers.py @@ -1,9 +1,7 @@ -from lsprotocol.types import Range, Position +from lsprotocol.types import Position, Range -from sqlmesh.core.linter.helpers import ( - Range as SQLMeshRange, - Position as SQLMeshPosition, -) +from sqlmesh.core.linter.helpers import Position as SQLMeshPosition +from sqlmesh.core.linter.helpers import Range as SQLMeshRange def to_sqlmesh_position(position: Position) -> SQLMeshPosition: diff --git a/sqlmesh/lsp/hints.py b/sqlmesh/lsp/hints.py index 611ce8608d..7667d0d935 100644 --- a/sqlmesh/lsp/hints.py +++ b/sqlmesh/lsp/hints.py @@ -3,9 +3,9 @@ import typing as t from lsprotocol import types - from sqlglot import exp from sqlglot.optimizer.normalize_identifiers import normalize_identifiers + from sqlmesh.core.model.definition import SqlModel from sqlmesh.lsp.context import LSPContext, ModelTarget from sqlmesh.lsp.uri import URI diff --git a/sqlmesh/lsp/main.py b/sqlmesh/lsp/main.py index b5623f3ff8..c21ef5f96b 100755 --- a/sqlmesh/lsp/main.py +++ b/sqlmesh/lsp/main.py @@ -1,94 +1,68 @@ #!/usr/bin/env python """A Language Server Protocol (LSP) server for SQL with SQLMesh integration, refactored without globals.""" -from itertools import chain import logging import typing as t -from pathlib import Path import urllib.parse import uuid +from dataclasses import dataclass, field +from itertools import chain +from pathlib import Path +from typing import Union from lsprotocol import types -from lsprotocol.types import ( - WorkspaceDiagnosticRefreshRequest, - WorkspaceInlayHintRefreshRequest, -) +from lsprotocol.types import (WorkspaceDiagnosticRefreshRequest, + WorkspaceInlayHintRefreshRequest) from pygls.server import LanguageServer from sqlglot import exp + from sqlmesh._version import __version__ from sqlmesh.core.context import Context -from sqlmesh.utils.date import to_timestamp -from sqlmesh.lsp.api import ( - API_FEATURE, - ApiRequest, - ApiResponseGetColumnLineage, - ApiResponseGetLineage, - ApiResponseGetModels, - ApiResponseGetTableDiff, -) - +from sqlmesh.lsp.api import (API_FEATURE, ApiRequest, + ApiResponseGetColumnLineage, + ApiResponseGetLineage, ApiResponseGetModels, + ApiResponseGetTableDiff) from sqlmesh.lsp.commands import EXTERNAL_MODEL_UPDATE_COLUMNS from sqlmesh.lsp.completions import get_sql_completions -from sqlmesh.lsp.context import ( - LSPContext, - ModelTarget, -) -from sqlmesh.lsp.custom import ( - ALL_MODELS_FEATURE, - ALL_MODELS_FOR_RENDER_FEATURE, - RENDER_MODEL_FEATURE, - SUPPORTED_METHODS_FEATURE, - FORMAT_PROJECT_FEATURE, - GET_ENVIRONMENTS_FEATURE, - GET_MODELS_FEATURE, - AllModelsRequest, - AllModelsResponse, - AllModelsForRenderRequest, - AllModelsForRenderResponse, - CustomMethodResponseBaseClass, - RenderModelRequest, - RenderModelResponse, - SupportedMethodsRequest, - SupportedMethodsResponse, - FormatProjectRequest, - FormatProjectResponse, - CustomMethod, - LIST_WORKSPACE_TESTS_FEATURE, - ListWorkspaceTestsRequest, - ListWorkspaceTestsResponse, - LIST_DOCUMENT_TESTS_FEATURE, - ListDocumentTestsRequest, - ListDocumentTestsResponse, - RUN_TEST_FEATURE, - RunTestRequest, - RunTestResponse, - GetEnvironmentsRequest, - GetEnvironmentsResponse, - EnvironmentInfo, - GetModelsRequest, - GetModelsResponse, - ModelInfo, -) +from sqlmesh.lsp.context import LSPContext, ModelTarget +from sqlmesh.lsp.custom import (ALL_MODELS_FEATURE, + ALL_MODELS_FOR_RENDER_FEATURE, + FORMAT_PROJECT_FEATURE, + GET_ENVIRONMENTS_FEATURE, GET_MODELS_FEATURE, + LIST_DOCUMENT_TESTS_FEATURE, + LIST_WORKSPACE_TESTS_FEATURE, + RENDER_MODEL_FEATURE, RUN_TEST_FEATURE, + SUPPORTED_METHODS_FEATURE, + AllModelsForRenderRequest, + AllModelsForRenderResponse, AllModelsRequest, + AllModelsResponse, CustomMethod, + CustomMethodResponseBaseClass, EnvironmentInfo, + FormatProjectRequest, FormatProjectResponse, + GetEnvironmentsRequest, + GetEnvironmentsResponse, GetModelsRequest, + GetModelsResponse, ListDocumentTestsRequest, + ListDocumentTestsResponse, + ListWorkspaceTestsRequest, + ListWorkspaceTestsResponse, ModelInfo, + RenderModelRequest, RenderModelResponse, + RunTestRequest, RunTestResponse, + SupportedMethodsRequest, + SupportedMethodsResponse) from sqlmesh.lsp.errors import ContextFailedError, context_error_to_diagnostic from sqlmesh.lsp.helpers import to_lsp_range, to_sqlmesh_position from sqlmesh.lsp.hints import get_hints -from sqlmesh.lsp.reference import ( - CTEReference, - ModelReference, - get_references, - get_all_references, -) -from sqlmesh.lsp.rename import prepare_rename, rename_symbol, get_document_highlights +from sqlmesh.lsp.reference import (CTEReference, ModelReference, + get_all_references, get_references) +from sqlmesh.lsp.rename import (get_document_highlights, prepare_rename, + rename_symbol) from sqlmesh.lsp.uri import URI +from sqlmesh.utils.date import to_timestamp from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.lineage import ExternalModelReference from sqlmesh.utils.pydantic import PydanticModel from web.server.api.endpoints.lineage import column_lineage, model_lineage from web.server.api.endpoints.models import get_models from web.server.api.endpoints.table_diff import _process_sample_data -from typing import Union -from dataclasses import dataclass, field - from web.server.models import RowDiff, SchemaDiff, TableDiff @@ -220,7 +194,9 @@ def _run_test( return RunTestResponse(success=False, response_error=str(e)) # All the custom LSP methods are registered here and prefixed with _custom - def _custom_all_models(self, ls: LanguageServer, params: AllModelsRequest) -> AllModelsResponse: + def _custom_all_models( + self, ls: LanguageServer, params: AllModelsRequest + ) -> AllModelsResponse: uri = URI(params.textDocument.uri) # Get the document content content = None @@ -293,7 +269,9 @@ def _custom_get_environments( default_target_environment="", ) - def _custom_get_models(self, ls: LanguageServer, params: GetModelsRequest) -> GetModelsResponse: + def _custom_get_models( + self, ls: LanguageServer, params: GetModelsRequest + ) -> GetModelsResponse: """Get all models available for table diff.""" try: context = self._context_get_or_load() @@ -315,9 +293,7 @@ def _custom_get_models(self, ls: LanguageServer, params: GetModelsRequest) -> Ge models=[], ) - def _custom_api( - self, ls: LanguageServer, request: ApiRequest - ) -> t.Union[ + def _custom_api(self, ls: LanguageServer, request: ApiRequest) -> t.Union[ ApiResponseGetModels, ApiResponseGetColumnLineage, ApiResponseGetLineage, @@ -339,7 +315,9 @@ def _custom_api( # /api/lineage/{model} model_name = urllib.parse.unquote(path_parts[2]) lineage = model_lineage(model_name, context.context) - non_set_lineage = {k: v for k, v in lineage.items() if v is not None} + non_set_lineage = { + k: v for k, v in lineage.items() if v is not None + } return ApiResponseGetLineage(data=non_set_lineage) if len(path_parts) == 4: @@ -348,7 +326,9 @@ def _custom_api( column = urllib.parse.unquote(path_parts[3]) models_only = False if hasattr(request, "params"): - models_only = bool(getattr(request.params, "models_only", False)) + models_only = bool( + getattr(request.params, "models_only", False) + ) column_lineage_response = column_lineage( model_name, column, models_only, context.context ) @@ -368,21 +348,29 @@ def _custom_api( getattr(params, "model_or_snapshot", None) if params else None ) where = getattr(params, "where", None) if params else None - temp_schema = getattr(params, "temp_schema", None) if params else None + temp_schema = ( + getattr(params, "temp_schema", None) if params else None + ) limit = getattr(params, "limit", 20) if params else 20 table_diffs = context.context.table_diff( source=source, target=target, on=exp.condition(on) if on else None, - select_models={model_or_snapshot} if model_or_snapshot else None, + select_models=( + {model_or_snapshot} if model_or_snapshot else None + ), where=where, limit=limit, show=False, ) if table_diffs: - diff = table_diffs[0] if isinstance(table_diffs, list) else table_diffs + diff = ( + table_diffs[0] + if isinstance(table_diffs, list) + else table_diffs + ) _schema_diff = diff.schema_diff() _row_diff = diff.row_diff(temp_schema=temp_schema) @@ -397,17 +385,27 @@ def _custom_api( ) # create a readable column-centric sample data structure - processed_sample_data = _process_sample_data(_row_diff, source, target) + processed_sample_data = _process_sample_data( + _row_diff, source, target + ) row_diff = RowDiff( source=_row_diff.source, target=_row_diff.target, stats=_row_diff.stats, sample=_row_diff.sample.replace({np.nan: None}).to_dict(), - joined_sample=_row_diff.joined_sample.replace({np.nan: None}).to_dict(), - s_sample=_row_diff.s_sample.replace({np.nan: None}).to_dict(), - t_sample=_row_diff.t_sample.replace({np.nan: None}).to_dict(), - column_stats=_row_diff.column_stats.replace({np.nan: None}).to_dict(), + joined_sample=_row_diff.joined_sample.replace( + {np.nan: None} + ).to_dict(), + s_sample=_row_diff.s_sample.replace( + {np.nan: None} + ).to_dict(), + t_sample=_row_diff.t_sample.replace( + {np.nan: None} + ).to_dict(), + column_stats=_row_diff.column_stats.replace( + {np.nan: None} + ).to_dict(), source_count=_row_diff.source_count, target_count=_row_diff.target_count, count_pct_change=_row_diff.count_pct_change, @@ -464,7 +462,9 @@ def _reload_context_and_publish_diagnostics( # If there's no context, reset to NoContext and try to create one from scratch ls.log_trace("No partial context available, attempting fresh creation") self.context_state = NoContext() - self.has_raised_loading_error = False # Reset error flag to show new errors + self.has_raised_loading_error = ( + False # Reset error flag to show new errors + ) try: self._ensure_context_for_document(uri) # If successful, context_state will be ContextLoaded @@ -514,7 +514,9 @@ def _register_features(self) -> None: for name, method in self._supported_custom_methods.items(): def create_function_call(method_func: t.Callable) -> t.Callable: - def function_call(ls: LanguageServer, params: t.Any) -> t.Dict[str, t.Any]: + def function_call( + ls: LanguageServer, params: t.Any + ) -> t.Dict[str, t.Any]: try: response = method_func(ls, params) except Exception as e: @@ -526,7 +528,9 @@ def function_call(ls: LanguageServer, params: t.Any) -> t.Dict[str, t.Any]: self.server.feature(name)(create_function_call(method)) @self.server.command(EXTERNAL_MODEL_UPDATE_COLUMNS) - def command_external_models_update_columns(ls: LanguageServer, raw: t.Any) -> None: + def command_external_models_update_columns( + ls: LanguageServer, raw: t.Any + ) -> None: try: if not isinstance(raw, list): raise ValueError("Invalid command parameters") @@ -543,7 +547,9 @@ def command_external_models_update_columns(ls: LanguageServer, raw: t.Any) -> No if model is None: raise ValueError(f"External model '{model_name}' not found") if model._path is None: - raise ValueError(f"External model '{model_name}' does not have a file path") + raise ValueError( + f"External model '{model_name}' does not have a file path" + ) uri = URI.from_path(model._path) updated = context.update_external_model_columns( ls=ls, @@ -560,7 +566,9 @@ def command_external_models_update_columns(ls: LanguageServer, raw: t.Any) -> No f"Columns for '{model_name}' are already up to date", ) except Exception as e: - ls.show_message(f"Error executing command: {e}", types.MessageType.Error) + ls.show_message( + f"Error executing command: {e}", types.MessageType.Error + ) return None @self.server.feature(types.INITIALIZE) @@ -569,13 +577,19 @@ def initialize(ls: LanguageServer, params: types.InitializeParams) -> None: try: # Check the custom options if params.initialization_options: - options = InitializationOptions.model_validate(params.initialization_options) + options = InitializationOptions.model_validate( + params.initialization_options + ) if options.project_paths is not None: - self.specified_paths = [Path(path) for path in options.project_paths] + self.specified_paths = [ + Path(path) for path in options.project_paths + ] # Check if the client supports pull diagnostics if params.capabilities and params.capabilities.text_document: - diagnostics = getattr(params.capabilities.text_document, "diagnostic", None) + diagnostics = getattr( + params.capabilities.text_document, "diagnostic", None + ) if diagnostics: self.client_supports_pull_diagnostics = True ls.log_trace("Client supports pull diagnostics") @@ -588,7 +602,8 @@ def initialize(ls: LanguageServer, params: types.InitializeParams) -> None: if params.workspace_folders: # Store all workspace folders for later use self.workspace_folders = [ - Path(self._uri_to_path(folder.uri)) for folder in params.workspace_folders + Path(self._uri_to_path(folder.uri)) + for folder in params.workspace_folders ] # Try to find a SQLMesh config file in any workspace folder (only at the root level) @@ -606,7 +621,9 @@ def initialize(ls: LanguageServer, params: types.InitializeParams) -> None: ) @self.server.feature(types.TEXT_DOCUMENT_DID_OPEN) - def did_open(ls: LanguageServer, params: types.DidOpenTextDocumentParams) -> None: + def did_open( + ls: LanguageServer, params: types.DidOpenTextDocumentParams + ) -> None: uri = URI(params.text_document.uri) context = self._context_get_or_load(uri) @@ -619,9 +636,13 @@ def did_open(ls: LanguageServer, params: types.DidOpenTextDocumentParams) -> Non ) @self.server.feature(types.TEXT_DOCUMENT_DID_SAVE) - def did_save(ls: LanguageServer, params: types.DidSaveTextDocumentParams) -> None: + def did_save( + ls: LanguageServer, params: types.DidSaveTextDocumentParams + ) -> None: uri = URI(params.text_document.uri) - self._reload_context_and_publish_diagnostics(ls, uri, params.text_document.uri) + self._reload_context_and_publish_diagnostics( + ls, uri, params.text_document.uri + ) @self.server.feature(types.TEXT_DOCUMENT_FORMATTING) def formatting( @@ -659,7 +680,9 @@ def formatting( start=types.Position(line=0, character=0), end=types.Position( line=len(document.lines), - character=len(document.lines[-1]) if document.lines else 0, + character=( + len(document.lines[-1]) if document.lines else 0 + ), ), ), new_text=after, @@ -670,18 +693,25 @@ def formatting( return [] @self.server.feature(types.TEXT_DOCUMENT_HOVER) - def hover(ls: LanguageServer, params: types.HoverParams) -> t.Optional[types.Hover]: + def hover( + ls: LanguageServer, params: types.HoverParams + ) -> t.Optional[types.Hover]: """Provide hover information for an object.""" try: uri = URI(params.text_document.uri) context = self._context_get_or_load(uri) document = ls.workspace.get_text_document(params.text_document.uri) - references = get_references(context, uri, to_sqlmesh_position(params.position)) + references = get_references( + context, uri, to_sqlmesh_position(params.position) + ) if not references: return None reference = references[0] - if isinstance(reference, CTEReference) or not reference.markdown_description: + if ( + isinstance(reference, CTEReference) + or not reference.markdown_description + ): return None return types.Hover( contents=types.MarkupContent( @@ -723,7 +753,9 @@ def goto_definition( uri = URI(params.text_document.uri) context = self._context_get_or_load(uri) - references = get_references(context, uri, to_sqlmesh_position(params.position)) + references = get_references( + context, uri, to_sqlmesh_position(params.position) + ) location_links = [] for reference in references: # Use target_range if available (CTEs, Macros, and external models in YAML) @@ -749,7 +781,9 @@ def goto_definition( ) if reference.target_range is not None: target_range = to_lsp_range(reference.target_range) - target_selection_range = to_lsp_range(reference.target_range) + target_selection_range = to_lsp_range( + reference.target_range + ) else: # CTEs and Macros always have target_range target_range = to_lsp_range(reference.target_range) @@ -766,7 +800,9 @@ def goto_definition( ) return location_links except Exception as e: - ls.show_message(f"Error getting references: {e}", types.MessageType.Error) + ls.show_message( + f"Error getting references: {e}", types.MessageType.Error + ) return [] @self.server.feature(types.TEXT_DOCUMENT_REFERENCES) @@ -784,14 +820,18 @@ def find_references( # Convert references to Location objects locations = [ - types.Location(uri=URI.from_path(ref.path).value, range=to_lsp_range(ref.range)) + types.Location( + uri=URI.from_path(ref.path).value, range=to_lsp_range(ref.range) + ) for ref in all_references if ref.path is not None ] return locations if locations else None except Exception as e: - ls.show_message(f"Error getting locations: {e}", types.MessageType.Error) + ls.show_message( + f"Error getting locations: {e}", types.MessageType.Error + ) return None @self.server.feature(types.TEXT_DOCUMENT_PREPARE_RENAME) @@ -816,10 +856,14 @@ def rename_handler( try: uri = URI(params.text_document.uri) context = self._context_get_or_load(uri) - workspace_edit = rename_symbol(context, uri, params.position, params.new_name) + workspace_edit = rename_symbol( + context, uri, params.position, params.new_name + ) return workspace_edit except Exception as e: - ls.show_message(f"Error performing rename: {e}", types.MessageType.Error) + ls.show_message( + f"Error performing rename: {e}", types.MessageType.Error + ) return None @self.server.feature(types.TEXT_DOCUMENT_DOCUMENT_HIGHLIGHT) @@ -846,7 +890,10 @@ def diagnostic( diagnostics, result_id = self._get_diagnostics_for_uri(uri) # Check if client provided a previous result ID - if hasattr(params, "previous_result_id") and params.previous_result_id == result_id: + if ( + hasattr(params, "previous_result_id") + and params.previous_result_id == result_id + ): # Return unchanged report if diagnostics haven't changed return types.RelatedUnchangedDocumentDiagnosticReport( kind=types.DocumentDiagnosticReportKind.Unchanged, @@ -890,7 +937,10 @@ def workspace_diagnostic( # Check if we have a previous result ID for this file previous_result_id = None - if hasattr(params, "previous_result_ids") and params.previous_result_ids: + if ( + hasattr(params, "previous_result_ids") + and params.previous_result_ids + ): for prev in params.previous_result_ids: if prev.uri == uri.value: previous_result_id = prev.value @@ -952,7 +1002,9 @@ def code_action( return None @self.server.feature(types.TEXT_DOCUMENT_CODE_LENS) - def code_lens(ls: LanguageServer, params: types.CodeLensParams) -> t.List[types.CodeLens]: + def code_lens( + ls: LanguageServer, params: types.CodeLensParams + ) -> t.List[types.CodeLens]: try: uri = URI(params.text_document.uri) context = self._context_get_or_load(uri) @@ -964,7 +1016,9 @@ def code_lens(ls: LanguageServer, params: types.CodeLensParams) -> t.List[types. @self.server.feature( types.TEXT_DOCUMENT_COMPLETION, - types.CompletionOptions(trigger_characters=["@"]), # advertise "@" for macros + types.CompletionOptions( + trigger_characters=["@"] + ), # advertise "@" for macros ) def completion( ls: LanguageServer, params: types.CompletionParams @@ -993,12 +1047,15 @@ def completion( label=model.name, kind=types.CompletionItemKind.Reference, detail="SQLMesh Model", - documentation=types.MarkupContent( - kind=types.MarkupKind.Markdown, - value=model.description or "No description available", - ) - if model.description - else None, + documentation=( + types.MarkupContent( + kind=types.MarkupKind.Markdown, + value=model.description + or "No description available", + ) + if model.description + else None + ), ) ) # Add macro completions @@ -1041,7 +1098,9 @@ def completion( get_sql_completions(None, URI(params.text_document.uri)) return None - def _get_diagnostics_for_uri(self, uri: URI) -> t.Tuple[t.List[types.Diagnostic], str]: + def _get_diagnostics_for_uri( + self, uri: URI + ) -> t.Tuple[t.List[types.Diagnostic], str]: """Get diagnostics for a specific URI, returning (diagnostics, result_id). Since we no longer track version numbers, we always return 0 as the result_id. @@ -1050,11 +1109,14 @@ def _get_diagnostics_for_uri(self, uri: URI) -> t.Tuple[t.List[types.Diagnostic] try: context = self._context_get_or_load(uri) diagnostics = context.lint_model(uri) - return LSPContext.diagnostics_to_lsp_diagnostics( - diagnostics - ), self.context_state.version_id + return ( + LSPContext.diagnostics_to_lsp_diagnostics(diagnostics), + self.context_state.version_id, + ) except ConfigError as config_error: - diagnostic, error = context_error_to_diagnostic(config_error, uri_filter=uri) + diagnostic, error = context_error_to_diagnostic( + config_error, uri_filter=uri + ) if diagnostic: location, diag = diagnostic if location == uri.value: @@ -1171,7 +1233,10 @@ def _create_lsp_context(self, paths: t.List[Path]) -> t.Optional[LSPContext]: context = None if isinstance(self.context_state, ContextLoaded): context = self.context_state.lsp_context.context - elif isinstance(self.context_state, ContextFailed) and self.context_state.context: + elif ( + isinstance(self.context_state, ContextFailed) + and self.context_state.context + ): context = self.context_state.context self.context_state = ContextFailed(error=e, context=context) return None diff --git a/sqlmesh/lsp/reference.py b/sqlmesh/lsp/reference.py index 5881e1ece7..e6487b7901 100644 --- a/sqlmesh/lsp/reference.py +++ b/sqlmesh/lsp/reference.py @@ -1,27 +1,21 @@ +import ast +import inspect import typing as t from pathlib import Path -from sqlmesh.core.audit import StandaloneAudit -from sqlmesh.core.linter.helpers import ( - TokenPositionDetails, -) -from sqlmesh.core.linter.rule import Range, Position -from sqlmesh.core.model.definition import SqlModel -from sqlmesh.lsp.context import LSPContext, ModelTarget, AuditTarget from sqlglot import exp -from sqlmesh.lsp.uri import URI -from sqlmesh.utils.lineage import ( - MacroReference, - CTEReference, - Reference, - ModelReference, - extract_references_from_query, -) -import ast -from sqlmesh.core.model import Model from sqlmesh import macro -import inspect +from sqlmesh.core.audit import StandaloneAudit +from sqlmesh.core.linter.helpers import TokenPositionDetails +from sqlmesh.core.linter.rule import Position, Range +from sqlmesh.core.model import Model +from sqlmesh.core.model.definition import SqlModel +from sqlmesh.lsp.context import AuditTarget, LSPContext, ModelTarget +from sqlmesh.lsp.uri import URI +from sqlmesh.utils.lineage import (CTEReference, MacroReference, + ModelReference, Reference, + extract_references_from_query) def by_position(position: Position) -> t.Callable[[Reference], bool]: @@ -162,8 +156,10 @@ def get_macro_definitions_for_a_path( # Process based on whether it's a model or standalone audit if isinstance(file_info, ModelTarget): # It's a model - target: t.Optional[t.Union[Model, StandaloneAudit]] = lsp_context.context.get_model( - model_or_snapshot=file_info.names[0], raise_if_missing=False + target: t.Optional[t.Union[Model, StandaloneAudit]] = ( + lsp_context.context.get_model( + model_or_snapshot=file_info.names[0], raise_if_missing=False + ) ) if target is None or not isinstance(target, SqlModel): return [] @@ -281,7 +277,9 @@ def get_macro_reference( return None -def get_built_in_macro_reference(macro_name: str, macro_range: Range) -> t.Optional[Reference]: +def get_built_in_macro_reference( + macro_name: str, macro_range: Range +) -> t.Optional[Reference]: """ Get a reference to a built-in macro by its name. @@ -333,7 +331,8 @@ def get_model_find_all_references( model_at_position = next( filter( lambda ref: ( - isinstance(ref, ModelReference) and _position_within_range(position, ref.range) + isinstance(ref, ModelReference) + and _position_within_range(position, ref.range) ), get_model_definitions_for_a_path(lint_context, document_uri), ), @@ -386,7 +385,8 @@ def get_model_find_all_references( # Get model references that point to the target model matching_refs = filter( - lambda ref: isinstance(ref, ModelReference) and ref.path == target_model_path, + lambda ref: isinstance(ref, ModelReference) + and ref.path == target_model_path, get_model_definitions_for_a_path(lint_context, file_uri), ) @@ -488,7 +488,8 @@ def get_macro_find_all_references( macro_at_position = next( filter( lambda ref: ( - isinstance(ref, MacroReference) and _position_within_range(position, ref.range) + isinstance(ref, MacroReference) + and _position_within_range(position, ref.range) ), get_macro_definitions_for_a_path(lsp_context, document_uri), ), @@ -563,11 +564,15 @@ def get_all_references( return cte_references # Then try model references (across files) - if model_references := get_model_find_all_references(lint_context, document_uri, position): + if model_references := get_model_find_all_references( + lint_context, document_uri, position + ): return model_references # Finally try macro references (across files) - if macro_references := get_macro_find_all_references(lint_context, document_uri, position): + if macro_references := get_macro_find_all_references( + lint_context, document_uri, position + ): return macro_references return [] @@ -577,8 +582,14 @@ def _position_within_range(position: Position, range: Range) -> bool: """Check if a position is within a given range.""" return ( range.start.line < position.line - or (range.start.line == position.line and range.start.character <= position.character) + or ( + range.start.line == position.line + and range.start.character <= position.character + ) ) and ( range.end.line > position.line - or (range.end.line == position.line and range.end.character >= position.character) + or ( + range.end.line == position.line + and range.end.character >= position.character + ) ) diff --git a/sqlmesh/lsp/rename.py b/sqlmesh/lsp/rename.py index 5675c4efca..165f6a55d0 100644 --- a/sqlmesh/lsp/rename.py +++ b/sqlmesh/lsp/rename.py @@ -1,20 +1,13 @@ import typing as t -from lsprotocol.types import ( - Position, - TextEdit, - WorkspaceEdit, - PrepareRenameResult_Type1, - DocumentHighlight, - DocumentHighlightKind, -) + +from lsprotocol.types import (DocumentHighlight, DocumentHighlightKind, + Position, PrepareRenameResult_Type1, TextEdit, + WorkspaceEdit) from sqlmesh.lsp.context import LSPContext -from sqlmesh.lsp.helpers import to_sqlmesh_position, to_lsp_range -from sqlmesh.lsp.reference import ( - _position_within_range, - get_cte_references, - CTEReference, -) +from sqlmesh.lsp.helpers import to_lsp_range, to_sqlmesh_position +from sqlmesh.lsp.reference import (CTEReference, _position_within_range, + get_cte_references) from sqlmesh.lsp.uri import URI @@ -125,7 +118,9 @@ def get_document_highlights( List of DocumentHighlight objects or None if no symbol found """ # Check if there's a CTE at this position - cte_references = get_cte_references(lsp_context, document_uri, to_sqlmesh_position(position)) + cte_references = get_cte_references( + lsp_context, document_uri, to_sqlmesh_position(position) + ) if cte_references: highlights = [] for ref in cte_references: @@ -136,7 +131,9 @@ def get_document_highlights( else DocumentHighlightKind.Read ) - highlights.append(DocumentHighlight(range=to_lsp_range(ref.range), kind=kind)) + highlights.append( + DocumentHighlight(range=to_lsp_range(ref.range), kind=kind) + ) return highlights # For now, only CTEs are supported diff --git a/sqlmesh/lsp/tests_ranges.py b/sqlmesh/lsp/tests_ranges.py index cbcb33d8b6..e84ba52b4e 100644 --- a/sqlmesh/lsp/tests_ranges.py +++ b/sqlmesh/lsp/tests_ranges.py @@ -2,12 +2,13 @@ Provides helper functions to get ranges of tests in SQLMesh LSP. """ +import typing as t from pathlib import Path -from sqlmesh.core.linter.rule import Range, Position from ruamel import yaml from ruamel.yaml.comments import CommentedMap -import typing as t + +from sqlmesh.core.linter.rule import Position, Range def get_test_ranges( @@ -28,7 +29,9 @@ def get_test_ranges( data = yaml_obj.load(content) if not isinstance(data, dict): - raise ValueError("Invalid test file format: expected a dictionary at the top level.") + raise ValueError( + "Invalid test file format: expected a dictionary at the top level." + ) # For each top-level key (test name), find its range for test_name in data: @@ -58,7 +61,8 @@ def get_test_ranges( test_ranges[test_name] = Range( start=Position(line=start_line, character=start_col), end=Position( - line=end_line, character=len(lines[end_line]) if end_line < len(lines) else 0 + line=end_line, + character=len(lines[end_line]) if end_line < len(lines) else 0, ), ) diff --git a/sqlmesh/lsp/uri.py b/sqlmesh/lsp/uri.py index f8f0a495db..f53fb38fa8 100644 --- a/sqlmesh/lsp/uri.py +++ b/sqlmesh/lsp/uri.py @@ -1,6 +1,7 @@ +import typing as t from pathlib import Path + from pygls.uris import from_fs_path, to_fs_path -import typing as t class URI: diff --git a/sqlmesh/magics.py b/sqlmesh/magics.py index ed6a1b62de..caba353812 100644 --- a/sqlmesh/magics.py +++ b/sqlmesh/magics.py @@ -1,13 +1,12 @@ from __future__ import annotations -from io import StringIO - import functools import logging import typing as t -from argparse import Namespace, SUPPRESS +from argparse import SUPPRESS, Namespace from collections import defaultdict from copy import deepcopy +from io import StringIO from pathlib import Path from hyperscript import h @@ -17,27 +16,25 @@ except ImportError: from IPython.display import display -from IPython.core.magic import ( - Magics, - cell_magic, - line_cell_magic, - line_magic, - magics_class, -) -from IPython.core.magic_arguments import argument, magic_arguments, parse_argstring +from IPython.core.magic import (Magics, cell_magic, line_cell_magic, + line_magic, magics_class) +from IPython.core.magic_arguments import (argument, magic_arguments, + parse_argstring) from IPython.utils.process import arg_split from rich.jupyter import JupyterRenderable + from sqlmesh.cli.project_init import ProjectTemplate, init_example_project from sqlmesh.core import analytics from sqlmesh.core.config import load_configs from sqlmesh.core.config.connection import INIT_DISPLAY_INFO_TO_TYPE -from sqlmesh.core.console import create_console, set_console, configure_console +from sqlmesh.core.console import configure_console, create_console, set_console from sqlmesh.core.context import Context from sqlmesh.core.dialect import format_model_expressions, parse from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.test import ModelTestMetadata -from sqlmesh.utils import yaml, Verbosity, optional_import -from sqlmesh.utils.errors import MagicError, MissingContextException, SQLMeshError +from sqlmesh.utils import Verbosity, optional_import, yaml +from sqlmesh.utils.errors import (MagicError, MissingContextException, + SQLMeshError) logger = logging.getLogger(__name__) @@ -85,7 +82,9 @@ def wrapper(self: SQLMeshMagics, *args: t.Any, **kwargs: t.Any) -> None: parser._defaults = original_parser_defaults command_args = {k for k, v in parsed_args.__dict__.items() if v is not None} - analytics.collector.on_magic_command(command_name=magic_name, command_args=command_args) + analytics.collector.on_magic_command( + command_name=magic_name, command_args=command_args + ) func(self, context, *args, **kwargs) @@ -174,9 +173,13 @@ def _shell(self) -> t.Any: @argument("--gateway", type=str, help="The name of the gateway.") @argument("--ignore-warnings", action="store_true", help="Ignore warnings.") @argument("--debug", action="store_true", help="Enable debug mode.") - @argument("--log-file-dir", type=str, help="The directory to write the log file to.") @argument( - "--dotenv", type=str, help="Path to a custom .env file to load environment variables from." + "--log-file-dir", type=str, help="The directory to write the log file to." + ) + @argument( + "--dotenv", + type=str, + help="Path to a custom .env file to load environment variables from.", ) @line_magic def context(self, line: str) -> None: @@ -209,10 +212,16 @@ def context(self, line: str) -> None: logger.exception("Failed to initialize SQLMesh context") raise - context.console.log_success(f"SQLMesh project context set to: {', '.join(args.paths)}") + context.console.log_success( + f"SQLMesh project context set to: {', '.join(args.paths)}" + ) @magic_arguments() - @argument("path", type=str, help="The path where the new SQLMesh project should be created.") + @argument( + "path", + type=str, + help="The path where the new SQLMesh project should be created.", + ) @argument( "engine", type=str, @@ -335,11 +344,15 @@ def model(self, context: Context, line: str, sql: t.Optional[str] = None) -> Non @magic_arguments() @argument("model", type=str, help="The model.") - @argument("test_name", type=str, nargs="?", default=None, help="The test name to display") + @argument( + "test_name", type=str, nargs="?", default=None, help="The test name to display" + ) @argument("--ls", action="store_true", help="List tests associated with a model") @line_cell_magic @pass_sqlmesh_context - def test(self, context: Context, line: str, test_def_raw: t.Optional[str] = None) -> None: + def test( + self, context: Context, line: str, test_def_raw: t.Optional[str] = None + ) -> None: """Allow the user to list tests for a model, output a specific test, and then write their changes back""" args = parse_argstring(self.test, line) if not args.test_name and not args.ls: @@ -645,7 +658,11 @@ def evaluate(self, context: Context, line: str) -> None: help="Whether or not to use expand materialized models, defaults to False. If 'true', all referenced models are expanded as raw queries. If a comma-separated list of model names, only those models are expanded as raw queries.", ) @argument("--dialect", type=str, help="SQL dialect to render.") - @argument("--no-format", action="store_true", help="Disable fancy formatting of the query.") + @argument( + "--no-format", + action="store_true", + help="Disable fancy formatting of the query.", + ) @format_arguments @line_magic @pass_sqlmesh_context @@ -705,7 +722,12 @@ def fetchdf(self, context: Context, line: str, sql: str) -> None: self.display(df) @magic_arguments() - @argument("--file", "-f", type=str, help="An optional file path to write the HTML output to.") + @argument( + "--file", + "-f", + type=str, + help="An optional file path to write the HTML output to.", + ) @argument( "--select-model", type=str, @@ -899,8 +921,12 @@ def dlt_refresh(self, context: Context, line: str) -> None: context, args.pipeline, list(args.table or []), args.force, args.dlt_path ) if sqlmesh_models: - model_names = "\n".join([f"- {model_name}" for model_name in sqlmesh_models]) - context.console.log_success(f"Updated SQLMesh project with models:\n{model_names}") + model_names = "\n".join( + [f"- {model_name}" for model_name in sqlmesh_models] + ) + context.console.log_success( + f"Updated SQLMesh project with models:\n{model_names}" + ) else: context.console.log_success("All SQLMesh models are up to date.") @@ -968,7 +994,9 @@ def format(self, context: Context, line: str) -> bool: return context.format(**{k: v for k, v in format_opts.items() if v is not None}) @magic_arguments() - @argument("environment", type=str, help="The environment to diff local state against.") + @argument( + "environment", type=str, help="The environment to diff local state against." + ) @line_magic @pass_sqlmesh_context def diff(self, context: Context, line: str) -> None: @@ -1049,7 +1077,9 @@ def create_test(self, context: Context, line: str) -> None: variables = iter(args.var) if args.var else None context.create_test( args.model, - input_queries={k: v.strip('"') for k, v in dict(zip(queries, queries)).items()}, + input_queries={ + k: v.strip('"') for k, v in dict(zip(queries, queries)).items() + }, overwrite=args.overwrite, variables=dict(zip(variables, variables)) if variables else None, path=args.path, @@ -1094,7 +1124,10 @@ def run_test(self, context: Context, line: str) -> None: @magic_arguments() @argument( - "models", type=str, nargs="*", help="A model to audit. Multiple models can be audited." + "models", + type=str, + nargs="*", + help="A model to audit. Multiple models can be audited.", ) @argument("--start", "-s", type=str, help="Start date to audit.") @argument("--end", "-e", type=str, help="End date to audit.") @@ -1105,11 +1138,19 @@ def audit(self, context: Context, line: str) -> bool: """Run audit(s)""" args = parse_argstring(self.audit, line) return context.audit( - models=args.models, start=args.start, end=args.end, execution_time=args.execution_time + models=args.models, + start=args.start, + end=args.end, + execution_time=args.execution_time, ) @magic_arguments() - @argument("environment", nargs="?", type=str, help="The environment to check intervals for.") + @argument( + "environment", + nargs="?", + type=str, + help="The environment to check intervals for.", + ) @argument( "--no-signals", action="store_true", @@ -1159,7 +1200,9 @@ def check_intervals(self, context: Context, line: str) -> None: def info(self, context: Context, line: str) -> None: """Display SQLMesh project information.""" args = parse_argstring(self.info, line) - context.print_info(skip_connection=args.skip_connection, verbosity=Verbosity(args.verbose)) + context.print_info( + skip_connection=args.skip_connection, verbosity=Verbosity(args.verbose) + ) @magic_arguments() @line_magic diff --git a/sqlmesh/migrations/v0000_baseline.py b/sqlmesh/migrations/v0000_baseline.py index abd316fcfe..5cbc0ab65e 100644 --- a/sqlmesh/migrations/v0000_baseline.py +++ b/sqlmesh/migrations/v0000_baseline.py @@ -1,6 +1,7 @@ """The baseline migration script that sets up the initial state tables.""" from sqlglot import exp + from sqlmesh.utils.migration import blob_text_type, index_text_type @@ -74,7 +75,9 @@ def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore engine_adapter.create_state_table( snapshots_table, snapshots_columns_to_types, primary_key=("name", "identifier") ) - engine_adapter.create_index(snapshots_table, "_snapshots_name_version_idx", ("name", "version")) + engine_adapter.create_index( + snapshots_table, "_snapshots_name_version_idx", ("name", "version") + ) # Create the environments table and its indexes. engine_adapter.create_state_table( @@ -88,7 +91,9 @@ def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore engine_adapter.create_index( intervals_table, "_intervals_name_identifier_idx", ("name", "identifier") ) - engine_adapter.create_index(intervals_table, "_intervals_name_version_idx", ("name", "version")) + engine_adapter.create_index( + intervals_table, "_intervals_name_version_idx", ("name", "version") + ) def migrate_rows(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0063_change_signals.py b/sqlmesh/migrations/v0063_change_signals.py index bbced547fd..9152536be0 100644 --- a/sqlmesh/migrations/v0063_change_signals.py +++ b/sqlmesh/migrations/v0063_change_signals.py @@ -4,7 +4,7 @@ from sqlglot import exp, parse_one -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0064_join_when_matched_strings.py b/sqlmesh/migrations/v0064_join_when_matched_strings.py index ffd4c94913..2cec27c4cd 100644 --- a/sqlmesh/migrations/v0064_join_when_matched_strings.py +++ b/sqlmesh/migrations/v0064_join_when_matched_strings.py @@ -4,7 +4,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0069_update_dev_table_suffix.py b/sqlmesh/migrations/v0069_update_dev_table_suffix.py index f69aac434e..6fffe40cf4 100644 --- a/sqlmesh/migrations/v0069_update_dev_table_suffix.py +++ b/sqlmesh/migrations/v0069_update_dev_table_suffix.py @@ -4,7 +4,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore @@ -116,7 +116,9 @@ def migrate_rows(engine_adapter, schema, **kwargs): # type: ignore _update_snapshot(s) if previous_finalized_snapshots: - parsed_previous_finalized_snapshots = json.loads(previous_finalized_snapshots) + parsed_previous_finalized_snapshots = json.loads( + previous_finalized_snapshots + ) for s in parsed_previous_finalized_snapshots: _update_snapshot(s) @@ -133,9 +135,11 @@ def migrate_rows(engine_adapter, schema, **kwargs): # type: ignore "promoted_snapshot_ids": promoted_snapshot_ids, "suffix_target": suffix_target, "catalog_name_override": catalog_name_override, - "previous_finalized_snapshots": json.dumps(parsed_previous_finalized_snapshots) - if previous_finalized_snapshots - else None, + "previous_finalized_snapshots": ( + json.dumps(parsed_previous_finalized_snapshots) + if previous_finalized_snapshots + else None + ), "normalize_name": normalize_name, "requirements": requirements, } diff --git a/sqlmesh/migrations/v0071_add_dev_version_to_intervals.py b/sqlmesh/migrations/v0071_add_dev_version_to_intervals.py index 61a49dc0b9..a9492670a6 100644 --- a/sqlmesh/migrations/v0071_add_dev_version_to_intervals.py +++ b/sqlmesh/migrations/v0071_add_dev_version_to_intervals.py @@ -1,11 +1,12 @@ """Add dev version to the intervals table.""" -import typing as t import json +import typing as t import zlib from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type + +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore @@ -198,7 +199,9 @@ def _migrate_snapshots( for previous_version in parsed_snapshot.get("previous_versions", []): previous_identifier = get_identifier(previous_version) previous_dev_version = get_dev_version(previous_version) - snapshot_ids_to_dev_versions[(name, previous_identifier)] = previous_dev_version + snapshot_ids_to_dev_versions[(name, previous_identifier)] = ( + previous_dev_version + ) new_snapshots.append( { diff --git a/sqlmesh/migrations/v0073_remove_symbolic_disable_restatement.py b/sqlmesh/migrations/v0073_remove_symbolic_disable_restatement.py index 708693ed61..02bab02963 100644 --- a/sqlmesh/migrations/v0073_remove_symbolic_disable_restatement.py +++ b/sqlmesh/migrations/v0073_remove_symbolic_disable_restatement.py @@ -3,7 +3,8 @@ import json from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type + +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0075_remove_validate_query.py b/sqlmesh/migrations/v0075_remove_validate_query.py index 9fdcca7ea6..abc18158af 100644 --- a/sqlmesh/migrations/v0075_remove_validate_query.py +++ b/sqlmesh/migrations/v0075_remove_validate_query.py @@ -4,8 +4,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type -from sqlmesh.utils.migration import blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0078_warn_if_non_migratable_python_env.py b/sqlmesh/migrations/v0078_warn_if_non_migratable_python_env.py index adf1e96dd0..b1a8b27d51 100644 --- a/sqlmesh/migrations/v0078_warn_if_non_migratable_python_env.py +++ b/sqlmesh/migrations/v0078_warn_if_non_migratable_python_env.py @@ -73,7 +73,9 @@ def migrate_rows(engine_adapter, schema, **kwargs): # type: ignore # We use try-except here as a conservative measure to avoid any unexpected exceptions try: if on_virtual_update := node.get("on_virtual_update"): - metadata_hash_statements.extend(parse_expression(on_virtual_update, dialect)) + metadata_hash_statements.extend( + parse_expression(on_virtual_update, dialect) + ) for _, audit_args in func_call_validator(node.get("audits") or []): metadata_hash_statements.extend(audit_args.values()) @@ -92,7 +94,9 @@ def migrate_rows(engine_adapter, schema, **kwargs): # type: ignore for macro_name in extract_used_macros(metadata_hash_statements): serialized_macro = python_env.get(macro_name) - if isinstance(serialized_macro, dict) and not serialized_macro.get("is_metadata"): + if isinstance(serialized_macro, dict) and not serialized_macro.get( + "is_metadata" + ): get_console().log_warning(warning) return except Exception: diff --git a/sqlmesh/migrations/v0081_update_partitioned_by.py b/sqlmesh/migrations/v0081_update_partitioned_by.py index 8740285bf0..a72c2d054d 100644 --- a/sqlmesh/migrations/v0081_update_partitioned_by.py +++ b/sqlmesh/migrations/v0081_update_partitioned_by.py @@ -4,8 +4,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type -from sqlmesh.utils.migration import blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0085_deterministic_repr.py b/sqlmesh/migrations/v0085_deterministic_repr.py index 81cb0f194e..53bf64f5b1 100644 --- a/sqlmesh/migrations/v0085_deterministic_repr.py +++ b/sqlmesh/migrations/v0085_deterministic_repr.py @@ -10,8 +10,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type - +from sqlmesh.utils.migration import blob_text_type, index_text_type logger = logging.getLogger(__name__) @@ -94,7 +93,9 @@ def migrate_rows(engine_adapter, schema, **kwargs): # type: ignore migration_needed = True except Exception: # If we still can't eval it, leave it as-is - logger.warning("Exception trying to eval payload", exc_info=True) + logger.warning( + "Exception trying to eval payload", exc_info=True + ) new_snapshots.append( { diff --git a/sqlmesh/migrations/v0086_check_deterministic_bug.py b/sqlmesh/migrations/v0086_check_deterministic_bug.py index f44e5b8e33..81380c9d4b 100644 --- a/sqlmesh/migrations/v0086_check_deterministic_bug.py +++ b/sqlmesh/migrations/v0086_check_deterministic_bug.py @@ -5,7 +5,6 @@ from sqlmesh.core.console import get_console - logger = logging.getLogger(__name__) KEYS_TO_MAKE_DETERMINISTIC = ["__sqlmesh__vars__", "__sqlmesh__blueprint__vars__"] @@ -81,4 +80,6 @@ def migrate_rows(engine_adapter, schema, **kwargs): # type: ignore get_console().log_warning(warning) return except Exception: - logger.warning("Exception trying to eval payload", exc_info=True) + logger.warning( + "Exception trying to eval payload", exc_info=True + ) diff --git a/sqlmesh/migrations/v0087_normalize_blueprint_variables.py b/sqlmesh/migrations/v0087_normalize_blueprint_variables.py index fe737861c2..467373cac6 100644 --- a/sqlmesh/migrations/v0087_normalize_blueprint_variables.py +++ b/sqlmesh/migrations/v0087_normalize_blueprint_variables.py @@ -17,9 +17,9 @@ from dataclasses import dataclass from sqlglot import exp -from sqlmesh.core.console import get_console -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.core.console import get_console +from sqlmesh.utils.migration import blob_text_type, index_text_type logger = logging.getLogger(__name__) diff --git a/sqlmesh/migrations/v0088_warn_about_variable_python_env_diffs.py b/sqlmesh/migrations/v0088_warn_about_variable_python_env_diffs.py index 0aa7171821..d60482eb08 100644 --- a/sqlmesh/migrations/v0088_warn_about_variable_python_env_diffs.py +++ b/sqlmesh/migrations/v0088_warn_about_variable_python_env_diffs.py @@ -32,7 +32,12 @@ SQLMESH_VARS = "__sqlmesh__vars__" SQLMESH_BLUEPRINT_VARS = "__sqlmesh__blueprint__vars__" -METADATA_HASH_EXPRESSIONS = {"on_virtual_update", "audits", "signals", "audit_definitions"} +METADATA_HASH_EXPRESSIONS = { + "on_virtual_update", + "audits", + "signals", + "audit_definitions", +} def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0090_add_forward_only_column.py b/sqlmesh/migrations/v0090_add_forward_only_column.py index 48253691ec..cc1c54a134 100644 --- a/sqlmesh/migrations/v0090_add_forward_only_column.py +++ b/sqlmesh/migrations/v0090_add_forward_only_column.py @@ -4,7 +4,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0094_add_dev_version_and_fingerprint_columns.py b/sqlmesh/migrations/v0094_add_dev_version_and_fingerprint_columns.py index 9d7adf21a3..90e065c885 100644 --- a/sqlmesh/migrations/v0094_add_dev_version_and_fingerprint_columns.py +++ b/sqlmesh/migrations/v0094_add_dev_version_and_fingerprint_columns.py @@ -4,7 +4,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0098_add_dbt_node_info_in_node.py b/sqlmesh/migrations/v0098_add_dbt_node_info_in_node.py index b69ba8fa6f..76be9aaf2b 100644 --- a/sqlmesh/migrations/v0098_add_dbt_node_info_in_node.py +++ b/sqlmesh/migrations/v0098_add_dbt_node_info_in_node.py @@ -1,8 +1,10 @@ """Replace 'dbt_name' with 'dbt_node_info' in the snapshot definition""" import json + from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type + +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/migrations/v0102_normalize_python_env_payloads.py b/sqlmesh/migrations/v0102_normalize_python_env_payloads.py index 12f7da86b4..8329aed551 100644 --- a/sqlmesh/migrations/v0102_normalize_python_env_payloads.py +++ b/sqlmesh/migrations/v0102_normalize_python_env_payloads.py @@ -23,7 +23,7 @@ from sqlglot import exp -from sqlmesh.utils.migration import index_text_type, blob_text_type +from sqlmesh.utils.migration import blob_text_type, index_text_type def migrate_schemas(engine_adapter, schema, **kwargs): # type: ignore diff --git a/sqlmesh/utils/__init__.py b/sqlmesh/utils/__init__.py index da8d4f85a4..08786714f1 100644 --- a/sqlmesh/utils/__init__.py +++ b/sqlmesh/utils/__init__.py @@ -12,16 +12,16 @@ import traceback import types import typing as t +import unicodedata import uuid -from dataclasses import dataclass from collections import defaultdict from contextlib import contextmanager from copy import deepcopy -from enum import IntEnum, Enum +from dataclasses import dataclass +from enum import Enum, IntEnum from functools import lru_cache, reduce, wraps from pathlib import Path -import unicodedata from sqlglot import exp from sqlglot.dialects.dialect import Dialects @@ -54,10 +54,14 @@ def optional_import(name: str) -> t.Optional[types.ModuleType]: def major_minor(version: str) -> t.Tuple[int, int]: """Returns a tuple of just the major.minor for a version string (major.minor.patch).""" - return t.cast(t.Tuple[int, int], tuple(int(part) for part in version.split(".")[0:2])) + return t.cast( + t.Tuple[int, int], tuple(int(part) for part in version.split(".")[0:2]) + ) -def unique(iterable: t.Iterable[T], by: t.Callable[[T], t.Any] = lambda i: i) -> t.List[T]: +def unique( + iterable: t.Iterable[T], by: t.Callable[[T], t.Any] = lambda i: i +) -> t.List[T]: return list({by(i): None for i in iterable}) @@ -320,7 +324,8 @@ def columns_to_types_to_struct( return exp.DataType( this=exp.DataType.Type.STRUCT, expressions=[ - exp.ColumnDef(this=exp.to_identifier(k), kind=v) for k, v in columns_to_types.items() + exp.ColumnDef(this=exp.to_identifier(k), kind=v) + for k, v in columns_to_types.items() ], nested=True, ) @@ -386,7 +391,8 @@ def is_nothing_to_do(self) -> bool: def to_snake_case(name: str) -> str: return "".join( - f"_{c.lower()}" if c.isupper() and idx != 0 else c.lower() for idx, c in enumerate(name) + f"_{c.lower()}" if c.isupper() and idx != 0 else c.lower() + for idx, c in enumerate(name) ) diff --git a/sqlmesh/utils/cache.py b/sqlmesh/utils/cache.py index e1ff59a4a7..7255e52da0 100644 --- a/sqlmesh/utils/cache.py +++ b/sqlmesh/utils/cache.py @@ -74,7 +74,9 @@ def __init__(self, path: Path, prefix: t.Optional[str] = None): # File was deleted between glob() and stat() — skip stale cache entries gracefully continue - def get_or_load(self, name: str, entry_id: str = "", *, loader: t.Callable[[], T]) -> T: + def get_or_load( + self, name: str, entry_id: str = "", *, loader: t.Callable[[], T] + ) -> T: """Returns an existing cached entry or loads and caches a new one. Args: @@ -125,7 +127,9 @@ def put(self, name: str, entry_id: str = "", *, value: T) -> None: if not self._path.is_dir(): raise SQLMeshError(f"Cache path '{self._path}' is not a directory.") - with gzip.open(self._cache_entry_path(name, entry_id), "wb", compresslevel=1) as fd: + with gzip.open( + self._cache_entry_path(name, entry_id), "wb", compresslevel=1 + ) as fd: pickle.dump(value, fd) def exists(self, name: str, entry_id: str = "") -> bool: @@ -144,7 +148,9 @@ def clear(self) -> None: pass def _cache_entry_path(self, name: str, entry_id: str = "") -> Path: - entry_file_name = "__".join(p for p in (self._cache_version, name, entry_id) if p) + entry_file_name = "__".join( + p for p in (self._cache_version, name, entry_id) if p + ) full_path = self._path / sanitize_name(entry_file_name, include_unicode=True) if IS_WINDOWS: # handle paths longer than 260 chars diff --git a/sqlmesh/utils/concurrency.py b/sqlmesh/utils/concurrency.py index c5f78645f6..7ef7e12fa2 100644 --- a/sqlmesh/utils/concurrency.py +++ b/sqlmesh/utils/concurrency.py @@ -84,7 +84,9 @@ def _process_node(self, node: H, executor: Executor) -> None: self._node_errors.append(error) self._skip_next_nodes(node) - def _submit_next_nodes(self, executor: Executor, processed_node: t.Optional[H] = None) -> None: + def _submit_next_nodes( + self, executor: Executor, processed_node: t.Optional[H] = None + ) -> None: if not self._unprocessed_nodes_num: self._finished_future.set_result(None) return @@ -105,7 +107,9 @@ def _skip_next_nodes(self, parent: H) -> None: self._finished_future.set_result(None) return - skipped_nodes = {node for node, deps in self._unprocessed_nodes.items() if parent in deps} + skipped_nodes = { + node for node, deps in self._unprocessed_nodes.items() if parent in deps + } while skipped_nodes: self._skipped_nodes.extend(skipped_nodes) diff --git a/sqlmesh/utils/config.py b/sqlmesh/utils/config.py index 248f3adcc7..68712886ce 100644 --- a/sqlmesh/utils/config.py +++ b/sqlmesh/utils/config.py @@ -3,7 +3,6 @@ from sqlmesh.core.config.connection import ConnectionConfig from sqlmesh.utils import yaml - # Fields that should be excluded from the configuration hash excluded_fields: Set[str] = { "concurrent_tasks", diff --git a/sqlmesh/utils/connection_pool.py b/sqlmesh/utils/connection_pool.py index 9a70db6885..11b631b4cc 100644 --- a/sqlmesh/utils/connection_pool.py +++ b/sqlmesh/utils/connection_pool.py @@ -131,7 +131,9 @@ def __init__( self._connection_factory = connection_factory self._thread_cursors: t.Dict[t.Hashable, t.Any] = {} self._thread_transactions: t.Set[t.Hashable] = set() - self._thread_attributes: t.Dict[t.Hashable, t.Dict[str, t.Any]] = defaultdict(dict) + self._thread_attributes: t.Dict[t.Hashable, t.Dict[str, t.Any]] = defaultdict( + dict + ) self._thread_cursors_lock = Lock() self._thread_transactions_lock = Lock() self._cursor_init = cursor_init @@ -348,9 +350,7 @@ def create_connection_pool( pool_class = ( ThreadLocalSharedConnectionPool if multithreaded and shared_connection - else ThreadLocalConnectionPool - if multithreaded - else SingletonConnectionPool + else ThreadLocalConnectionPool if multithreaded else SingletonConnectionPool ) return pool_class(connection_factory, cursor_init=cursor_init) diff --git a/sqlmesh/utils/cron.py b/sqlmesh/utils/cron.py index 904202db7c..4ff50e681f 100644 --- a/sqlmesh/utils/cron.py +++ b/sqlmesh/utils/cron.py @@ -34,7 +34,12 @@ def interval_seconds(cron: str) -> int: class CroniterCache: - def __init__(self, cron: str, time: t.Optional[TimeLike] = None, tz: t.Optional[tzinfo] = None): + def __init__( + self, + cron: str, + time: t.Optional[TimeLike] = None, + tz: t.Optional[tzinfo] = None, + ): self.cron = cron self.tz = tz self.curr: datetime = to_datetime(now() if time is None else time, tz=self.tz) @@ -44,12 +49,16 @@ def get_next(self, estimate: bool = False) -> datetime: if estimate and self.interval_seconds: self.curr = self.curr + timedelta(seconds=self.interval_seconds) else: - self.curr = to_datetime(croniter(self.cron, self.curr).get_next() * 1000, tz=self.tz) + self.curr = to_datetime( + croniter(self.cron, self.curr).get_next() * 1000, tz=self.tz + ) return self.curr def get_prev(self, estimate: bool = False) -> datetime: if estimate and self.interval_seconds: self.curr = self.curr - timedelta(seconds=self.interval_seconds) else: - self.curr = to_datetime(croniter(self.cron, self.curr).get_prev() * 1000, tz=self.tz) + self.curr = to_datetime( + croniter(self.cron, self.curr).get_prev() * 1000, tz=self.tz + ) return self.curr diff --git a/sqlmesh/utils/dag.py b/sqlmesh/utils/dag.py index c39fd2a1d2..fa4f713507 100644 --- a/sqlmesh/utils/dag.py +++ b/sqlmesh/utils/dag.py @@ -99,7 +99,9 @@ def upstream(self, node: T) -> t.Set[T]: return self._upstream[node] - def _find_cycle_path(self, nodes_in_cycle: t.Dict[T, t.Set[T]]) -> t.Optional[t.List[T]]: + def _find_cycle_path( + self, nodes_in_cycle: t.Dict[T, t.Set[T]] + ) -> t.Optional[t.List[T]]: """Find the exact cycle path using DFS when a cycle is detected. Args: @@ -169,7 +171,9 @@ def sorted(self) -> t.List[T]: cycle_candidates: t.Collection = unprocessed_nodes while unprocessed_nodes: - next_nodes = {node for node, deps in unprocessed_nodes.items() if not deps} + next_nodes = { + node for node, deps in unprocessed_nodes.items() if not deps + } if not next_nodes: # A cycle was detected - find the exact cycle path @@ -189,8 +193,9 @@ def sorted(self) -> t.List[T]: ) cycle_msg = cycle_candidates_msg if last_processed_nodes: - last_processed_msg = "\nLast nodes added to the DAG: " + ", ".join( - str(node) for node in last_processed_nodes + last_processed_msg = ( + "\nLast nodes added to the DAG: " + + ", ".join(str(node) for node in last_processed_nodes) ) raise SQLMeshError( diff --git a/sqlmesh/utils/date.py b/sqlmesh/utils/date.py index f8df65c352..17d8873b3a 100644 --- a/sqlmesh/utils/date.py +++ b/sqlmesh/utils/date.py @@ -4,7 +4,6 @@ import time import typing as t import warnings - from datetime import date, datetime, timedelta, timezone, tzinfo import dateparser @@ -16,7 +15,6 @@ if t.TYPE_CHECKING: import pandas as pd - from sqlglot.dialects.dialect import DialectType UTC = timezone.utc @@ -179,10 +177,13 @@ def to_datetime( if epoch is None: relative_base = relative_base or now() expression = str(value) - if check_categorical_relative_expression and is_categorical_relative_expression( - expression + if ( + check_categorical_relative_expression + and is_categorical_relative_expression(expression) ): - relative_base = relative_base.replace(hour=0, minute=0, second=0, microsecond=0) + relative_base = relative_base.replace( + hour=0, minute=0, second=0, microsecond=0 + ) # note: we hardcode TIMEZONE: UTC to work around this bug: https://github.com/scrapinghub/dateparser/issues/896 # where dateparser just silently fails if it cant interpret the contents of /etc/localtime @@ -248,7 +249,10 @@ def date_dict( execution_dt = to_datetime(execution_time) prefixes = [ - ("latest", execution_dt), # TODO: Preserved for backward compatibility. Remove in 1.0.0. + ( + "latest", + execution_dt, + ), # TODO: Preserved for backward compatibility. Remove in 1.0.0. ("execution", execution_dt), ] @@ -283,7 +287,11 @@ def to_ds(obj: TimeLike, relative_base: t.Optional[datetime] = None) -> str: def to_ts(obj: TimeLike, relative_base: t.Optional[datetime] = None) -> str: """Converts a TimeLike object into YYYY-MM-DD HH:MM:SS formatted string.""" - return to_datetime(obj, relative_base=relative_base).replace(tzinfo=None).isoformat(sep=" ") + return ( + to_datetime(obj, relative_base=relative_base) + .replace(tzinfo=None) + .isoformat(sep=" ") + ) def to_tstz(obj: TimeLike, relative_base: t.Optional[datetime] = None) -> str: @@ -333,7 +341,9 @@ def make_inclusive( return (to_datetime(start), make_inclusive_end(end, dialect=dialect)) -def make_inclusive_end(end: TimeLike, dialect: t.Optional[DialectType] = "") -> datetime: +def make_inclusive_end( + end: TimeLike, dialect: t.Optional[DialectType] = "" +) -> datetime: import pandas as pd exclusive_end = make_exclusive(end) @@ -412,7 +422,10 @@ def to_time_column( ) -> exp.Expr: """Convert a TimeLike object to the same time format and type as the model's time column.""" if dialect == "clickhouse" and time_column_type.is_type( - *(exp.DataType.TEMPORAL_TYPES - {exp.DataType.Type.DATE, exp.DataType.Type.DATE32}) + *( + exp.DataType.TEMPORAL_TYPES + - {exp.DataType.Type.DATE, exp.DataType.Type.DATE32} + ) ): if time_column_type.is_type(exp.DataType.Type.DATETIME64): if nullable: diff --git a/sqlmesh/utils/errors.py b/sqlmesh/utils/errors.py index ca3e1bfb05..6f686e8238 100644 --- a/sqlmesh/utils/errors.py +++ b/sqlmesh/utils/errors.py @@ -27,7 +27,9 @@ class SQLMeshError(Exception): class ConfigError(SQLMeshError): location: t.Optional[Path] = None - def __init__(self, message: str | Exception, location: t.Optional[Path] = None) -> None: + def __init__( + self, message: str | Exception, location: t.Optional[Path] = None + ) -> None: super().__init__(message) if location: self.location = Path(location) if isinstance(location, str) else location @@ -253,12 +255,15 @@ def _format_schema_change_msg( dialect: SQL dialect for formatting error: Whether this is an error or warning """ - from sqlmesh.core.schema_diff import get_dropped_column_names, get_additive_column_names + from sqlmesh.core.schema_diff import (get_additive_column_names, + get_dropped_column_names) change_type = "destructive" if is_destructive else "additive" setting_name = "on_destructive_change" if is_destructive else "on_additive_change" action_verb = "drops" if is_destructive else "adds" - cli_flag = "--allow-destructive-model" if is_destructive else "--allow-additive-model" + cli_flag = ( + "--allow-destructive-model" if is_destructive else "--allow-additive-model" + ) column_names = ( get_dropped_column_names(alter_operations) @@ -278,9 +283,7 @@ def _format_schema_change_msg( ) # Main warning message - warning_msg = ( - f"Plan requires {change_type} change to forward-only model '{snapshot_name}'s schema" - ) + warning_msg = f"Plan requires {change_type} change to forward-only model '{snapshot_name}'s schema" if error: permissive_values = "`warn`, `allow`, or `ignore`" diff --git a/sqlmesh/utils/git.py b/sqlmesh/utils/git.py index cdb9d4e2d5..f03d0ad0a9 100644 --- a/sqlmesh/utils/git.py +++ b/sqlmesh/utils/git.py @@ -22,11 +22,16 @@ def list_uncommitted_changed_files(self) -> t.List[Path]: def list_committed_changed_files(self, target_branch: str = "main") -> t.List[Path]: return self._execute_list_output( - ["diff", "--name-only", "--diff-filter=d", f"{target_branch}..."], self._git_root + ["diff", "--name-only", "--diff-filter=d", f"{target_branch}..."], + self._git_root, ) - def _execute_list_output(self, commands: t.List[str], base_path: Path) -> t.List[Path]: - return [(base_path / o).absolute() for o in self._execute(commands).split("\n") if o] + def _execute_list_output( + self, commands: t.List[str], base_path: Path + ) -> t.List[Path]: + return [ + (base_path / o).absolute() for o in self._execute(commands).split("\n") if o + ] def _execute(self, commands: t.List[str]) -> str: result = subprocess.run( @@ -41,7 +46,11 @@ def _execute(self, commands: t.List[str]) -> str: if result.returncode != 0: stderr_output = result.stderr.decode("utf-8").strip() error_message = next( - (line for line in stderr_output.splitlines() if line.lower().startswith("fatal:")), + ( + line + for line in stderr_output.splitlines() + if line.lower().startswith("fatal:") + ), stderr_output, ) raise RuntimeError(f"Git error: {error_message}") diff --git a/sqlmesh/utils/jinja.py b/sqlmesh/utils/jinja.py index 829981db1c..82368887ce 100644 --- a/sqlmesh/utils/jinja.py +++ b/sqlmesh/utils/jinja.py @@ -11,7 +11,7 @@ from sys import exc_info from traceback import walk_tb -from jinja2 import Environment, Template, nodes, UndefinedError +from jinja2 import Environment, Template, UndefinedError, nodes from jinja2.runtime import Macro from sqlglot import Dialect, Parser, TokenType from sqlglot.expressions import Expression @@ -19,9 +19,9 @@ from sqlmesh.core import constants as c from sqlmesh.core import dialect as d from sqlmesh.utils import AttributeDict -from sqlmesh.utils.pydantic import PRIVATE_FIELDS, PydanticModel, field_serializer, field_validator from sqlmesh.utils.metaprogramming import SqlValue - +from sqlmesh.utils.pydantic import (PRIVATE_FIELDS, PydanticModel, + field_serializer, field_validator) if t.TYPE_CHECKING: CallNames = t.Tuple[t.Tuple[str, ...], t.Union[nodes.Call, nodes.Getattr]] @@ -133,7 +133,9 @@ def extract(self, jinja: str, dialect: str = "") -> t.Dict[str, MacroInfo]: macro_str = self._find_sql(macro_start, self._next) macros[name] = MacroInfo( definition=macro_str, - depends_on=list(extract_macro_references_and_variables(macro_str)[0]), + depends_on=list( + extract_macro_references_and_variables(macro_str)[0] + ), ) self._advance() @@ -144,7 +146,9 @@ def _advance(self, times: int = 1) -> None: super()._advance(times) self._tag = ( self._curr.text.upper() - if self._curr and self._prev and self._prev.token_type == TokenType.BLOCK_START + if self._curr + and self._prev + and self._prev.token_type == TokenType.BLOCK_START else "" ) @@ -165,7 +169,9 @@ def render_jinja(query: str, methods: t.Optional[t.Dict[str, t.Any]] = None) -> return ENVIRONMENT.from_string(query).render(methods or {}) -def find_call_names(node: nodes.Node, vars_in_scope: t.Set[str]) -> t.Iterator[CallNames]: +def find_call_names( + node: nodes.Node, vars_in_scope: t.Set[str] +) -> t.Iterator[CallNames]: vars_in_scope = vars_in_scope.copy() for child_node in node.iter_child_nodes(): if "target" in child_node.fields: @@ -186,7 +192,8 @@ def find_call_names(node: nodes.Node, vars_in_scope: t.Set[str]) -> t.Iterator[C for arg in child_node.args: vars_in_scope.add(arg.name) elif isinstance(child_node, nodes.Call) or ( - isinstance(child_node, nodes.Getattr) and not isinstance(child_node.node, nodes.Getattr) + isinstance(child_node, nodes.Getattr) + and not isinstance(child_node.node, nodes.Getattr) ): name = call_name(child_node) if name[0][0] != "'" and name[0] not in vars_in_scope: @@ -197,7 +204,8 @@ def find_call_names(node: nodes.Node, vars_in_scope: t.Set[str]) -> t.Iterator[C def extract_call_names( - jinja_str: str, cache: t.Optional[t.Dict[str, t.Tuple[t.List[CallNames], bool]]] = None + jinja_str: str, + cache: t.Optional[t.Dict[str, t.Tuple[t.List[CallNames], bool]]] = None, ) -> t.List[CallNames]: def parse() -> t.List[CallNames]: return list(find_call_names(ENVIRONMENT.parse(jinja_str), set())) @@ -246,7 +254,9 @@ def extract_macro_references_and_variables( elif len(call_name) == 1: macro_references.add(MacroReference(name=call_name[0])) elif len(call_name) == 2: - macro_references.add(MacroReference(package=call_name[0], name=call_name[1])) + macro_references.add( + MacroReference(package=call_name[0], name=call_name[1]) + ) return macro_references, variables @@ -324,7 +334,10 @@ def _serialize_attribute_dict( def _convert( val: t.Union[t.Dict[str, JinjaGlobalAttribute], t.Dict[str, t.Any]], ) -> t.Dict[str, t.Any]: - return {k: _convert(v) if isinstance(v, AttributeDict) else v for k, v in val.items()} + return { + k: _convert(v) if isinstance(v, AttributeDict) else v + for k, v in val.items() + } return _convert(value) @@ -332,7 +345,9 @@ def _convert( def trimmed(self) -> bool: return self._trimmed - def add_macros(self, macros: t.Dict[str, MacroInfo], package: t.Optional[str] = None) -> None: + def add_macros( + self, macros: t.Dict[str, MacroInfo], package: t.Optional[str] = None + ) -> None: """Adds macros to the target package. Args: @@ -359,7 +374,9 @@ def add_globals(self, globals: t.Dict[str, JinjaGlobalAttribute]) -> None: globals.pop("flat_graph", None) self.global_objs.update(**self._validate_global_objs(globals)) - def build_macro(self, reference: MacroReference, **kwargs: t.Any) -> t.Optional[t.Callable]: + def build_macro( + self, reference: MacroReference, **kwargs: t.Any + ) -> t.Optional[t.Callable]: """Builds a Python callable for a macro with the given reference. Args: @@ -386,7 +403,9 @@ def build_environment(self, **kwargs: t.Any) -> Environment: package_macros: t.Dict[str, t.Any] = defaultdict(AttributeDict) for package_name, macros in self.packages.items(): for macro_name, macro in macros.items(): - macro_wrapper = self._MacroWrapper(macro_name, package_name, self, context) + macro_wrapper = self._MacroWrapper( + macro_name, package_name, self, context + ) package_macros[package_name][macro_name] = macro_wrapper if macro.is_top_level and macro_name not in root_macros: root_macros[macro_name] = macro_wrapper @@ -483,7 +502,8 @@ def merge(self, other: JinjaMacroRegistry) -> JinjaMacroRegistry: packages=packages, root_macros=root_macros, global_objs=global_objs, - create_builtins_module=self.create_builtins_module or other.create_builtins_module, + create_builtins_module=self.create_builtins_module + or other.create_builtins_module, root_package_name=self.root_package_name or other.root_package_name, top_level_packages=[*self.top_level_packages, *other.top_level_packages], ) @@ -492,7 +512,9 @@ def to_expressions(self) -> t.List[Expression]: output: t.List[Expression] = [] filtered_objs = { - k: v for k, v in self.global_objs.items() if k in ("refs", "sources", "vars") + k: v + for k, v in self.global_objs.items() + if k in ("refs", "sources", "vars") } if filtered_objs: output.append( @@ -527,13 +549,17 @@ def data_hash_values(self) -> t.List[str]: data.append(macro.definition) trimmed_global_objs = { - k: self.global_objs[k] for k in ("refs", "sources", "vars") if k in self.global_objs + k: self.global_objs[k] + for k in ("refs", "sources", "vars") + if k in self.global_objs } data.append(json.dumps(trimmed_global_objs, sort_keys=True)) return data - def __deepcopy__(self, memo: t.Optional[t.Dict[int, t.Any]] = None) -> JinjaMacroRegistry: + def __deepcopy__( + self, memo: t.Optional[t.Dict[int, t.Any]] = None + ) -> JinjaMacroRegistry: return JinjaMacroRegistry.parse_obj(self.dict()) def _parse_macro(self, name: str, package: t.Optional[str]) -> Template: @@ -565,7 +591,9 @@ def _trim_macros( if visited is None: visited = defaultdict(set) - macros = self.packages.get(package, {}) if package is not None else self.root_macros + macros = ( + self.packages.get(package, {}) if package is not None else self.root_macros + ) trimmed_macros = {} dependencies: t.Dict[t.Optional[str], t.Set[str]] = defaultdict(set) @@ -598,9 +626,15 @@ def _macro_exists(self, name: str, package: t.Optional[str]) -> bool: ) def _get_macro(self, name: str, package: t.Optional[str]) -> MacroInfo: - return self.packages[package][name] if package is not None else self.root_macros[name] + return ( + self.packages[package][name] + if package is not None + else self.root_macros[name] + ) - def _to_non_private_macro_def(self, name: str, template: nodes.Template) -> nodes.Template: + def _to_non_private_macro_def( + self, name: str, template: nodes.Template + ) -> nodes.Template: for node in template.find_all((nodes.Macro, nodes.Call)): if isinstance(node, nodes.Macro): node.name = _non_private_name(name) @@ -609,7 +643,9 @@ def _to_non_private_macro_def(self, name: str, template: nodes.Template) -> node return template - def _create_builtin_globals(self, global_vars: t.Dict[str, t.Any]) -> t.Dict[str, t.Any]: + def _create_builtin_globals( + self, global_vars: t.Dict[str, t.Any] + ) -> t.Dict[str, t.Any]: """Creates Jinja builtin globals using a factory function defined in the provided module.""" engine_adapter = global_vars.pop("engine_adapter", None) global_vars = {**self.global_objs, **global_vars} @@ -687,7 +723,10 @@ def _var(var_name: str, default: t.Optional[t.Any] = None) -> t.Optional[t.Any]: def create_builtin_globals( - jinja_macros: JinjaMacroRegistry, global_vars: t.Dict[str, t.Any], *args: t.Any, **kwargs: t.Any + jinja_macros: JinjaMacroRegistry, + global_vars: t.Dict[str, t.Any], + *args: t.Any, + **kwargs: t.Any, ) -> t.Dict[str, t.Any]: global_vars.pop(c.GATEWAY, None) variables = global_vars.pop(c.SQLMESH_VARS, None) or {} @@ -701,7 +740,9 @@ def create_builtin_globals( def make_jinja_registry( - jinja_macros: JinjaMacroRegistry, package_name: str, jinja_references: t.Set[MacroReference] + jinja_macros: JinjaMacroRegistry, + package_name: str, + jinja_references: t.Set[MacroReference], ) -> JinjaMacroRegistry: """ Creates a Jinja macro registry for a specific package. diff --git a/sqlmesh/utils/lineage.py b/sqlmesh/utils/lineage.py index f63395708d..b0ded45029 100644 --- a/sqlmesh/utils/lineage.py +++ b/sqlmesh/utils/lineage.py @@ -2,24 +2,20 @@ from pathlib import Path from pydantic import Field - -from sqlmesh.core.dialect import normalize_model_name -from sqlmesh.core.linter.helpers import ( - TokenPositionDetails, -) -from sqlmesh.core.linter.rule import Range, Position -from sqlmesh.core.model.definition import SqlModel, ExternalModel, PythonModel, SeedModel +from ruamel.yaml import YAML from sqlglot import exp -from sqlglot.optimizer.scope import build_scope - from sqlglot.optimizer.normalize_identifiers import normalize_identifiers -from ruamel.yaml import YAML +from sqlglot.optimizer.scope import build_scope +from sqlmesh.core.dialect import normalize_model_name +from sqlmesh.core.linter.helpers import TokenPositionDetails +from sqlmesh.core.linter.rule import Position, Range +from sqlmesh.core.model.definition import (ExternalModel, PythonModel, + SeedModel, SqlModel) from sqlmesh.utils.pydantic import PydanticModel if t.TYPE_CHECKING: - from sqlmesh.core.context import Context - from sqlmesh.core.context import GenericContext + from sqlmesh.core.context import Context, GenericContext class ModelReference(PydanticModel): @@ -163,7 +159,9 @@ def extract_references_from_query( # If there's a catalog or database qualifier, adjust the start position catalog_or_db = table.args.get("catalog") or table.args.get("db") if catalog_or_db is not None: - catalog_or_db_meta = TokenPositionDetails.from_meta(catalog_or_db.meta) + catalog_or_db_meta = TokenPositionDetails.from_meta( + catalog_or_db.meta + ) catalog_or_db_range_sqlmesh = catalog_or_db_meta.to_range(read_file) start_pos_sqlmesh = catalog_or_db_range_sqlmesh.start @@ -371,7 +369,9 @@ def _get_column_table_range(column: exp.Column, read_file: t.List[str]) -> Range table_parts = column.parts[:-1] - start_range = TokenPositionDetails.from_meta(table_parts[0].meta).to_range(read_file) + start_range = TokenPositionDetails.from_meta(table_parts[0].meta).to_range( + read_file + ) end_range = TokenPositionDetails.from_meta(table_parts[-1].meta).to_range(read_file) return Range( @@ -419,7 +419,9 @@ def get_yaml_model_name_ranges(path: Path) -> t.Optional[t.Dict[str, Range]]: if isinstance(item, dict): position_data = item.lc.data["name"] # type: ignore start = Position(line=position_data[2], character=position_data[3]) - end = Position(line=position_data[2], character=position_data[3] + len(item["name"])) + end = Position( + line=position_data[2], character=position_data[3] + len(item["name"]) + ) name = item.get("name") if not name: continue diff --git a/sqlmesh/utils/metaprogramming.py b/sqlmesh/utils/metaprogramming.py index a5bd376566..cc25ae2632 100644 --- a/sqlmesh/utils/metaprogramming.py +++ b/sqlmesh/utils/metaprogramming.py @@ -30,7 +30,9 @@ LITERALS = (Number, str, bytes, tuple, list, dict, set, bool) -def _is_relative_to(path: t.Optional[Path | str], other: t.Optional[Path | str]) -> bool: +def _is_relative_to( + path: t.Optional[Path | str], other: t.Optional[Path | str] +) -> bool: if path is None or other is None: return False @@ -89,18 +91,27 @@ def func_globals(func: t.Callable) -> t.Dict[str, t.Any]: if hasattr(func, "__code__"): root_node = parse_source(func) - func_args = next(node for node in ast.walk(root_node) if isinstance(node, ast.arguments)) - arg_defaults = (d for d in func_args.defaults + func_args.kw_defaults if d is not None) + func_args = next( + node for node in ast.walk(root_node) if isinstance(node, ast.arguments) + ) + arg_defaults = ( + d for d in func_args.defaults + func_args.kw_defaults if d is not None + ) # ast.Name corresponds to variable references, such as foo or x.foo. The former is # represented as Name(id=foo), and the latter as Attribute(value=Name(id=x) attr=foo) arg_globals = [ - n.id for default in arg_defaults for n in ast.walk(default) if isinstance(n, ast.Name) + n.id + for default in arg_defaults + for n in ast.walk(default) + if isinstance(n, ast.Name) ] code = func.__code__ for var in ( - arg_globals + list(_code_globals(code)) + decorator_vars(func, root_node=root_node) + arg_globals + + list(_code_globals(code)) + + decorator_vars(func, root_node=root_node) ): if var in func.__globals__: variables[var] = func.__globals__[var] @@ -209,7 +220,9 @@ def join_source(lnum: int) -> str: obj = obj.__code__ if hasattr(obj, "co_firstlineno"): lnum = obj.co_firstlineno - 1 - pat = re.compile(r"^(\s*def\s)|(\s*async\s+def\s)|(.*(? 0: try: line = lines[lnum] @@ -234,7 +247,9 @@ def _decorator_name(decorator: ast.expr) -> str: return node.id if isinstance(node, ast.Name) else "" -def decorator_vars(func: t.Callable, root_node: t.Optional[ast.Module] = None) -> t.List[str]: +def decorator_vars( + func: t.Callable, root_node: t.Optional[ast.Module] = None +) -> t.List[str]: """ Returns a list of all the decorators of a callable, as well as names of objects that are referenced in their argument list. These objects may be transitive dependencies @@ -437,7 +452,10 @@ def is_value(self) -> bool: @classmethod def value( - cls, v: t.Any, is_metadata: t.Optional[bool] = None, sort_root_dict: bool = False + cls, + v: t.Any, + is_metadata: t.Optional[bool] = None, + sort_root_dict: bool = False, ) -> Executable: payload = _dict_sort(v) if sort_root_dict else repr(v) return Executable( @@ -530,7 +548,9 @@ def serialize_env(env: t.Dict[str, t.Any], path: Path) -> t.Dict[str, Executable # # [1]: https://github.com/jd/tenacity/blob/0d40e76f7d06d631fb127e1ec58c8bd776e70d49/tenacity/__init__.py#L322-L346 # [2]: https://github.com/python/cpython/blob/f502c8f6a6db4be27c97a0e5466383d117859b7f/Lib/functools.py#L33-L57 - if not relative_obj_file_path and (wrapped := getattr(v, "__wrapped__", None)): + if not relative_obj_file_path and ( + wrapped := getattr(v, "__wrapped__", None) + ): v = wrapped file_path = Path(inspect.getfile(wrapped)) relative_obj_file_path = _is_relative_to(file_path, path) @@ -544,7 +564,9 @@ def serialize_env(env: t.Dict[str, t.Any], path: Path) -> t.Dict[str, Executable payload=normalize_source(v), kind=ExecutableKind.DEFINITION, # Do `as_posix` to serialize windows path back to POSIX - path=t.cast(Path, file_path).relative_to(path.absolute()).as_posix(), + path=t.cast(Path, file_path) + .relative_to(path.absolute()) + .as_posix(), alias=k if name != k else None, is_metadata=is_metadata, ) @@ -643,9 +665,7 @@ def format_evaluated_code_exception( executable = python_env[func] indent = error_line[: eval_code_match.start()] - error_line = ( - f"{indent}File '{executable.path}' (or imported file), line {line_num}, in {func}" - ) + error_line = f"{indent}File '{executable.path}' (or imported file), line {line_num}, in {func}" code = executable.payload formatted = [] diff --git a/sqlmesh/utils/pandas.py b/sqlmesh/utils/pandas.py index 43851e861a..08fa461d3b 100644 --- a/sqlmesh/utils/pandas.py +++ b/sqlmesh/utils/pandas.py @@ -11,8 +11,8 @@ @lru_cache() def get_pandas_type_mappings() -> t.Dict[t.Any, exp.DataType]: - import pandas as pd import numpy as np + import pandas as pd mappings = { np.dtype("int8"): exp.DataType.build("tinyint"), @@ -59,7 +59,9 @@ def columns_to_types_from_dtypes( result = {} for column_name, column_type in dtypes: exp_type: t.Optional[exp.DataType] = None - if hasattr(pd, "DatetimeTZDtype") and isinstance(column_type, pd.DatetimeTZDtype): + if hasattr(pd, "DatetimeTZDtype") and isinstance( + column_type, pd.DatetimeTZDtype + ): exp_type = exp.DataType.build("timestamptz") else: exp_type = get_pandas_type_mappings().get(column_type) diff --git a/sqlmesh/utils/process.py b/sqlmesh/utils/process.py index 453fee78f5..1cb6644557 100644 --- a/sqlmesh/utils/process.py +++ b/sqlmesh/utils/process.py @@ -1,8 +1,9 @@ # mypy: disable-error-code=no-untyped-def -from concurrent.futures import Future, ProcessPoolExecutor -import typing as t import multiprocessing as mp +import typing as t +from concurrent.futures import Future, ProcessPoolExecutor + from sqlmesh.utils.windows import IS_WINDOWS @@ -13,7 +14,9 @@ class SynchronousPoolExecutor: with forking in test environments or when forking isn't possible (non-posix). """ - def __init__(self, max_workers=None, mp_context=None, initializer=None, initargs=()): + def __init__( + self, max_workers=None, mp_context=None, initializer=None, initargs=() + ): if initializer is not None: try: initializer(*initargs) diff --git a/sqlmesh/utils/pydantic.py b/sqlmesh/utils/pydantic.py index 4e5cfc3dc1..41da5d84b6 100644 --- a/sqlmesh/utils/pydantic.py +++ b/sqlmesh/utils/pydantic.py @@ -24,9 +24,9 @@ T = t.TypeVar("T") DEFAULT_ARGS = {"exclude_none": True, "by_alias": True} PRIVATE_FIELDS = "__pydantic_private__" -PYDANTIC_MAJOR_VERSION, PYDANTIC_MINOR_VERSION = [int(p) for p in pydantic.__version__.split(".")][ - :2 -] +PYDANTIC_MAJOR_VERSION, PYDANTIC_MINOR_VERSION = [ + int(p) for p in pydantic.__version__.split(".") +][:2] def field_validator(*args: t.Any, **kwargs: t.Any) -> t.Callable[[t.Any], t.Any]: @@ -124,7 +124,9 @@ def parse_obj(cls: t.Type["Model"], obj: t.Any) -> "Model": return super().model_validate(obj) @classmethod - def parse_raw(cls: t.Type["Model"], b: t.Union[str, bytes], **kwargs: t.Any) -> "Model": + def parse_raw( + cls: t.Type["Model"], b: t.Union[str, bytes], **kwargs: t.Any + ) -> "Model": return super().model_validate_json(b, **kwargs) @classmethod @@ -134,7 +136,9 @@ def missing_required_fields( return cls.required_fields() - provided_fields @classmethod - def extra_fields(cls: t.Type["PydanticModel"], provided_fields: t.Set[str]) -> t.Set[str]: + def extra_fields( + cls: t.Type["PydanticModel"], provided_fields: t.Set[str] + ) -> t.Set[str]: return provided_fields - cls.all_fields() @classmethod @@ -169,13 +173,18 @@ def __eq__(self, other: t.Any) -> bool: def __hash__(self) -> int: if (PYDANTIC_MAJOR_VERSION, PYDANTIC_MINOR_VERSION) < (2, 6): - obj = {k: v for k, v in self.__dict__.items() if k in self.all_field_infos()} + obj = { + k: v for k, v in self.__dict__.items() if k in self.all_field_infos() + } return hash(self.__class__) + hash(tuple(obj.values())) - from pydantic._internal._model_construction import make_hash_func # type: ignore + from pydantic._internal._model_construction import \ + make_hash_func # type: ignore if self.__class__ not in PydanticModel._hash_func_mapping: - PydanticModel._hash_func_mapping[self.__class__] = make_hash_func(self.__class__) + PydanticModel._hash_func_mapping[self.__class__] = make_hash_func( + self.__class__ + ) return PydanticModel._hash_func_mapping[self.__class__](self) @@ -258,7 +267,9 @@ def _get_field( else: expression = parse_one(v, dialect=dialect) - expression = exp.column(expression) if isinstance(expression, exp.Identifier) else expression + expression = ( + exp.column(expression) if isinstance(expression, exp.Identifier) else expression + ) expression = quote_identifiers( normalize_identifiers(expression, dialect=dialect), dialect=dialect ) @@ -358,13 +369,18 @@ def get_concrete_types_from_typehint(typehint: type[t.Any]) -> set[type[t.Any]]: else: from pydantic.functional_validators import BeforeValidator - SQLGlotListOfStrings = t.Annotated[t.List[str], BeforeValidator(validate_list_of_strings)] + SQLGlotListOfStrings = t.Annotated[ + t.List[str], BeforeValidator(validate_list_of_strings) + ] SQLGlotString = t.Annotated[str, BeforeValidator(validate_string)] SQLGlotBool = t.Annotated[bool, BeforeValidator(bool_validator)] SQLGlotPositiveInt = t.Annotated[int, BeforeValidator(positive_int_validator)] SQLGlotColumn = t.Annotated[exp.Expr, BeforeValidator(column_validator)] - SQLGlotListOfFields = t.Annotated[t.List[exp.Expr], BeforeValidator(list_of_fields_validator)] + SQLGlotListOfFields = t.Annotated[ + t.List[exp.Expr], BeforeValidator(list_of_fields_validator) + ] SQLGlotListOfFieldsOrStar = t.Annotated[ - t.Union[SQLGlotListOfFields, exp.Star], BeforeValidator(list_of_fields_or_star_validator) + t.Union[SQLGlotListOfFields, exp.Star], + BeforeValidator(list_of_fields_or_star_validator), ] SQLGlotCron = t.Annotated[str, BeforeValidator(cron_validator)] diff --git a/sqlmesh/utils/rich.py b/sqlmesh/utils/rich.py index 0b43e3d87c..5437805cae 100644 --- a/sqlmesh/utils/rich.py +++ b/sqlmesh/utils/rich.py @@ -1,14 +1,13 @@ from __future__ import annotations -import typing as t - import re +import typing as t +from rich.align import Align from rich.console import Console from rich.progress import Column, ProgressColumn, Task, Text -from rich.theme import Theme from rich.table import Table -from rich.align import Align +from rich.theme import Theme if t.TYPE_CHECKING: import pandas as pd @@ -78,13 +77,17 @@ def df_to_table( Returns: Table: The rich Table instance passed, populated with the DataFrame values.""" - rich_table = Table(title=f"[bold red]{header}[/bold red]", show_lines=True, min_width=60) + rich_table = Table( + title=f"[bold red]{header}[/bold red]", show_lines=True, min_width=60 + ) if show_index: index_name = str(index_name) if index_name else "" rich_table.add_column(Align.center(index_name)) for column in df.columns: - column_name = column if isinstance(column, str) else ": ".join(str(col) for col in column) + column_name = ( + column if isinstance(column, str) else ": ".join(str(col) for col in column) + ) # Color coding unit test columns (expected/actual), can be removed or refactored if df_to_table is used elswhere too lower = column_name.lower() diff --git a/sqlmesh_dbt/cli.py b/sqlmesh_dbt/cli.py index 278daa5370..529ee523bf 100644 --- a/sqlmesh_dbt/cli.py +++ b/sqlmesh_dbt/cli.py @@ -1,15 +1,19 @@ -import typing as t +import functools import sys +import typing as t +from pathlib import Path + import click + +from sqlmesh_dbt.error import ErrorHandlingGroup, cli_global_error_handler from sqlmesh_dbt.operations import DbtOperations, create -from sqlmesh_dbt.error import cli_global_error_handler, ErrorHandlingGroup -from pathlib import Path from sqlmesh_dbt.options import YamlParamType -import functools def _get_dbt_operations( - ctx: click.Context, vars: t.Optional[t.Dict[str, t.Any]], threads: t.Optional[int] = None + ctx: click.Context, + vars: t.Optional[t.Dict[str, t.Any]], + threads: t.Optional[int] = None, ) -> DbtOperations: if not isinstance(ctx.obj, functools.partial): raise ValueError(f"Unexpected click context object: {type(ctx.obj)}") @@ -46,7 +50,9 @@ def _cleanup() -> None: multiple=True, help="Specify the model nodes to include; other nodes are excluded.", ) -exclude_option = click.option("--exclude", multiple=True, help="Specify the nodes to exclude.") +exclude_option = click.option( + "--exclude", multiple=True, help="Specify the nodes to exclude." +) # TODO: expand this out into --resource-type/--resource-types and --exclude-resource-type/--exclude-resource-types resource_types = [ @@ -70,7 +76,9 @@ def _cleanup() -> None: @click.group(cls=ErrorHandlingGroup, invoke_without_command=True) -@click.option("--profile", help="Which existing profile to load. Overrides output.profile") +@click.option( + "--profile", help="Which existing profile to load. Overrides output.profile" +) @click.option("-t", "--target", help="Which target to load for the given profile") @click.option( "-d", @@ -153,7 +161,9 @@ def dbt( help="Run against a specific Virtual Data Environment (VDE) instead of the main environment", ) @click.option( - "--empty/--no-empty", default=False, help="If specified, limit input refs and sources" + "--empty/--no-empty", + default=False, + help="If specified, limit input refs and sources", ) @click.option( "--threads", @@ -180,7 +190,9 @@ def run( @resource_type_option @vars_option @click.pass_context -def list_(ctx: click.Context, vars: t.Optional[t.Dict[str, t.Any]], **kwargs: t.Any) -> None: +def list_( + ctx: click.Context, vars: t.Optional[t.Dict[str, t.Any]], **kwargs: t.Any +) -> None: """List the resources in your project""" _get_dbt_operations(ctx, vars).list_(**kwargs) diff --git a/sqlmesh_dbt/console.py b/sqlmesh_dbt/console.py index 6bf7a1618f..820df5f094 100644 --- a/sqlmesh_dbt/console.py +++ b/sqlmesh_dbt/console.py @@ -1,8 +1,10 @@ import typing as t + +from rich.tree import Tree + from sqlmesh.core.console import TerminalConsole from sqlmesh.core.model import Model from sqlmesh.core.snapshot.definition import Node -from rich.tree import Tree class DbtCliConsole(TerminalConsole): diff --git a/sqlmesh_dbt/error.py b/sqlmesh_dbt/error.py index 49a2f8195b..8349ad7944 100644 --- a/sqlmesh_dbt/error.py +++ b/sqlmesh_dbt/error.py @@ -1,8 +1,9 @@ -import typing as t import logging +import sys +import typing as t from functools import wraps + import click -import sys logger = logging.getLogger(__name__) @@ -17,9 +18,10 @@ def wrapper(*args: t.List[t.Any], **kwargs: t.Any) -> t.Any: except Exception as ex: # these imports are deliberately deferred to avoid the penalty of importing the `sqlmesh` # package up front for every CLI command - from sqlmesh.utils.errors import SQLMeshError from sqlglot.errors import SqlglotError + from sqlmesh.utils.errors import SQLMeshError + if isinstance(ex, (SQLMeshError, SqlglotError, ValueError)): click.echo(click.style("Error: " + str(ex), fg="red")) sys.exit(1) diff --git a/sqlmesh_dbt/operations.py b/sqlmesh_dbt/operations.py index 576d8e090b..6844aa654a 100644 --- a/sqlmesh_dbt/operations.py +++ b/sqlmesh_dbt/operations.py @@ -1,23 +1,28 @@ from __future__ import annotations + +import logging import typing as t -from rich.progress import Progress from pathlib import Path -import logging + +from rich.progress import Progress + from sqlmesh_dbt import selectors if t.TYPE_CHECKING: # important to gate these to be able to defer importing sqlmesh until we need to from sqlmesh.core.context import Context - from sqlmesh.dbt.project import Project - from sqlmesh_dbt.console import DbtCliConsole from sqlmesh.core.model import Model from sqlmesh.core.plan import Plan, PlanBuilder + from sqlmesh.dbt.project import Project + from sqlmesh_dbt.console import DbtCliConsole logger = logging.getLogger(__name__) class DbtOperations: - def __init__(self, sqlmesh_context: Context, dbt_project: Project, debug: bool = False): + def __init__( + self, sqlmesh_context: Context, dbt_project: Project, debug: bool = False + ): self.context = sqlmesh_context self.project = dbt_project self.debug = debug @@ -103,7 +108,9 @@ def _selected_models( resource_type: t.Optional[str] = None, ) -> t.Dict[str, Model]: if sqlmesh_selector := selectors.to_sqlmesh( - *selectors.consolidate(select or [], exclude or [], models or [], resource_type) + *selectors.consolidate( + select or [], exclude or [], models or [], resource_type + ) ): if self.debug: self.console.print(f"dbt --select: {select}") @@ -195,15 +202,19 @@ def _plan_builder_options( # --full-refresh is implemented in terms of "add every model as a restatement" # however, `--empty` sets skip_backfill=True, which causes the BackfillStage of the plan to be skipped. # the re-processing of data intervals happens in the BackfillStage, so if it gets skipped, restatements become a no-op - raise ValueError("`--full-refresh` alongside `--empty` is not currently supported.") + raise ValueError( + "`--full-refresh` alongside `--empty` is not currently supported." + ) if full_refresh: options.update( dict( # Add every selected model as a restatement to force them to get repopulated from scratch - restate_models=[m.dbt_fqn for m in self.context.models.values() if m.dbt_fqn] - if not select_models - else select_models, + restate_models=( + [m.dbt_fqn for m in self.context.models.values() if m.dbt_fqn] + if not select_models + else select_models + ), # by default in SQLMesh, restatements only operate on what has been committed to state. # in order to emulate dbt, we need to use the local filesystem instead, so we override this default always_include_local_changes=True, @@ -245,12 +256,12 @@ def create( load_task_id = progress.add_task("Loading engine", total=None) from sqlmesh import configure_logging + from sqlmesh.core.console import set_console from sqlmesh.core.context import Context + from sqlmesh.core.selector import DbtSelector from sqlmesh.dbt.loader import DbtLoader - from sqlmesh.core.console import set_console - from sqlmesh_dbt.console import DbtCliConsole from sqlmesh.utils.errors import SQLMeshError - from sqlmesh.core.selector import DbtSelector + from sqlmesh_dbt.console import DbtCliConsole # clear any existing handlers set up by click/rich as defaults so that once SQLMesh logging config is applied, # we dont get duplicate messages logged from things like console.log_warning() @@ -300,12 +311,17 @@ def init_project_if_required(project_dir: Path, start: t.Optional[str] = None) - This is preferable to trying to inject config into `dbt_project.yml` because it means we have full control over the file and dont need to worry about accidentally reformatting it or accidentally clobbering other config """ - from sqlmesh.cli.project_init import init_example_project, ProjectTemplate + from sqlmesh.cli.project_init import ProjectTemplate, init_example_project from sqlmesh.core.config.common import ALL_CONFIG_FILENAMES from sqlmesh.core.console import get_console - if not any(f.exists() for f in [project_dir / file for file in ALL_CONFIG_FILENAMES]): + if not any( + f.exists() for f in [project_dir / file for file in ALL_CONFIG_FILENAMES] + ): get_console().log_warning("No existing SQLMesh config detected; creating one") init_example_project( - path=project_dir, engine_type=None, template=ProjectTemplate.DBT, start=start + path=project_dir, + engine_type=None, + template=ProjectTemplate.DBT, + start=start, ) diff --git a/sqlmesh_dbt/options.py b/sqlmesh_dbt/options.py index 5a7cabe93b..d51a273700 100644 --- a/sqlmesh_dbt/options.py +++ b/sqlmesh_dbt/options.py @@ -1,4 +1,5 @@ import typing as t + import click from click.core import Context, Parameter @@ -20,6 +21,10 @@ def convert( self.fail(f"String '{value}' is not valid YAML", param, ctx) if not isinstance(parsed, dict): - self.fail(f"String '{value}' did not evaluate to a dict, got: {parsed}", param, ctx) + self.fail( + f"String '{value}' did not evaluate to a dict, got: {parsed}", + param, + ctx, + ) return parsed diff --git a/sqlmesh_dbt/selectors.py b/sqlmesh_dbt/selectors.py index 5821586ad3..dddde7b8cf 100644 --- a/sqlmesh_dbt/selectors.py +++ b/sqlmesh_dbt/selectors.py @@ -1,5 +1,5 @@ -import typing as t import logging +import typing as t logger = logging.getLogger(__name__) @@ -21,7 +21,9 @@ def consolidate( raise ValueError('"models" and "select" are mutually exclusive arguments') if models and resource_type: - raise ValueError('"models" and "resource_type" are mutually exclusive arguments') + raise ValueError( + '"models" and "resource_type" are mutually exclusive arguments' + ) if models: # --models implies resource_type:model @@ -84,7 +86,9 @@ def to_sqlmesh(dbt_select: t.List[str], dbt_exclude: t.List[str]) -> t.Optional[ return None select_expr = " | ".join(_to_sqlmesh(expr) for expr in dbt_select) - select_expr = _wrap(select_expr) if dbt_exclude and len(dbt_select) > 1 else select_expr + select_expr = ( + _wrap(select_expr) if dbt_exclude and len(dbt_select) > 1 else select_expr + ) exclude_expr = "" @@ -118,7 +122,9 @@ def _to_sqlmesh(selector_str: str) -> str: return " | ".join([expr for expr in [union_expr, intersection_expr] if expr]) -def _split_unions_and_intersections(selector_str: str) -> t.Tuple[t.List[str], t.List[str]]: +def _split_unions_and_intersections( + selector_str: str, +) -> t.Tuple[t.List[str], t.List[str]]: # break space-separated items like: "my_first_model my_second_model" into a list of selectors to union # and comma-separated items like: "my_first_model,my_second_model" into a list of selectors to intersect # but, take into account brackets, eg "(my_first_model & my_second_model)" should not be split diff --git a/tests/cli/test_cli.py b/tests/cli/test_cli.py index c625cb084d..8808efe7f1 100644 --- a/tests/cli/test_cli.py +++ b/tests/cli/test_cli.py @@ -1,22 +1,24 @@ import json import os -import pytest import string -import time_machine from os import getcwd, path, remove from pathlib import Path from shutil import rmtree from unittest.mock import MagicMock +import pytest +import time_machine from click import ClickException from click.testing import CliRunner + from sqlmesh import RuntimeEnv -from sqlmesh.cli.project_init import ProjectTemplate, init_example_project from sqlmesh.cli.main import cli +from sqlmesh.cli.project_init import ProjectTemplate, init_example_project +from sqlmesh.core.config.connection import DIALECT_TO_TYPE from sqlmesh.core.context import Context from sqlmesh.integrations.dlt import generate_dlt_models -from sqlmesh.utils.date import now_ds, time_like_to_str, timedelta, to_datetime, yesterday_ds -from sqlmesh.core.config.connection import DIALECT_TO_TYPE +from sqlmesh.utils.date import (now_ds, time_like_to_str, timedelta, + to_datetime, yesterday_ds) FREEZE_TIME = "2023-01-01 00:00:00 UTC" @@ -25,7 +27,9 @@ @pytest.fixture(autouse=True) def mock_runtime_env(monkeypatch): - monkeypatch.setattr("sqlmesh.RuntimeEnv.get", MagicMock(return_value=RuntimeEnv.TERMINAL)) + monkeypatch.setattr( + "sqlmesh.RuntimeEnv.get", MagicMock(return_value=RuntimeEnv.TERMINAL) + ) @pytest.fixture(scope="session") @@ -41,8 +45,7 @@ def create_example_project(temp_dir, template=ProjectTemplate.DEFAULT) -> None: """ init_example_project(temp_dir, engine_type="duckdb", template=template) with open(temp_dir / "config.yaml", "w", encoding="utf-8") as f: - f.write( - f"""gateways: + f.write(f"""gateways: local: connection: type: duckdb @@ -55,14 +58,14 @@ def create_example_project(temp_dir, template=ProjectTemplate.DEFAULT) -> None: plan: no_prompts: false -""" - ) +""") def update_incremental_model(temp_dir) -> None: - with open(temp_dir / "models" / "incremental_model.sql", "w", encoding="utf-8") as f: - f.write( - """ + with open( + temp_dir / "models" / "incremental_model.sql", "w", encoding="utf-8" + ) as f: + f.write(""" MODEL ( name sqlmesh_example.incremental_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -82,14 +85,12 @@ def update_incremental_model(temp_dir) -> None: sqlmesh_example.seed_model WHERE event_date between @start_date and @end_date -""" - ) +""") def update_full_model(temp_dir) -> None: with open(temp_dir / "models" / "full_model.sql", "w", encoding="utf-8") as f: - f.write( - """ + f.write(""" MODEL ( name sqlmesh_example.full_model, kind FULL, @@ -104,8 +105,7 @@ def update_full_model(temp_dir) -> None: FROM sqlmesh_example.incremental_model GROUP BY item_id -""" - ) +""") def init_prod_and_backfill(runner, temp_dir) -> None: @@ -158,7 +158,9 @@ def test_version(runner, tmp_path): def test_plan_no_config(runner, tmp_path): # Error if no SQLMesh project config is found - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan"] + ) assert result.exit_code == 1 assert "Error: SQLMesh project config could not be found" in result.output @@ -175,8 +177,13 @@ def test_plan(runner, tmp_path): ) assert_plan_success(result) # 'Models needing backfill' section and eval progress bar should display the same inclusive intervals - assert "sqlmesh_example.incremental_model: [2020-01-01 - 2022-12-31]" in result.output - assert "sqlmesh_example.incremental_model [insert 2020-01-01 - 2022-12-31]" in result.output + assert ( + "sqlmesh_example.incremental_model: [2020-01-01 - 2022-12-31]" in result.output + ) + assert ( + "sqlmesh_example.incremental_model [insert 2020-01-01 - 2022-12-31]" + in result.output + ) def test_plan_skip_tests(runner, tmp_path): @@ -185,7 +192,9 @@ def test_plan_skip_tests(runner, tmp_path): # Successful test run message should not appear with `--skip-tests` # Input: `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--skip-tests"], input="y\n" + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--skip-tests"], + input="y\n", ) assert result.exit_code == 0 assert "Successfully Ran 1 tests against duckdb" not in result.output @@ -197,16 +206,16 @@ def test_plan_skip_linter(runner, tmp_path): create_example_project(tmp_path) with open(tmp_path / "config.yaml", "a", encoding="utf-8") as f: - f.write( - """linter: + f.write("""linter: enabled: True rules: "ALL" - """ - ) + """) # Input: `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--skip-linter"], input="y\n" + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--skip-linter"], + input="y\n", ) assert result.exit_code == 0 @@ -247,9 +256,14 @@ def test_plan_skip_backfill(runner, tmp_path, flag): create_example_project(tmp_path) # plan for `prod` errors if `--skip-backfill` is passed without --no-gaps - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", flag]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", flag] + ) assert result.exit_code == 1 - assert "Skipping the backfill stage for production can lead to unexpected" in result.output + assert ( + "Skipping the backfill stage for production can lead to unexpected" + in result.output + ) # plan executes virtual update without executing model batches # Input: `y` to perform virtual update @@ -269,7 +283,15 @@ def test_plan_min_intervals(runner, tmp_path): # build prod so the dev plan below has a baseline to diff against runner.invoke( cli, - ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--no-prompts", "--auto-apply"], + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "--no-prompts", + "--auto-apply", + ], ) update_incremental_model(tmp_path) @@ -295,7 +317,16 @@ def test_plan_min_intervals(runner, tmp_path): # a non-integer value is rejected by click, not surfaced as a traceback result = runner.invoke( cli, - ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev", "--min-intervals", "abc"], + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "dev", + "--min-intervals", + "abc", + ], ) assert result.exit_code == 2 assert "is not a valid integer" in result.output @@ -320,7 +351,9 @@ def test_plan_verbose(runner, tmp_path): # Input: `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--verbose"], input="y\n" + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--verbose"], + input="y\n", ) assert_plan_success(result) assert "sqlmesh_example.seed_model created" in result.output @@ -334,7 +367,9 @@ def test_plan_verbose(runner, tmp_path): # Input: `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--verbose"], input="y\n" + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--verbose"], + input="y\n", ) assert result.exit_code == 0 assert_backfill_success(result) @@ -371,7 +406,9 @@ def test_plan_dev(runner, tmp_path): # Input: enter for backfill start date prompt, enter for end date prompt, `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev"], input="\n\ny\n" + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev"], + input="\n\ny\n", ) assert_plan_success(result, "dev") @@ -382,7 +419,16 @@ def test_plan_dev_start_date(runner, tmp_path): # Input: enter for backfill end date prompt, `y` to apply and backfill result = runner.invoke( cli, - ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev", "--start", "2023-01-01"], + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "dev", + "--start", + "2023-01-01", + ], input="\ny\n", ) assert_plan_success(result, "dev") @@ -396,12 +442,24 @@ def test_plan_dev_end_date(runner, tmp_path): # Input: enter for backfill start date prompt, `y` to apply and backfill result = runner.invoke( cli, - ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev", "--end", "2023-01-01"], + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "dev", + "--end", + "2023-01-01", + ], input="\ny\n", ) assert_plan_success(result, "dev") assert "sqlmesh_example__dev.full_model: [full refresh]" in result.output - assert "sqlmesh_example__dev.incremental_model: [2020-01-01 - 2023-01-01]" in result.output + assert ( + "sqlmesh_example__dev.incremental_model: [2020-01-01 - 2023-01-01]" + in result.output + ) def test_plan_dev_create_from_virtual(runner, tmp_path): @@ -538,7 +596,16 @@ def test_plan_dev_no_prompts(runner, tmp_path): # plan for non-prod environment doesn't prompt for dates but prompts to apply result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev", "--no-prompts"] + cli, + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "dev", + "--no-prompts", + ], ) assert "Apply - Backfill Tables [y/n]: " in result.output assert "Physical layer updated" not in result.output @@ -552,7 +619,15 @@ def test_plan_dev_auto_apply(runner, tmp_path): # Input: enter for backfill start date prompt, enter for end date prompt result = runner.invoke( cli, - ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev", "--auto-apply"], + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "dev", + "--auto-apply", + ], input="\n\n", ) assert_plan_success(result, "dev") @@ -563,7 +638,9 @@ def test_plan_dev_no_changes(runner, tmp_path): init_prod_and_backfill(runner, tmp_path) # Error if no changes made and `--include-unmodified` is not passed - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev"] + ) assert result.exit_code == 1 assert ( "Error: Creating a new environment requires a change, but project files match the `prod` environment. Make a change or use the --include-unmodified flag to create a new environment without changes." @@ -574,7 +651,15 @@ def test_plan_dev_no_changes(runner, tmp_path): # Input: `y` to apply and virtual update result = runner.invoke( cli, - ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "dev", "--include-unmodified"], + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "dev", + "--include-unmodified", + ], input="y\n", ) assert result.exit_code == 0 @@ -645,7 +730,10 @@ def test_plan_nonbreaking(runner, tmp_path): assert result.exit_code == 0 assert "Differences from the `prod` environment" in result.output assert "+ 'a' AS new_col" in result.output - assert "Directly Modified: sqlmesh_example.incremental_model (Non-breaking)" in result.output + assert ( + "Directly Modified: sqlmesh_example.incremental_model (Non-breaking)" + in result.output + ) assert "sqlmesh_example.full_model (Indirect Non-breaking)" in result.output assert "sqlmesh_example.incremental_model [insert" in result.output assert "sqlmesh_example.full_model [full refresh" not in result.output @@ -661,7 +749,14 @@ def test_plan_nonbreaking_noautocategorization(runner, tmp_path): # Input: `2` to classify change as non-breaking, `y` to apply and backfill result = runner.invoke( cli, - ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--no-auto-categorization"], + [ + "--log-file-dir", + tmp_path, + "--paths", + tmp_path, + "plan", + "--no-auto-categorization", + ], input="2\ny\n", ) assert result.exit_code == 0 @@ -684,7 +779,9 @@ def test_plan_nonbreaking_nodiff(runner, tmp_path): # Input: `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--no-diff"], input="y\n" + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--no-diff"], + input="y\n", ) assert result.exit_code == 0 assert "+ 'a' AS new_col" not in result.output @@ -700,7 +797,9 @@ def test_plan_breaking(runner, tmp_path): # full_model change makes test fail, so we pass `--skip-tests` # Input: `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--skip-tests"], input="y\n" + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--skip-tests"], + input="y\n", ) assert result.exit_code == 0 assert "+ item_id + 1 AS item_id," in result.output @@ -738,11 +837,15 @@ def test_plan_dev_select(runner, tmp_path): # incremental_model diff present assert "+ 'a' AS new_col" in result.output assert ( - "Directly Modified: sqlmesh_example__dev.incremental_model (Non-breaking)" in result.output + "Directly Modified: sqlmesh_example__dev.incremental_model (Non-breaking)" + in result.output ) # full_model diff not present assert "+ item_id + 1 AS item_id," not in result.output - assert "Directly Modified: sqlmesh_example__dev.full_model (Breaking)" not in result.output + assert ( + "Directly Modified: sqlmesh_example__dev.full_model (Breaking)" + not in result.output + ) # only incremental_model backfilled assert "sqlmesh_example__dev.incremental_model [insert" in result.output assert "sqlmesh_example__dev.full_model [full refresh" not in result.output @@ -777,10 +880,13 @@ def test_plan_dev_backfill(runner, tmp_path): assert_new_env(result, "dev", initialize=False) # both model diffs present assert "+ item_id + 1 AS item_id," in result.output - assert "Directly Modified: sqlmesh_example__dev.full_model (Breaking)" in result.output + assert ( + "Directly Modified: sqlmesh_example__dev.full_model (Breaking)" in result.output + ) assert "+ 'a' AS new_col" in result.output assert ( - "Directly Modified: sqlmesh_example__dev.incremental_model (Non-breaking)" in result.output + "Directly Modified: sqlmesh_example__dev.incremental_model (Non-breaking)" + in result.output ) # only incremental_model backfilled assert "sqlmesh_example__dev.incremental_model [insert" in result.output @@ -792,7 +898,9 @@ def test_run_no_prod(runner, tmp_path): create_example_project(tmp_path) # Error if no env specified and `prod` doesn't exist - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run"] + ) assert result.exit_code == 1 assert "Error: Environment 'prod' was not found." in result.output @@ -811,7 +919,9 @@ def test_run_dev(runner, tmp_path, flag): ) # Confirm backfill occurs when we run non-backfilled dev env - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run", "dev"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run", "dev"] + ) assert result.exit_code == 0 assert_model_batches_executed(result) @@ -822,7 +932,9 @@ def test_run_cron_not_elapsed(runner, tmp_path, caplog): init_prod_and_backfill(runner, tmp_path) # No error if `prod` environment exists and cron has not elapsed - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run"] + ) assert result.exit_code == 0 assert ( @@ -841,7 +953,9 @@ def test_run_cron_elapsed(runner, tmp_path): # Run `prod` environment with daily cron elapsed traveler.move_to("2023-01-02 00:01:00 UTC") - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "run"] + ) assert result.exit_code == 0 assert_model_batches_executed(result) @@ -858,7 +972,9 @@ def test_clean(runner, tmp_path): assert len(list(cache_path.iterdir())) > 0 # Invoke the clean command - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "clean"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "clean"] + ) # Confirm cache was cleared assert result.exit_code == 0 @@ -881,14 +997,18 @@ def test_table_name(runner, tmp_path): ], ) assert result.exit_code == 0 - assert result.output.startswith("db.sqlmesh__sqlmesh_example.sqlmesh_example__full_model__") + assert result.output.startswith( + "db.sqlmesh__sqlmesh_example.sqlmesh_example__full_model__" + ) def test_info_on_new_project_does_not_create_state_sync(runner, tmp_path): create_example_project(tmp_path) # Invoke the info command - result = runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "info"]) + result = runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "info"] + ) assert result.exit_code == 0 context = Context(paths=tmp_path) @@ -911,7 +1031,16 @@ def test_dlt_pipeline_errors(runner, tmp_path): # Error if the pipeline provided is not correct result = runner.invoke( cli, - ["--paths", tmp_path, "init", "-t", "dlt", "--dlt-pipeline", "missing_pipeline", "duckdb"], + [ + "--paths", + tmp_path, + "init", + "-t", + "dlt", + "--dlt-pipeline", + "missing_pipeline", + "duckdb", + ], ) assert "Error: Could not attach to pipeline" in result.output @@ -1033,7 +1162,9 @@ def test_dlt_pipeline(runner, tmp_path): exec(file.read()) # This should fail since it won't be able to locate the pipeline in this path - with pytest.raises(ClickException, match=r".*Could not attach to pipeline*") as excinfo: + with pytest.raises( + ClickException, match=r".*Could not attach to pipeline*" + ) as excinfo: init_example_project( tmp_path, "duckdb", @@ -1051,7 +1182,11 @@ def test_dlt_pipeline(runner, tmp_path): # By setting the pipelines path where the pipeline directory is located, it should work dlt_path = get_dlt_pipelines_dir() init_example_project( - tmp_path, "duckdb", template=ProjectTemplate.DLT, pipeline="sushi", dlt_path=dlt_path + tmp_path, + "duckdb", + template=ProjectTemplate.DLT, + pipeline="sushi", + dlt_path=dlt_path, ) expected_config = f"""# --- Gateway Connection --- @@ -1108,7 +1243,9 @@ def test_dlt_pipeline(runner, tmp_path): dlt_sushi_types_model_path = tmp_path / "models/incremental_sushi_types.sql" dlt_loads_model_path = tmp_path / "models/incremental__dlt_loads.sql" dlt_waiters_model_path = tmp_path / "models/incremental_waiters.sql" - dlt_sushi_fillings_model_path = tmp_path / "models/incremental_sushi_menu__fillings.sql" + dlt_sushi_fillings_model_path = ( + tmp_path / "models/incremental_sushi_menu__fillings.sql" + ) dlt_sushi_twice_nested_model_path = ( tmp_path / "models/incremental_sushi_menu__details__ingredients.sql" ) @@ -1180,7 +1317,8 @@ def test_dlt_pipeline(runner, tmp_path): try: # Plan prod and backfill result = runner.invoke( - cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--auto-apply"] + cli, + ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan", "--auto-apply"], ) assert result.exit_code == 0 @@ -1206,9 +1344,9 @@ def test_dlt_pipeline(runner, tmp_path): # Update to generate a specific model: sushi_types. # Also validate using the dlt_path that the pipelines are located. - assert generate_dlt_models(context, "sushi", ["sushi_types"], False, dlt_path) == [ - "sushi_dataset_sqlmesh.incremental_sushi_types" - ] + assert generate_dlt_models( + context, "sushi", ["sushi_types"], False, dlt_path + ) == ["sushi_dataset_sqlmesh.incremental_sushi_types"] # Only the sushi_types should be generated now assert not dlt_waiters_model_path.exists() @@ -1290,12 +1428,17 @@ def test_environments(runner, tmp_path): ], ) assert result.exit_code == 0 - assert f"Number of SQLMesh environments are: 2\ndev - {ttl}\ndev2 - {ttl}\n" in result.output + assert ( + f"Number of SQLMesh environments are: 2\ndev - {ttl}\ndev2 - {ttl}\n" + in result.output + ) # Example project models have start dates, so there are no date prompts # for the `prod` environment. # Input: `y` to apply and backfill - runner.invoke(cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan"], input="y\n") + runner.invoke( + cli, ["--log-file-dir", tmp_path, "--paths", tmp_path, "plan"], input="y\n" + ) result = runner.invoke( cli, [ @@ -1317,12 +1460,10 @@ def test_lint(runner, tmp_path): create_example_project(tmp_path) with open(tmp_path / "config.yaml", "a", encoding="utf-8") as f: - f.write( - """linter: + f.write("""linter: enabled: True rules: "ALL" -""" - ) +""") result = runner.invoke(cli, ["--paths", tmp_path, "lint"]) assert result.output.count("Linter errors for") == 2 @@ -1373,7 +1514,15 @@ def test_state_export(runner: CliRunner, tmp_path: Path) -> None: # export it result = runner.invoke( cli, - ["--paths", str(tmp_path), "state", "export", "-o", str(state_export_file), "--no-confirm"], + [ + "--paths", + str(tmp_path), + "state", + "export", + "-o", + str(state_export_file), + "--no-confirm", + ], catch_exceptions=False, ) assert result.exit_code == 0 @@ -1406,16 +1555,14 @@ def test_state_export_specific_environments(runner: CliRunner, tmp_path: Path) - ) assert result.exit_code == 0 - (tmp_path / "models" / "new_model.sql").write_text( - """ + (tmp_path / "models" / "new_model.sql").write_text(""" MODEL ( name sqlmesh_example.new_model, kind FULL ); SELECT 1; - """ - ) + """) # create dev env with new model result = runner.invoke( @@ -1551,7 +1698,15 @@ def test_state_import(runner: CliRunner, tmp_path: Path) -> None: # export it result = runner.invoke( cli, - ["--paths", str(tmp_path), "state", "export", "-o", str(state_export_file), "--no-confirm"], + [ + "--paths", + str(tmp_path), + "state", + "export", + "-o", + str(state_export_file), + "--no-confirm", + ], catch_exceptions=False, ) assert result.exit_code == 0 @@ -1559,7 +1714,15 @@ def test_state_import(runner: CliRunner, tmp_path: Path) -> None: # import it back result = runner.invoke( cli, - ["--paths", str(tmp_path), "state", "import", "-i", str(state_export_file), "--no-confirm"], + [ + "--paths", + str(tmp_path), + "state", + "import", + "-i", + str(state_export_file), + "--no-confirm", + ], catch_exceptions=False, ) assert result.exit_code == 0 @@ -1602,16 +1765,14 @@ def test_state_import_replace(runner: CliRunner, tmp_path: Path) -> None: ) assert result.exit_code == 0 - (tmp_path / "models" / "new_model.sql").write_text( - """ + (tmp_path / "models" / "new_model.sql").write_text(""" MODEL ( name sqlmesh_example.new_model, kind FULL ); SELECT 1; - """ - ) + """) # create dev with new model result = runner.invoke( @@ -1715,7 +1876,15 @@ def test_state_import_local(runner: CliRunner, tmp_path: Path) -> None: # import should fail - local state is not importable result = runner.invoke( cli, - ["--paths", str(tmp_path), "state", "import", "-i", str(state_export_file), "--no-confirm"], + [ + "--paths", + str(tmp_path), + "state", + "import", + "-i", + str(state_export_file), + "--no-confirm", + ], catch_exceptions=False, ) assert result.exit_code == 1 @@ -1751,11 +1920,20 @@ def test_ignore_warnings(runner: CliRunner, tmp_path: Path) -> None: select 1 as a; """) - audit_warning = "[WARNING] sqlmesh_example.full_model: 'full_nonblocking_audit' audit error: " + audit_warning = ( + "[WARNING] sqlmesh_example.full_model: 'full_nonblocking_audit' audit error: " + ) result = runner.invoke( cli, - ["--paths", str(tmp_path), "plan", "--no-prompts", "--auto-apply", "--skip-tests"], + [ + "--paths", + str(tmp_path), + "plan", + "--no-prompts", + "--auto-apply", + "--skip-tests", + ], ) assert result.exit_code == 0 assert audit_warning in result.output @@ -1800,7 +1978,15 @@ def test_table_diff_schema_diff_ignore_case(runner: CliRunner, tmp_path: Path): # ignore case result = runner.invoke( cli, - ["--paths", str(tmp_path), "table_diff", "t1:t2", "-o", "id", "--schema-diff-ignore-case"], + [ + "--paths", + str(tmp_path), + "table_diff", + "t1:t2", + "-o", + "id", + "--schema-diff-ignore-case", + ], ) assert result.exit_code == 0 stripped_output = "".join((x for x in result.output if x in string.printable)) @@ -1824,7 +2010,10 @@ def test_init_bad_template(runner: CliRunner, tmp_path: Path): ["--paths", str(tmp_path), "init", "-t", "invalid_template"], ) assert result.exit_code == 1 - assert "Invalid project template 'invalid_template'. Please specify one of " in result.output + assert ( + "Invalid project template 'invalid_template'. Please specify one of " + in result.output + ) # empty template should not produce example project files @@ -1872,7 +2061,8 @@ def test_init_interactive_invalid_int(runner: CliRunner, tmp_path: Path): ) assert result.exit_code == 0 assert ( - "'0' is not a valid project type number - please enter a number between 1" in result.output + "'0' is not a valid project type number - please enter a number between 1" + in result.output ) @@ -1916,7 +2106,9 @@ def test_init_interactive_cli_mode_simple(runner: CliRunner, tmp_path: Path): assert "no_diff: true" in config_path.read_text() -def test_init_interactive_engine_install_msg(runner: CliRunner, tmp_path: Path, monkeypatch): +def test_init_interactive_engine_install_msg( + runner: CliRunner, tmp_path: Path, monkeypatch +): monkeypatch.setattr("sqlmesh.utils.rich.console.width", 80) # Engine install text should not appear for built-in engines like DuckDB @@ -2076,8 +2268,7 @@ def test_signals(runner: CliRunner, tmp_path: Path): signals_dir.mkdir(exist_ok=True) # Create signal definitions - (signals_dir / "signal.py").write_text( - """from sqlmesh import signal + (signals_dir / "signal.py").write_text("""from sqlmesh import signal @signal() def only_first_two_ready(batch): if len(batch) > 2: @@ -2087,12 +2278,10 @@ def only_first_two_ready(batch): @signal() def none_ready(batch): return False -""" - ) +""") # Create model with signals - (tmp_path / "models" / "model_with_signals.sql").write_text( - """MODEL ( + (tmp_path / "models" / "model_with_signals.sql").write_text("""MODEL ( name sqlmesh_example.model_with_signals, kind INCREMENTAL_BY_TIME_RANGE ( time_column ds @@ -2115,12 +2304,10 @@ def none_ready(batch): ('2023-01-01') AS t(ds) WHERE ds::DATE BETWEEN @start_ds AND @end_ds -""" - ) +""") # Create model with no ready intervals - (tmp_path / "models" / "model_with_unready.sql").write_text( - """MODEL ( + (tmp_path / "models" / "model_with_unready.sql").write_text("""MODEL ( name sqlmesh_example.model_with_unready, kind INCREMENTAL_BY_TIME_RANGE ( time_column ds @@ -2143,8 +2330,7 @@ def none_ready(batch): ('2023-01-01') AS t(ds) WHERE ds::DATE BETWEEN @start_ds AND @end_ds -""" - ) +""") # Test 1: Normal plan flow with --no-prompts --auto-apply result = runner.invoke( @@ -2205,7 +2391,9 @@ def none_ready(batch): # for the `prod` environment. # Input: `y` to apply and backfill result = runner.invoke( - cli, ["--log-file-dir", str(tmp_path), "--paths", str(tmp_path), "plan"], input="y\n" + cli, + ["--log-file-dir", str(tmp_path), "--paths", str(tmp_path), "plan"], + input="y\n", ) assert_plan_success(result) @@ -2311,17 +2499,25 @@ def _setup_local_only_project(tmp_path, mocker): def test_format_runs_without_state(runner: CliRunner, tmp_path: Path, mocker): mock = _setup_local_only_project(tmp_path, mocker) result = runner.invoke(cli, ["--paths", str(tmp_path), "format"]) - assert result.exit_code == 0, f"Format failed: {result.output}\nException: {result.exception}" + assert ( + result.exit_code == 0 + ), f"Format failed: {result.output}\nException: {result.exception}" mock.assert_not_called() -def test_format_runs_without_state_multi_repo_partial(runner: CliRunner, copy_to_temp_path, mocker): +def test_format_runs_without_state_multi_repo_partial( + runner: CliRunner, copy_to_temp_path, mocker +): """Format one repo of a multi-repo project whose upstream models live only in prod state.""" repo_2 = copy_to_temp_path("examples/multi")[0] / "repo_2" mock = _patch_state_access(mocker) - result = runner.invoke(cli, ["--gateway", "memory", "--paths", str(repo_2), "format"]) - assert result.exit_code == 0, f"Format failed: {result.output}\nException: {result.exception}" + result = runner.invoke( + cli, ["--gateway", "memory", "--paths", str(repo_2), "format"] + ) + assert ( + result.exit_code == 0 + ), f"Format failed: {result.output}\nException: {result.exception}" mock.assert_not_called() @@ -2334,12 +2530,12 @@ def test_lint_still_loads_state(runner: CliRunner, tmp_path: Path, mocker): assert init_spy.called, "Context was never constructed" for call in init_spy.call_args_list: - assert "load_state" in call.kwargs, ( - "CLI didn't pass load_state= explicitly; missing kwarg defaults to True silently" - ) - assert call.kwargs["load_state"] is True, ( - f"Context was constructed with load_state={call.kwargs['load_state']} for `lint`" - ) + assert ( + "load_state" in call.kwargs + ), "CLI didn't pass load_state= explicitly; missing kwarg defaults to True silently" + assert ( + call.kwargs["load_state"] is True + ), f"Context was constructed with load_state={call.kwargs['load_state']} for `lint`" assert mock.called, "state-sync was never accessed during `lint`" @@ -2349,15 +2545,17 @@ def test_lint_local_runs_without_state(runner: CliRunner, tmp_path: Path, mocker result = runner.invoke(cli, ["--paths", str(tmp_path), "lint", "--local"]) - assert result.exit_code == 0, f"Lint failed: {result.output}\nException: {result.exception}" + assert ( + result.exit_code == 0 + ), f"Lint failed: {result.output}\nException: {result.exception}" assert init_spy.called, "Context was never constructed" for call in init_spy.call_args_list: - assert "load_state" in call.kwargs, ( - "CLI didn't pass load_state= explicitly; missing kwarg defaults to True silently" - ) - assert call.kwargs["load_state"] is False, ( - f"Context was constructed with load_state={call.kwargs['load_state']} for `lint --local`" - ) + assert ( + "load_state" in call.kwargs + ), "CLI didn't pass load_state= explicitly; missing kwarg defaults to True silently" + assert ( + call.kwargs["load_state"] is False + ), f"Context was constructed with load_state={call.kwargs['load_state']} for `lint --local`" mock.assert_not_called() @@ -2371,10 +2569,12 @@ def test_local_only_commands_skip_state_multiple_paths( _create_local_only_project(project_b, "proj_b") mock = _patch_state_access(mocker) - result = runner.invoke(cli, ["--paths", str(project_a), "--paths", str(project_b), command]) - assert result.exit_code == 0, ( - f"{command} failed: {result.output}\nException: {result.exception}" + result = runner.invoke( + cli, ["--paths", str(project_a), "--paths", str(project_b), command] ) + assert ( + result.exit_code == 0 + ), f"{command} failed: {result.output}\nException: {result.exception}" mock.assert_not_called() @@ -2387,12 +2587,12 @@ def test_plan_still_loads_state(runner: CliRunner, tmp_path: Path, mocker): assert init_spy.called, "Context was never constructed" for call in init_spy.call_args_list: - assert "load_state" in call.kwargs, ( - "CLI didn't pass load_state= explicitly; missing kwarg defaults to True silently" - ) - assert call.kwargs["load_state"] is True, ( - f"Context was constructed with load_state={call.kwargs['load_state']} for `plan`" - ) + assert ( + "load_state" in call.kwargs + ), "CLI didn't pass load_state= explicitly; missing kwarg defaults to True silently" + assert ( + call.kwargs["load_state"] is True + ), f"Context was constructed with load_state={call.kwargs['load_state']} for `plan`" assert mock.called, "state-sync was never accessed during `plan`" @@ -2440,5 +2640,7 @@ def test_format_does_not_open_state_connection( ) result = runner.invoke(cli, ["--paths", str(tmp_path), "format"]) - assert result.exit_code == 0, f"Format failed: {result.output}\nException: {result.exception}" + assert ( + result.exit_code == 0 + ), f"Format failed: {result.output}\nException: {result.exception}" mock.assert_not_called() diff --git a/tests/cli/test_integration_cli.py b/tests/cli/test_integration_cli.py index 5d000b9d8b..19ef674932 100644 --- a/tests/cli/test_integration_cli.py +++ b/tests/cli/test_integration_cli.py @@ -1,12 +1,14 @@ +import shutil +import site +import subprocess import typing as t +import uuid from pathlib import Path + import pytest -import subprocess + from sqlmesh.cli.project_init import init_example_project from sqlmesh.utils import yaml -import shutil -import site -import uuid pytestmark = pytest.mark.slow @@ -29,7 +31,9 @@ def invoke_cli(tmp_path: Path) -> InvokeCliType: ["which", "sqlmesh"], capture_output=True, text=True ).stdout.strip() - def _invoke(sqlmesh_args: t.List[str], **kwargs: t.Any) -> subprocess.CompletedProcess: + def _invoke( + sqlmesh_args: t.List[str], **kwargs: t.Any + ) -> subprocess.CompletedProcess: return subprocess.run( args=[sqlmesh_bin] + sqlmesh_args, # set the working directory to the isolated temp dir for this test @@ -146,7 +150,9 @@ def do_something(evaluator): # render the query to ensure our macro is being invoked result = invoke_cli(["render", "example.test_model"]) assert result.returncode == 0 - assert """SELECT 'value from site-packages' AS "a\"""" in " ".join(result.stdout.split()) + assert """SELECT 'value from site-packages' AS "a\"""" in " ".join( + result.stdout.split() + ) # clear cache to ensure we are forced to reload everything assert invoke_cli(["clean"]).returncode == 0 @@ -266,7 +272,13 @@ def do_something(evaluator): # the invalid snapshot in state should not prevent a plan if --select-model is used on it (since the local version can be rendered) result = invoke_cli( - ["plan", "--select-model", "sqlmesh_example.test_model", "--no-prompts", "--skip-tests"], + [ + "plan", + "--select-model", + "sqlmesh_example.test_model", + "--no-prompts", + "--skip-tests", + ], input="n", # for the apply backfill (y/n) prompt ) assert result.returncode == 0 @@ -301,7 +313,8 @@ def do_something(evaluator): log_file_contents = last_log_file_contents() assert f"ModuleNotFoundError: No module named '{package_name}'" in log_file_contents assert ( - "The above exception was the direct cause of the following exception:" in log_file_contents + "The above exception was the direct cause of the following exception:" + in log_file_contents ) diff --git a/tests/cli/test_project_init.py b/tests/cli/test_project_init.py index 12b42705e1..972a59adc6 100644 --- a/tests/cli/test_project_init.py +++ b/tests/cli/test_project_init.py @@ -1,17 +1,21 @@ -import pytest from pathlib import Path -from sqlmesh.utils.errors import SQLMeshError -from sqlmesh.cli.project_init import init_example_project, ProjectTemplate -from sqlmesh.utils import yaml -from sqlmesh.core.context import Context + +import pytest + +from sqlmesh.cli.project_init import ProjectTemplate, init_example_project from sqlmesh.core.config.common import VirtualEnvironmentMode +from sqlmesh.core.context import Context +from sqlmesh.utils import yaml +from sqlmesh.utils.errors import SQLMeshError def test_project_init_dbt(tmp_path: Path): assert not len(list(tmp_path.glob("**/*"))) with pytest.raises(SQLMeshError, match=r"Required dbt project file.*not found"): - init_example_project(path=tmp_path, engine_type=None, template=ProjectTemplate.DBT) + init_example_project( + path=tmp_path, engine_type=None, template=ProjectTemplate.DBT + ) with (tmp_path / "dbt_project.yml").open("w") as f: yaml.dump({"name": "jaffle_shop"}, f) @@ -26,7 +30,10 @@ def test_project_init_dbt(tmp_path: Path): assert "start: " in sqlmesh_config.read_text() with (tmp_path / "profiles.yml").open("w") as f: - yaml.dump({"jaffle_shop": {"target": "dev", "outputs": {"dev": {"type": "duckdb"}}}}, f) + yaml.dump( + {"jaffle_shop": {"target": "dev", "outputs": {"dev": {"type": "duckdb"}}}}, + f, + ) ctx = Context(paths=tmp_path) assert ctx.config.model_defaults.start diff --git a/tests/conftest.py b/tests/conftest.py index 4d1bb23577..8781bd6de5 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,8 +1,9 @@ from __future__ import annotations - import datetime import logging +import os +import shutil import typing as t import uuid from contextlib import nullcontext @@ -11,8 +12,6 @@ from tempfile import TemporaryDirectory from unittest import mock from unittest.mock import PropertyMock -import os -import shutil import duckdb # noqa: TID253 import pandas as pd # noqa: TID253 @@ -23,29 +22,26 @@ from sqlglot.helper import ensure_list from sqlglot.optimizer.normalize_identifiers import normalize_identifiers -from sqlmesh.core.config import Config, BaseDuckDBConnectionConfig, DuckDBConnectionConfig +from sqlmesh.core import lineage +from sqlmesh.core.config import (BaseDuckDBConnectionConfig, Config, + DuckDBConnectionConfig) from sqlmesh.core.config.connection import ConnectionConfig from sqlmesh.core.context import Context from sqlmesh.core.engine_adapter import MSSQLEngineAdapter, SparkEngineAdapter from sqlmesh.core.engine_adapter.base import EngineAdapter +from sqlmesh.core.engine_adapter.shared import CatalogSupport from sqlmesh.core.environment import EnvironmentNamingInfo -from sqlmesh.core import lineage from sqlmesh.core.macros import macro from sqlmesh.core.model import IncrementalByTimeRangeKind, SqlModel, model -from sqlmesh.core.model.kind import OnDestructiveChange, OnAdditiveChange -from sqlmesh.core.plan import BuiltInPlanEvaluator, Plan, stages as plan_stages -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Node, - Snapshot, - SnapshotChangeCategory, - SnapshotDataVersion, - SnapshotFingerprint, -) +from sqlmesh.core.model.kind import OnAdditiveChange, OnDestructiveChange +from sqlmesh.core.plan import BuiltInPlanEvaluator, Plan +from sqlmesh.core.plan import stages as plan_stages +from sqlmesh.core.snapshot import (DeployabilityIndex, Node, Snapshot, + SnapshotChangeCategory, SnapshotDataVersion, + SnapshotFingerprint) from sqlmesh.utils import random_id from sqlmesh.utils.date import TimeLike, to_date from sqlmesh.utils.windows import IS_WINDOWS, fix_windows_path -from sqlmesh.core.engine_adapter.shared import CatalogSupport T = t.TypeVar("T", bound=EngineAdapter) @@ -128,7 +124,9 @@ def _system_schema_filter(self, col: str) -> str: return f"{col} not in ('information_schema', 'pg_catalog', 'main')" @staticmethod - def _get_single_col(query: str, col: str, engine_adapter: EngineAdapter) -> t.List[t.Any]: + def _get_single_col( + query: str, col: str, engine_adapter: EngineAdapter + ) -> t.List[t.Any]: return list(engine_adapter.fetchdf(query)[col].to_dict().values()) @@ -139,7 +137,9 @@ def __init__(self, engine_adapter: EngineAdapter, sushi_schema_name: str): @classmethod def from_context(cls, context: Context, sushi_schema_name: str = "sushi"): - return cls(engine_adapter=context.engine_adapter, sushi_schema_name=sushi_schema_name) + return cls( + engine_adapter=context.engine_adapter, sushi_schema_name=sushi_schema_name + ) def validate( self, @@ -179,7 +179,9 @@ def validate( "sushi.customer_revenue_lifetime", ): env_name = f"__{env_name}" if env_name else "" - full_table_path = f"{self.sushi_schema_name}{env_name}.customer_revenue_lifetime" + full_table_path = ( + f"{self.sushi_schema_name}{env_name}.customer_revenue_lifetime" + ) query = f"SELECT event_date, count(*) AS the_count FROM {full_table_path} group by event_date order by 2 desc, 1 desc" results = self.engine_adapter.fetchdf( parse_one(query), quote_identifiers=True @@ -190,7 +192,8 @@ def validate( # this creates Pandas Timestamp objects expected_dates = [ - pd.to_datetime(end_date - datetime.timedelta(days=x)) for x in range(num_days_diff) + pd.to_datetime(end_date - datetime.timedelta(days=x)) + for x in range(num_days_diff) ] # all engines but duckdb and clickhouse fetch dates as datetime.date objects if dialect and dialect not in ("duckdb", "clickhouse"): @@ -223,7 +226,9 @@ def pytest_collection_modifyitems(items, *args, **kwargs): # Ignore all local config files @pytest.fixture(scope="session", autouse=True) def ignore_local_config_files(): - with mock.patch("sqlmesh.core.constants.SQLMESH_PATH", Path(TemporaryDirectory().name)): + with mock.patch( + "sqlmesh.core.constants.SQLMESH_PATH", Path(TemporaryDirectory().name) + ): yield @@ -261,7 +266,7 @@ def rescope_lineage_cache(request): @pytest.fixture(autouse=True) def reset_console(): - from sqlmesh.core.console import set_console, NoopConsole, get_console + from sqlmesh.core.console import NoopConsole, get_console, set_console orig_console = get_console() set_console(NoopConsole()) @@ -291,12 +296,16 @@ def push_plan(context: Context, plan: Plan) -> None: plan_evaluator.visit_create_snapshot_records_stage(stage, evaluatable_plan) elif isinstance(stage, plan_stages.PhysicalLayerSchemaCreationStage): stage.deployability_index = deployability_index - plan_evaluator.visit_physical_layer_schema_creation_stage(stage, evaluatable_plan) + plan_evaluator.visit_physical_layer_schema_creation_stage( + stage, evaluatable_plan + ) elif isinstance(stage, plan_stages.PhysicalLayerUpdateStage): stage.deployability_index = deployability_index plan_evaluator.visit_physical_layer_update_stage(stage, evaluatable_plan) elif isinstance(stage, plan_stages.EnvironmentRecordUpdateStage): - plan_evaluator.visit_environment_record_update_stage(stage, evaluatable_plan) + plan_evaluator.visit_environment_record_update_stage( + stage, evaluatable_plan + ) elif isinstance(stage, plan_stages.VirtualLayerUpdateStage): stage.deployability_index = deployability_index plan_evaluator.visit_virtual_layer_update_stage(stage, evaluatable_plan) @@ -351,7 +360,9 @@ def sushi_test_dbt_context(init_and_plan_context) -> Context: @pytest.fixture() -def sushi_no_default_catalog(mocker: MockerFixture, init_and_plan_context: t.Callable) -> Context: +def sushi_no_default_catalog( + mocker: MockerFixture, init_and_plan_context: t.Callable +) -> Context: mocker.patch( "sqlmesh.core.engine_adapter.base.EngineAdapter.default_catalog", PropertyMock(return_value=None), @@ -402,7 +413,9 @@ def _assert_exp_eq( @pytest.fixture def make_snapshot() -> t.Callable: - def _make_function(node: Node, version: t.Optional[str] = None, **kwargs) -> Snapshot: + def _make_function( + node: Node, version: t.Optional[str] = None, **kwargs + ) -> Snapshot: return Snapshot.from_node( node, **{ # type: ignore @@ -430,7 +443,9 @@ def _make_function( dialect="duckdb", query=parse_one(old_query), kind=IncrementalByTimeRangeKind( - time_column="ds", forward_only=True, on_destructive_change=on_destructive_change + time_column="ds", + forward_only=True, + on_destructive_change=on_destructive_change, ), ) ) @@ -441,7 +456,9 @@ def _make_function( dialect="duckdb", query=parse_one(new_query), kind=IncrementalByTimeRangeKind( - time_column="ds", forward_only=True, on_destructive_change=on_destructive_change + time_column="ds", + forward_only=True, + on_destructive_change=on_destructive_change, ), ) ) @@ -476,7 +493,9 @@ def _make_function( dialect="duckdb", query=parse_one(old_query), kind=IncrementalByTimeRangeKind( - time_column="ds", forward_only=True, on_additive_change=on_additive_change + time_column="ds", + forward_only=True, + on_additive_change=on_additive_change, ), ) ) @@ -487,7 +506,9 @@ def _make_function( dialect="duckdb", query=parse_one(new_query), kind=IncrementalByTimeRangeKind( - time_column="ds", forward_only=True, on_additive_change=on_additive_change + time_column="ds", + forward_only=True, + on_additive_change=on_additive_change, ), ) ) @@ -519,7 +540,9 @@ def sushi_data_validator(sushi_context: Context) -> SushiDataValidator: @pytest.fixture -def sushi_fixed_date_data_validator(sushi_context_fixed_date: Context) -> SushiDataValidator: +def sushi_fixed_date_data_validator( + sushi_context_fixed_date: Context, +) -> SushiDataValidator: return SushiDataValidator.from_context(sushi_context_fixed_date) @@ -552,7 +575,9 @@ def _make_function( if isinstance(adapter, MSSQLEngineAdapter): mocker.patch( "sqlmesh.core.engine_adapter.mssql.MSSQLEngineAdapter.catalog_support", - new_callable=PropertyMock(return_value=CatalogSupport.REQUIRES_SET_CATALOG), + new_callable=PropertyMock( + return_value=CatalogSupport.REQUIRES_SET_CATALOG + ), ) if patch_get_data_objects: mocker.patch.object(adapter, "_get_data_objects", return_value=[]) @@ -571,7 +596,9 @@ def ignore(src, names): return [name for name in names if name == ".cache"] def _make_function( - paths: t.Union[str, Path, t.List[t.Union[str, Path]], t.Tuple[t.Union[str, Path], ...]], + paths: t.Union[ + str, Path, t.List[t.Union[str, Path]], t.Tuple[t.Union[str, Path], ...] + ], ) -> t.List[Path]: paths = ensure_list(paths) all_paths = [Path(p) for p in paths] @@ -625,7 +652,9 @@ def delete_cache(project_paths: str | t.List[str]) -> None: def make_temp_table_name(mocker: MockerFixture) -> t.Callable: def _make_function(table_name: str, random_id: str) -> exp.Table: temp_table = exp.to_table(table_name) - temp_table.set("this", exp.to_identifier(f"__temp_{temp_table.name}_{random_id}")) + temp_table.set( + "this", exp.to_identifier(f"__temp_{temp_table.name}_{random_id}") + ) return temp_table return _make_function @@ -641,14 +670,18 @@ def set_default_connection(request): else: original_get_connection = Config.get_connection - def _lax_get_connection(self, gateway_name: t.Optional[str] = None) -> ConnectionConfig: + def _lax_get_connection( + self, gateway_name: t.Optional[str] = None + ) -> ConnectionConfig: try: connection = original_get_connection(self, gateway_name) except: connection = DuckDBConnectionConfig() return connection - ctx = mock.patch("sqlmesh.core.config.Config.get_connection", _lax_get_connection) + ctx = mock.patch( + "sqlmesh.core.config.Config.get_connection", _lax_get_connection + ) with ctx: yield diff --git a/tests/core/analytics/test_collector.py b/tests/core/analytics/test_collector.py index e87704e694..2eaa416f46 100644 --- a/tests/core/analytics/test_collector.py +++ b/tests/core/analytics/test_collector.py @@ -26,7 +26,9 @@ def collector(mocker: MockerFixture) -> AnalyticsCollector: "hybrid", ], ) -def test_on_project_loaded(collector: AnalyticsCollector, mocker: MockerFixture, project_type): +def test_on_project_loaded( + collector: AnalyticsCollector, mocker: MockerFixture, project_type +): collector.on_project_loaded( project_type=project_type, models_count=1, @@ -43,7 +45,9 @@ def test_on_project_loaded(collector: AnalyticsCollector, mocker: MockerFixture, from dbt.version import __version__ as dbt_version - version = ', "dbt_version": "' + dbt_version + '"' if project_type != c.NATIVE else "" + version = ( + ', "dbt_version": "' + dbt_version + '"' if project_type != c.NATIVE else "" + ) collector._dispatcher.add_event.assert_has_calls( # type: ignore [ call( @@ -65,10 +69,16 @@ def test_on_project_loaded(collector: AnalyticsCollector, mocker: MockerFixture, def test_on_command(collector: AnalyticsCollector, mocker: MockerFixture): - collector.on_python_api_command(command_name="test_python_api", command_args=["arg_1", "arg_2"]) - collector.on_magic_command(command_name="test_magic", command_args=["arg_1", "arg_2"]) + collector.on_python_api_command( + command_name="test_python_api", command_args=["arg_1", "arg_2"] + ) + collector.on_magic_command( + command_name="test_magic", command_args=["arg_1", "arg_2"] + ) collector.on_cli_command( - command_name="test_cli", command_args=["arg_1", "arg_2"], parent_command_names=[] + command_name="test_cli", + command_args=["arg_1", "arg_2"], + parent_command_names=[], ) collector.flush() @@ -155,7 +165,9 @@ def test_on_cicd_command(collector: AnalyticsCollector, mocker: MockerFixture): @pytest.mark.slow def test_on_plan_apply( - collector: AnalyticsCollector, mocker: MockerFixture, init_and_plan_context: t.Callable + collector: AnalyticsCollector, + mocker: MockerFixture, + init_and_plan_context: t.Callable, ): context, plan = init_and_plan_context("examples/sushi") @@ -209,7 +221,9 @@ def test_on_plan_apply( @pytest.mark.slow def test_on_snapshots_created( - collector: AnalyticsCollector, mocker: MockerFixture, init_and_plan_context: t.Callable + collector: AnalyticsCollector, + mocker: MockerFixture, + init_and_plan_context: t.Callable, ): context, _ = init_and_plan_context("examples/sushi") @@ -286,7 +300,10 @@ def test_on_run(collector: AnalyticsCollector, mocker: MockerFixture): run_id = collector.on_run_start(engine_type="bigquery", state_sync_type="mysql") collector.on_run_end(run_id=run_id, succeeded=True, interrupted=False) collector.on_run_end( - run_id=run_id, succeeded=False, interrupted=False, error=SQLMeshError("test_error") + run_id=run_id, + succeeded=False, + interrupted=False, + error=SQLMeshError("test_error"), ) collector.on_run_end(run_id=run_id, succeeded=False, interrupted=True) diff --git a/tests/core/analytics/test_dispatcher.py b/tests/core/analytics/test_dispatcher.py index 14f55084b8..65f0701a07 100644 --- a/tests/core/analytics/test_dispatcher.py +++ b/tests/core/analytics/test_dispatcher.py @@ -5,7 +5,8 @@ import pytest from pytest_mock.plugin import MockerFixture -from sqlmesh.core.analytics.dispatcher import AsyncEventDispatcher, EventEmitter +from sqlmesh.core.analytics.dispatcher import (AsyncEventDispatcher, + EventEmitter) from sqlmesh.utils.errors import ApiClientError, SQLMeshError diff --git a/tests/core/engine_adapter/integration/__init__.py b/tests/core/engine_adapter/integration/__init__.py index 11bf95f3d6..7da3ce1bae 100644 --- a/tests/core/engine_adapter/integration/__init__.py +++ b/tests/core/engine_adapter/integration/__init__.py @@ -3,33 +3,35 @@ import os import pathlib import sys -import typing as t import time +import typing as t from contextlib import contextmanager +from dataclasses import dataclass import pandas as pd # noqa: TID253 import pytest +from _pytest.mark import MarkDecorator +from _pytest.mark.structures import ParameterSet from sqlglot import exp, parse_one from sqlglot.optimizer.normalize_identifiers import normalize_identifiers +import sqlmesh.core.dialect as d from sqlmesh import Config, Context, EngineAdapter from sqlmesh.core.config import load_config_from_paths from sqlmesh.core.config.connection import AthenaConnectionConfig from sqlmesh.core.dialect import normalize_model_name -import sqlmesh.core.dialect as d -from sqlmesh.core.engine_adapter import SparkEngineAdapter, TrinoEngineAdapter, AthenaEngineAdapter +from sqlmesh.core.engine_adapter import (AthenaEngineAdapter, + SparkEngineAdapter, + TrinoEngineAdapter) from sqlmesh.core.engine_adapter.shared import DataObject from sqlmesh.core.model.definition import SqlModel, load_sql_based_model from sqlmesh.utils import random_id from sqlmesh.utils.date import to_ds from sqlmesh.utils.pydantic import PydanticModel from tests.utils.pandas import compare_dataframes -from dataclasses import dataclass -from _pytest.mark import MarkDecorator -from _pytest.mark.structures import ParameterSet if t.TYPE_CHECKING: - from sqlmesh.core._typing import TableName, SchemaName + from sqlmesh.core._typing import SchemaName, TableName from sqlmesh.core.engine_adapter._typing import Query TEST_SCHEMA = "test_schema" @@ -73,7 +75,9 @@ def pytest_marks(self) -> t.List[MarkDecorator]: IntegrationTestEngine("postgres"), IntegrationTestEngine("mysql"), IntegrationTestEngine("mssql"), - IntegrationTestEngine("trino", catalog_types=["hive", "iceberg", "delta", "nessie"]), + IntegrationTestEngine( + "trino", catalog_types=["hive", "iceberg", "delta", "nessie"] + ), IntegrationTestEngine("spark", native_dataframe_type="pyspark"), IntegrationTestEngine("clickhouse", catalog_types=["standalone", "cluster"]), IntegrationTestEngine("risingwave"), @@ -127,7 +131,9 @@ def generate_pytest_params( catalogs = engine.catalog_types if engine.catalog_types else [""] for catalog in catalogs: gateway = ( - f"inttest_{engine.engine}_{catalog}" if catalog else f"inttest_{engine.engine}" + f"inttest_{engine.engine}_{catalog}" + if catalog + else f"inttest_{engine.engine}" ) if engine.engine == "athena": # athena only has a single gateway defined, not a gateway per catalog @@ -184,7 +190,11 @@ def from_data_objects(cls, data_objects: t.List[DataObject]) -> MetadataResults: @property def non_temp_tables(self) -> t.List[str]: - return [x for x in self.tables if not x.startswith("__temp") and not x.startswith("temp")] + return [ + x + for x in self.tables + if not x.startswith("__temp") and not x.startswith("temp") + ] class TestContext: @@ -208,12 +218,12 @@ def __init__( self.test_id = random_id(short=True) self._context: t.Optional[Context] = None self.is_remote = is_remote - self._schemas: t.List[ - str - ] = [] # keep track of any schemas returned from self.schema() / self.table() so we can drop them at the end - self._catalogs: t.List[ - str - ] = [] # keep track of any catalogs created via self.create_catalog() so we can drop them at the end + self._schemas: t.List[str] = ( + [] + ) # keep track of any schemas returned from self.schema() / self.table() so we can drop them at the end + self._catalogs: t.List[str] = ( + [] + ) # keep track of any catalogs created via self.create_catalog() so we can drop them at the end self.tmp_path = tmp_path @property @@ -300,7 +310,10 @@ def supports_merge(self) -> bool: assert isinstance(self.engine_adapter, TrinoEngineAdapter) # Trino supports MERGE on Delta and Iceberg but not Hive return ( - self.engine_adapter.get_catalog_type(self.engine_adapter.default_catalog) != "hive" + self.engine_adapter.get_catalog_type( + self.engine_adapter.default_catalog + ) + != "hive" ) if self.dialect == "athena": @@ -317,15 +330,21 @@ def supports_merge(self) -> bool: @property def default_table_format(self) -> t.Optional[str]: if self.dialect in {"athena", "trino"} and "_" in self.mark: - return self.mark.split("_", 1)[-1] # take eg 'athena_iceberg' and return 'iceberg' + return self.mark.split("_", 1)[ + -1 + ] # take eg 'athena_iceberg' and return 'iceberg' return None def add_test_suffix(self, value: str) -> str: return f"{value}_{self.test_id}" - def get_metadata_results(self, schema: t.Optional[SchemaName] = None) -> MetadataResults: + def get_metadata_results( + self, schema: t.Optional[SchemaName] = None + ) -> MetadataResults: schema = schema if schema else self.schema(TEST_SCHEMA) - return MetadataResults.from_data_objects(self.engine_adapter.get_data_objects(schema)) + return MetadataResults.from_data_objects( + self.engine_adapter.get_data_objects(schema) + ) def _init_engine_adapter(self) -> None: schema = self.schema(TEST_SCHEMA) @@ -342,7 +361,9 @@ def _format_df(self, data: pd.DataFrame, to_datetime: bool = True) -> pd.DataFra return data def init(self): - if self.df_type == "pyspark" and not hasattr(self.engine_adapter, "is_pyspark_df"): + if self.df_type == "pyspark" and not hasattr( + self.engine_adapter, "is_pyspark_df" + ): pytest.skip(f"Engine adapter {self.engine_adapter} doesn't support pyspark") self._init_engine_adapter() @@ -396,16 +417,24 @@ def physical_properties( self, properties_for_dialect: t.Dict[str, t.Dict[str, str | exp.Expr]] ) -> t.Dict[str, exp.Expr]: if props := properties_for_dialect.get(self.dialect): - return {k: exp.Literal.string(v) if isinstance(v, str) else v for k, v in props.items()} + return { + k: exp.Literal.string(v) if isinstance(v, str) else v + for k, v in props.items() + } return {} - def schema(self, schema_name: str = TEST_SCHEMA, catalog_name: t.Optional[str] = None) -> str: + def schema( + self, schema_name: str = TEST_SCHEMA, catalog_name: t.Optional[str] = None + ) -> str: schema_name = exp.table_name( normalize_model_name( self.add_test_suffix( ".".join( p - for p in (catalog_name or self.engine_adapter.default_catalog, schema_name) + for p in ( + catalog_name or self.engine_adapter.default_catalog, + schema_name, + ) if p ) if "." not in schema_name @@ -419,7 +448,9 @@ def schema(self, schema_name: str = TEST_SCHEMA, catalog_name: t.Optional[str] = return schema_name def get_current_data(self, table: exp.Table) -> pd.DataFrame: - df = self.engine_adapter.fetchdf(exp.select("*").from_(table), quote_identifiers=True) + df = self.engine_adapter.fetchdf( + exp.select("*").from_(table), quote_identifiers=True + ) if self.dialect == "snowflake" and "id" in df.columns: df["id"] = df["id"].apply(lambda x: x if pd.isna(x) else int(x)) return self._format_df(df) @@ -722,11 +753,15 @@ def create_catalog(self, catalog_name: str): # Use the engine adapter's built-in catalog creation functionality self.engine_adapter.create_catalog(catalog_name) elif self.dialect == "snowflake": - self.engine_adapter.execute(f'CREATE DATABASE IF NOT EXISTS "{catalog_name}"') + self.engine_adapter.execute( + f'CREATE DATABASE IF NOT EXISTS "{catalog_name}"' + ) elif self.dialect == "duckdb": try: # Only applies to MotherDuck - self.engine_adapter.execute(f'CREATE DATABASE IF NOT EXISTS "{catalog_name}"') + self.engine_adapter.execute( + f'CREATE DATABASE IF NOT EXISTS "{catalog_name}"' + ) except Exception: pass @@ -736,7 +771,9 @@ def drop_catalog(self, catalog_name: str): if self.dialect == "bigquery": return # bigquery cannot create/drop catalogs if self.dialect == "databricks": - self.engine_adapter.execute(f"DROP CATALOG IF EXISTS {catalog_name} CASCADE") + self.engine_adapter.execute( + f"DROP CATALOG IF EXISTS {catalog_name} CASCADE" + ) elif self.dialect == "fabric": # Use the engine adapter's built-in catalog dropping functionality self.engine_adapter.drop_catalog(catalog_name) @@ -796,14 +833,20 @@ def _get_create_user_or_role( # - sqlmesh-test-user@{project-id}.iam.gserviceaccount.com # - sqlmesh-test-writer@{project-id}.iam.gserviceaccount.com role_name = ( - username.replace(f"_{self.test_id}", "").replace("test_", "").replace("_", "-") + username.replace(f"_{self.test_id}", "") + .replace("test_", "") + .replace("_", "-") ) project_id = self.engine_adapter.get_current_catalog() - service_account = f"sqlmesh-test-{role_name}@{project_id}.iam.gserviceaccount.com" + service_account = ( + f"sqlmesh-test-{role_name}@{project_id}.iam.gserviceaccount.com" + ) return f"serviceAccount:{service_account}", None raise ValueError(f"User creation not supported for dialect: {self.dialect}") - def _create_user_or_role(self, username: str, password: t.Optional[str] = None) -> str: + def _create_user_or_role( + self, username: str, password: t.Optional[str] = None + ) -> str: username, create_user_sql = self._get_create_user_or_role(username, password) if create_user_sql: self.engine_adapter.execute(create_user_sql) @@ -821,9 +864,7 @@ def create_users_or_roles(self, *role_names: str) -> t.Iterator[t.Dict[str, str] ).sql(dialect=self.dialect) password = random_id() if self.dialect == "redshift": - password += ( - "A" # redshift requires passwords to have at least one uppercase letter - ) + password += "A" # redshift requires passwords to have at least one uppercase letter user_name = self._create_user_or_role(user_name, password) created_users.append(user_name) roles[role_name] = user_name diff --git a/tests/core/engine_adapter/integration/conftest.py b/tests/core/engine_adapter/integration/conftest.py index 3fb4bc15f1..058fb5059c 100644 --- a/tests/core/engine_adapter/integration/conftest.py +++ b/tests/core/engine_adapter/integration/conftest.py @@ -1,28 +1,24 @@ from __future__ import annotations +import logging +import os +import pathlib import typing as t + import pytest -import pathlib -import os -import logging from pytest import FixtureRequest from sqlmesh import Config, EngineAdapter +from sqlmesh.core.config import load_config_from_paths +from sqlmesh.core.config.connection import (AthenaConnectionConfig, + ConnectionConfig, + DuckDBConnectionConfig) from sqlmesh.core.constants import SQLMESH_PATH -from sqlmesh.core.config.connection import ( - ConnectionConfig, - AthenaConnectionConfig, - DuckDBConnectionConfig, -) from sqlmesh.core.engine_adapter import AthenaEngineAdapter -from sqlmesh.core.config import load_config_from_paths - -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES, - IntegrationTestEngine, -) +from tests.core.engine_adapter.integration import (ENGINES, + IntegrationTestEngine, + TestContext, + generate_pytest_params) logger = logging.getLogger(__name__) @@ -78,7 +74,9 @@ def _create(engine_name: str, gateway: str) -> EngineAdapter: if engine_name == "duckdb": assert isinstance(connection_config, DuckDBConnectionConfig) for raw_path in [ - v for v in (connection_config.catalogs or {}).values() if isinstance(v, str) + v + for v in (connection_config.catalogs or {}).values() + if isinstance(v, str) ]: pathlib.Path(raw_path).unlink(missing_ok=True) @@ -131,11 +129,15 @@ def _create( @pytest.fixture( - params=list(generate_pytest_params(ENGINES, query=True, show_variant_in_test_id=False)) + params=list( + generate_pytest_params(ENGINES, query=True, show_variant_in_test_id=False) + ) ) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -143,7 +145,9 @@ def ctx( @pytest.fixture(params=list(generate_pytest_params(ENGINES, query=False, df=True))) def ctx_df( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -151,6 +155,8 @@ def ctx_df( @pytest.fixture(params=list(generate_pytest_params(ENGINES, query=True, df=True))) def ctx_query_and_df( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) diff --git a/tests/core/engine_adapter/integration/test_freshness.py b/tests/core/engine_adapter/integration/test_freshness.py index e5ee574e7e..d3faf7e05a 100644 --- a/tests/core/engine_adapter/integration/test_freshness.py +++ b/tests/core/engine_adapter/integration/test_freshness.py @@ -4,22 +4,17 @@ import pathlib import typing as t from datetime import datetime, timedelta -from IPython.utils.capture import capture_output - -import time_machine -from pytest_mock.plugin import MockerFixture import pytest import time_machine +from IPython.utils.capture import capture_output +from pytest_mock.plugin import MockerFixture import sqlmesh from sqlmesh import Config, Context from sqlmesh.utils.date import now, to_datetime from sqlmesh.utils.errors import SignalEvalError -from tests.core.engine_adapter.integration import ( - TestContext, - TEST_SCHEMA, -) +from tests.core.engine_adapter.integration import TEST_SCHEMA, TestContext from tests.utils.test_helpers import use_terminal_console EVALUATION_SPY = None @@ -39,7 +34,9 @@ def _skip_snowflake(ctx: TestContext): @pytest.fixture(autouse=True, scope="function") def _install_evaluation_spy(mocker: MockerFixture): global EVALUATION_SPY - EVALUATION_SPY = mocker.spy(sqlmesh.core.snapshot.evaluator.SnapshotEvaluator, "evaluate") + EVALUATION_SPY = mocker.spy( + sqlmesh.core.snapshot.evaluator.SnapshotEvaluator, "evaluate" + ) yield EVALUATION_SPY = None @@ -58,9 +55,9 @@ def assert_snapshot_last_altered_ts( if snapshot.is_external: return - assert to_datetime(snapshot.last_altered_ts).replace(microsecond=0) == last_altered_ts.replace( + assert to_datetime(snapshot.last_altered_ts).replace( microsecond=0 - ) + ) == last_altered_ts.replace(microsecond=0) if dev_last_altered_ts: assert to_datetime(snapshot.dev_last_altered_ts).replace( @@ -69,7 +66,10 @@ def assert_snapshot_last_altered_ts( def assert_model_evaluation( - lambda_func, was_evaluated: bool = True, day_delta: int = 0, model_evaluations: int = 1 + lambda_func, + was_evaluated: bool = True, + day_delta: int = 0, + model_evaluations: int = 1, ): """ Ensure that a model was evaluated by checking the freshness signal and that @@ -104,8 +104,7 @@ def create_model( model_name = f"{schema}.{name}" model_path = path / "models" / f"{name}.sql" (path / "models").mkdir(parents=True, exist_ok=True) - model_path.write_text( - f""" + model_path.write_text(f""" MODEL ( name {model_name}, start '2024-01-01', @@ -116,8 +115,7 @@ def create_model( ); {query} - """ - ) + """) return model_name, model_path @@ -130,7 +128,9 @@ def initialize_context( """ adapter = ctx.engine_adapter if not adapter.SUPPORTS_METADATA_TABLE_LAST_MODIFIED_TS: - pytest.skip("This test only runs for engines that support metadata-based freshness") + pytest.skip( + "This test only runs for engines that support metadata-based freshness" + ) # Create & initialize schema schema = ctx.add_test_suffix(TEST_SCHEMA) @@ -149,15 +149,12 @@ def initialize_context( quote_identifiers=False, ) - yaml_content = ( - yaml_content - + f""" + yaml_content = yaml_content + f""" - name: {external_table} columns: col{i}: int """ - ) external_models_yaml = tmp_path / "external_models.yaml" external_models_yaml.write_text(yaml_content) @@ -172,7 +169,9 @@ def _set_config(gateway: str, config: Config) -> None: @use_terminal_console -def test_external_model_freshness(ctx: TestContext, tmp_path: pathlib.Path, mocker: MockerFixture): +def test_external_model_freshness( + ctx: TestContext, tmp_path: pathlib.Path, mocker: MockerFixture +): adapter = ctx.engine_adapter context, schema, (external_table1, external_table2) = initialize_context( ctx, tmp_path, num_external_models=2 @@ -194,7 +193,9 @@ def test_external_model_freshness(ctx: TestContext, tmp_path: pathlib.Path, mock ) prod_snapshot_id = next(iter(prod_plan_1.context_diff.new_snapshots)) - assert_snapshot_last_altered_ts(context, prod_snapshot_id, last_altered_ts=prod_plan_ts_1) + assert_snapshot_last_altered_ts( + context, prod_snapshot_id, last_altered_ts=prod_plan_ts_1 + ) # Case 2: Model is NOT evaluated on run if external models are not fresh assert_model_evaluation(lambda: context.run(), was_evaluated=False, day_delta=1) @@ -218,14 +219,20 @@ def test_external_model_freshness(ctx: TestContext, tmp_path: pathlib.Path, mock last_altered_ts=prod_plan_ts_1, dev_last_altered_ts=dev_plan_ts, ) - assert_snapshot_last_altered_ts(context, prod_snapshot_id, last_altered_ts=prod_plan_ts_1) + assert_snapshot_last_altered_ts( + context, prod_snapshot_id, last_altered_ts=prod_plan_ts_1 + ) # Case 4: Model is evaluated on run if any external model is fresh - adapter.execute(f"INSERT INTO {external_table2} (col2) VALUES (3)", quote_identifiers=False) + adapter.execute( + f"INSERT INTO {external_table2} (col2) VALUES (3)", quote_identifiers=False + ) assert_model_evaluation(lambda: context.run(), day_delta=2) # Case 5: Model is evaluated if changed (case 3) even if the external model is not fresh - model_path.write_text(model_path.read_text().replace("col1 + col2", "col1 * col2 * 5")) + model_path.write_text( + model_path.read_text().replace("col1 + col2", "col1 * col2 * 5") + ) context.load() assert_model_evaluation( lambda: context.plan(auto_apply=True, no_prompts=True), @@ -234,7 +241,9 @@ def test_external_model_freshness(ctx: TestContext, tmp_path: pathlib.Path, mock # Case 6: Model is evaluated on a restatement plan even if the external model is not fresh assert_model_evaluation( - lambda: context.plan(restate_models=[model_name], auto_apply=True, no_prompts=True), + lambda: context.plan( + restate_models=[model_name], auto_apply=True, no_prompts=True + ), day_delta=4, ) @@ -246,7 +255,9 @@ def test_mixed_model_freshness(ctx: TestContext, tmp_path: pathlib.Path): """ adapter = ctx.engine_adapter - context, schema, (external_table,) = initialize_context(ctx, tmp_path, num_external_models=1) + context, schema, (external_table,) = initialize_context( + ctx, tmp_path, num_external_models=1 + ) # Create parent model that depends on the external model parent_model_name, _ = create_model( @@ -289,10 +300,14 @@ def test_mixed_model_freshness(ctx: TestContext, tmp_path: pathlib.Path): ) for new_snapshot in prod_plan_1.context_diff.new_snapshots: - assert_snapshot_last_altered_ts(context, new_snapshot, last_altered_ts=prod_plan_ts_1) + assert_snapshot_last_altered_ts( + context, new_snapshot, last_altered_ts=prod_plan_ts_1 + ) # Case 2: Mixed models are evaluated if the upstream models (sqlmesh or external) become fresh - adapter.execute(f"INSERT INTO {external_table} (col1) VALUES (2)", quote_identifiers=False) + adapter.execute( + f"INSERT INTO {external_table} (col1) VALUES (2)", quote_identifiers=False + ) assert_model_evaluation( lambda: context.run(), was_evaluated=True, day_delta=1, model_evaluations=3 @@ -317,7 +332,9 @@ def test_mixed_model_freshness(ctx: TestContext, tmp_path: pathlib.Path): assert prod_plan_2.context_diff.modified_snapshots assert_snapshot_last_altered_ts( - context, next(iter(prod_plan_2.context_diff.new_snapshots)), last_altered_ts=prod_plan_ts_2 + context, + next(iter(prod_plan_2.context_diff.new_snapshots)), + last_altered_ts=prod_plan_ts_2, ) diff --git a/tests/core/engine_adapter/integration/test_integration.py b/tests/core/engine_adapter/integration/test_integration.py index 44f680dafb..b1b870f900 100644 --- a/tests/core/engine_adapter/integration/test_integration.py +++ b/tests/core/engine_adapter/integration/test_integration.py @@ -1,18 +1,15 @@ # type: ignore from __future__ import annotations +import logging import pathlib import re +import shutil import sys import typing as t -import shutil -from datetime import datetime, timedelta, date +from datetime import date, datetime, timedelta from unittest import mock from unittest.mock import patch -import logging - - -import time_machine import numpy as np # noqa: TID253 import pandas as pd # noqa: TID253 @@ -23,31 +20,28 @@ from sqlglot.optimizer.normalize_identifiers import normalize_identifiers from sqlglot.optimizer.qualify_columns import quote_identifiers +import sqlmesh.core.dialect as d from sqlmesh import Config, Context from sqlmesh.cli.project_init import init_example_project from sqlmesh.core.config.common import VirtualEnvironmentMode from sqlmesh.core.config.connection import ConnectionConfig -import sqlmesh.core.dialect as d -from sqlmesh.core.environment import EnvironmentSuffixTarget from sqlmesh.core.dialect import select_from_values -from sqlmesh.core.model import Model, load_sql_based_model +from sqlmesh.core.engine_adapter.mixins import LogicalMergeMixin, RowDiffMixin from sqlmesh.core.engine_adapter.shared import DataObject, DataObjectType -from sqlmesh.core.engine_adapter.mixins import RowDiffMixin, LogicalMergeMixin +from sqlmesh.core.environment import EnvironmentSuffixTarget +from sqlmesh.core.model import Model, load_sql_based_model from sqlmesh.core.model.definition import create_sql_model from sqlmesh.core.plan import Plan -from sqlmesh.core.state_sync.db import EngineAdapterStateSync from sqlmesh.core.snapshot import Snapshot, SnapshotChangeCategory -from sqlmesh.utils.date import now, to_date, to_time_column +from sqlmesh.core.state_sync.db import EngineAdapterStateSync from sqlmesh.core.table_diff import TableDiff +from sqlmesh.utils.date import now, to_date, to_time_column from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.pydantic import PydanticModel from tests.conftest import SushiDataValidator -from tests.core.engine_adapter.integration import ( - TestContext, - MetadataResults, - TEST_SCHEMA, - wait_until, -) +from tests.core.engine_adapter.integration import (TEST_SCHEMA, + MetadataResults, + TestContext, wait_until) DATA_TYPE = exp.DataType.Type VARCHAR_100 = exp.DataType.build("varchar(100)") @@ -71,10 +65,18 @@ def create(cls, plan: Plan, ctx: TestContext, schema_name: str): ) def snapshot_for(self, model: Model) -> Snapshot: - return next((s for s in list(self.plan.snapshots.values()) if s.name == model.fqn)) + return next( + (s for s in list(self.plan.snapshots.values()) if s.name == model.fqn) + ) def modified_snapshot_for(self, model: Model) -> Snapshot: - return next((s for s in list(self.plan.modified_snapshots.values()) if s.name == model.fqn)) + return next( + ( + s + for s in list(self.plan.modified_snapshots.values()) + if s.name == model.fqn + ) + ) def table_name_for( self, snapshot_or_model: Snapshot | Model, is_deployable: bool = True @@ -136,7 +138,9 @@ def drop_schema_and_validate(schema_name: str): def create_objects_and_validate(schema_name: str): ctx.engine_adapter.create_schema(schema_name) - ctx.engine_adapter.create_view(f"{schema_name}.test_view", parse_one("SELECT 1 as col")) + ctx.engine_adapter.create_view( + f"{schema_name}.test_view", parse_one("SELECT 1 as col") + ) ctx.engine_adapter.create_table( f"{schema_name}.test_table", {"col": exp.DataType.build("int")} ) @@ -167,14 +171,17 @@ def create_objects_and_validate(schema_name: str): if ctx.dialect == "bigquery": catalog_name = ctx.engine_adapter.get_current_catalog() - catalog_name = normalize_identifiers(catalog_name, dialect=ctx.dialect).sql(dialect=ctx.dialect) + catalog_name = normalize_identifiers(catalog_name, dialect=ctx.dialect).sql( + dialect=ctx.dialect + ) ctx.create_catalog(catalog_name) schema = ctx.schema("drop_schema_catalog_test", catalog_name) if ctx.engine_adapter.catalog_support.is_single_catalog_only: with pytest.raises( - SQLMeshError, match="requires that all catalog operations be against a single catalog" + SQLMeshError, + match="requires that all catalog operations be against a single catalog", ): drop_schema_and_validate(schema) create_objects_and_validate(schema) @@ -208,7 +215,9 @@ def test_temp_table(ctx_query_and_df: TestContext): ctx.compare_with_current(table_name, input_data) results = ctx.get_metadata_results() - assert len(results.views) == len(results.tables) == len(results.non_temp_tables) == 0 + assert ( + len(results.views) == len(results.tables) == len(results.non_temp_tables) == 0 + ) def test_create_table(ctx: TestContext): @@ -342,8 +351,12 @@ def test_create_view(ctx_query_and_df: TestContext): ctx.compare_with_current(view, input_data) if ctx.engine_adapter.COMMENT_CREATION_VIEW.is_supported: - table_description = ctx.get_table_comment(view.db, "test_view", table_kind="VIEW") - column_comments = ctx.get_column_comments(view.db, "test_view", table_kind="VIEW") + table_description = ctx.get_table_comment( + view.db, "test_view", table_kind="VIEW" + ) + column_comments = ctx.get_column_comments( + view.db, "test_view", table_kind="VIEW" + ) # Query: # In the query test, columns_to_types are not available when the view is created. Since we @@ -396,8 +409,12 @@ def test_create_view_source_columns(ctx_query_and_df: TestContext): ctx.compare_with_current(view, expected_data) if ctx.engine_adapter.COMMENT_CREATION_VIEW.is_supported: - table_description = ctx.get_table_comment(view.db, "test_view", table_kind="VIEW") - column_comments = ctx.get_column_comments(view.db, "test_view", table_kind="VIEW") + table_description = ctx.get_table_comment( + view.db, "test_view", table_kind="VIEW" + ) + column_comments = ctx.get_column_comments( + view.db, "test_view", table_kind="VIEW" + ) assert table_description == "test view description" assert column_comments == {"id": "test id column description"} @@ -406,7 +423,9 @@ def test_create_view_source_columns(ctx_query_and_df: TestContext): def test_materialized_view(ctx_query_and_df: TestContext): ctx = ctx_query_and_df if not ctx.engine_adapter.SUPPORTS_MATERIALIZED_VIEWS: - pytest.skip(f"Engine adapter {ctx.engine_adapter} doesn't support materialized views") + pytest.skip( + f"Engine adapter {ctx.engine_adapter} doesn't support materialized views" + ) if ctx.engine_adapter.dialect == "databricks": pytest.skip( "Databricks requires DBSQL Serverless or Pro warehouse to test materialized views which we do not have setup" @@ -427,7 +446,9 @@ def test_materialized_view(ctx_query_and_df: TestContext): ] ) source_table = ctx.table("source_table") - ctx.engine_adapter.ctas(source_table, ctx.input_data(input_data), ctx.columns_to_types) + ctx.engine_adapter.ctas( + source_table, ctx.input_data(input_data), ctx.columns_to_types + ) view = ctx.table("test_view") view_query = exp.select(*ctx.columns_to_types).from_(source_table) ctx.engine_adapter.create_view(view, view_query, materialized=True) @@ -518,9 +539,9 @@ def test_replace_query(ctx_query_and_df: TestContext): # provided then it checks the table itself for types. This is fine within SQLMesh since we always know the tables # exist prior to evaluation but when running these tests that isn't the case. As a result we just pass in # columns_to_types for these two engines so we can still test inference on the other ones - target_columns_to_types=ctx.columns_to_types - if ctx.dialect in ["spark", "databricks"] - else None, + target_columns_to_types=( + ctx.columns_to_types if ctx.dialect in ["spark", "databricks"] else None + ), table_format=ctx.default_table_format, ) results = ctx.get_metadata_results() @@ -571,7 +592,9 @@ def test_replace_query_source_columns(ctx_query_and_df: TestContext): {"id": 3, "ds": "2022-01-03", "ignored_source": "ignored_value"}, ] ) - ctx.engine_adapter.create_table(table, columns_to_types, table_format=ctx.default_table_format) + ctx.engine_adapter.create_table( + table, columns_to_types, table_format=ctx.default_table_format + ) ctx.engine_adapter.replace_query( table, ctx.input_data(input_data), @@ -639,9 +662,9 @@ def test_replace_query_batched(ctx_query_and_df: TestContext): # provided then it checks the table itself for types. This is fine within SQLMesh since we always know the tables # exist prior to evaluation but when running these tests that isn't the case. As a result we just pass in # columns_to_types for these two engines so we can still test inference on the other ones - target_columns_to_types=ctx.columns_to_types - if ctx.dialect in ["spark", "databricks"] - else None, + target_columns_to_types=( + ctx.columns_to_types if ctx.dialect in ["spark", "databricks"] else None + ), table_format=ctx.default_table_format, ) results = ctx.get_metadata_results() @@ -722,7 +745,9 @@ def test_insert_append_source_columns(ctx_query_and_df: TestContext): table = ctx.table("test_table") columns_to_types = ctx.columns_to_types.copy() columns_to_types["ignored_column"] = exp.DataType.build("int") - ctx.engine_adapter.create_table(table, columns_to_types, table_format=ctx.default_table_format) + ctx.engine_adapter.create_table( + table, columns_to_types, table_format=ctx.default_table_format + ) # Initial Load input_data = pd.DataFrame( [ @@ -773,7 +798,9 @@ def test_insert_append_source_columns(ctx_query_and_df: TestContext): assert len(results.tables) in [1, 2, 3] assert len(results.non_temp_tables) == 1 assert results.non_temp_tables[0] == table.name - ctx.compare_with_current(table, pd.concat([expected_data, append_expected_data])) + ctx.compare_with_current( + table, pd.concat([expected_data, append_expected_data]) + ) def test_insert_overwrite_by_time_partition(ctx_query_and_df: TestContext): @@ -866,7 +893,9 @@ def test_insert_overwrite_by_time_partition(ctx_query_and_df: TestContext): ) -def test_insert_overwrite_by_time_partition_source_columns(ctx_query_and_df: TestContext): +def test_insert_overwrite_by_time_partition_source_columns( + ctx_query_and_df: TestContext, +): ctx = ctx_query_and_df ds_type = "string" if ctx.dialect == "bigquery": @@ -931,9 +960,21 @@ def test_insert_overwrite_by_time_partition_source_columns(ctx_query_and_df: Tes if ctx.test_type == "df": overwrite_data = pd.DataFrame( [ - {"id": 10, ctx.time_column: "2022-01-03", "ignored_source": "ignored_value"}, - {"id": 4, ctx.time_column: "2022-01-04", "ignored_source": "ignored_value"}, - {"id": 5, ctx.time_column: "2022-01-05", "ignored_source": "ignored_value"}, + { + "id": 10, + ctx.time_column: "2022-01-03", + "ignored_source": "ignored_value", + }, + { + "id": 4, + ctx.time_column: "2022-01-04", + "ignored_source": "ignored_value", + }, + { + "id": 5, + ctx.time_column: "2022-01-05", + "ignored_source": "ignored_value", + }, ] ) ctx.engine_adapter.insert_overwrite_by_time_partition( @@ -979,7 +1020,9 @@ def test_merge(ctx_query_and_df: TestContext): # And it cant fall back to a logical merge on Hive tables because it cant delete records table_format = "iceberg" if ctx.dialect == "athena" else None - ctx.engine_adapter.create_table(table, ctx.columns_to_types, table_format=table_format) + ctx.engine_adapter.create_table( + table, ctx.columns_to_types, table_format=table_format + ) input_data = pd.DataFrame( [ {"id": 1, "ds": "2022-01-01"}, @@ -1128,7 +1171,9 @@ def test_scd_type_2_by_time(ctx_query_and_df: TestContext): } table = ctx.table("test_table") input_schema = { - k: v for k, v in ctx.columns_to_types.items() if k not in ("valid_from", "valid_to") + k: v + for k, v in ctx.columns_to_types.items() + if k not in ("valid_from", "valid_to") } ctx.engine_adapter.create_table( @@ -1286,10 +1331,14 @@ def test_scd_type_2_by_time_source_columns(ctx_query_and_df: TestContext): table = ctx.table("test_table") input_schema = { - k: v for k, v in ctx.columns_to_types.items() if k not in ("valid_from", "valid_to") + k: v + for k, v in ctx.columns_to_types.items() + if k not in ("valid_from", "valid_to") } - ctx.engine_adapter.create_table(table, columns_to_types, table_format=ctx.default_table_format) + ctx.engine_adapter.create_table( + table, columns_to_types, table_format=ctx.default_table_format + ) input_data = pd.DataFrame( [ { @@ -1481,7 +1530,9 @@ def test_scd_type_2_by_column(ctx_query_and_df: TestContext): } table = ctx.table("test_table") input_schema = { - k: v for k, v in ctx.columns_to_types.items() if k not in ("valid_from", "valid_to") + k: v + for k, v in ctx.columns_to_types.items() + if k not in ("valid_from", "valid_to") } ctx.engine_adapter.create_table( @@ -1661,16 +1712,40 @@ def test_scd_type_2_by_column_source_columns(ctx_query_and_df: TestContext): table = ctx.table("test_table") input_schema = { - k: v for k, v in ctx.columns_to_types.items() if k not in ("valid_from", "valid_to") + k: v + for k, v in ctx.columns_to_types.items() + if k not in ("valid_from", "valid_to") } - ctx.engine_adapter.create_table(table, columns_to_types, table_format=ctx.default_table_format) + ctx.engine_adapter.create_table( + table, columns_to_types, table_format=ctx.default_table_format + ) input_data = pd.DataFrame( [ - {"id": 1, "name": "a", "status": "active", "ignored_source": "ignored_value"}, - {"id": 2, "name": "b", "status": "inactive", "ignored_source": "ignored_value"}, - {"id": 3, "name": "c", "status": "active", "ignored_source": "ignored_value"}, - {"id": 4, "name": "d", "status": "active", "ignored_source": "ignored_value"}, + { + "id": 1, + "name": "a", + "status": "active", + "ignored_source": "ignored_value", + }, + { + "id": 2, + "name": "b", + "status": "inactive", + "ignored_source": "ignored_value", + }, + { + "id": 3, + "name": "c", + "status": "active", + "ignored_source": "ignored_value", + }, + { + "id": 4, + "name": "d", + "status": "active", + "ignored_source": "ignored_value", + }, ] ) ctx.engine_adapter.scd_type_2_by_column( @@ -1739,15 +1814,35 @@ def test_scd_type_2_by_column_source_columns(ctx_query_and_df: TestContext): current_data = pd.DataFrame( [ # Change `a` to `x` - {"id": 1, "name": "x", "status": "active", "ignored_source": "ignored_value"}, + { + "id": 1, + "name": "x", + "status": "active", + "ignored_source": "ignored_value", + }, # Delete # {"id": 2, "name": "b", status: "inactive", "ignored_source": "ignored_value"}, # No change - {"id": 3, "name": "c", "status": "active", "ignored_source": "ignored_value"}, + { + "id": 3, + "name": "c", + "status": "active", + "ignored_source": "ignored_value", + }, # Change status to inactive - {"id": 4, "name": "d", "status": "inactive", "ignored_source": "ignored_value"}, + { + "id": 4, + "name": "d", + "status": "inactive", + "ignored_source": "ignored_value", + }, # Add - {"id": 5, "name": "e", "status": "inactive", "ignored_source": "ignored_value"}, + { + "id": 5, + "name": "e", + "status": "inactive", + "ignored_source": "ignored_value", + }, ] ) ctx.engine_adapter.scd_type_2_by_column( @@ -1854,7 +1949,9 @@ def test_get_data_objects(ctx_query_and_df: TestContext): schema = ctx.schema(TEST_SCHEMA) - assert sorted(ctx.engine_adapter.get_data_objects(schema), key=lambda o: o.name) == [ + assert sorted( + ctx.engine_adapter.get_data_objects(schema), key=lambda o: o.name + ) == [ DataObject( name=table.name, schema=table.db, @@ -1930,7 +2027,9 @@ def test_truncate_table(ctx: TestContext): def test_transaction(ctx: TestContext): if ctx.engine_adapter.SUPPORTS_TRANSACTIONS is False: - pytest.skip(f"Engine adapter {ctx.engine_adapter.dialect} doesn't support transactions") + pytest.skip( + f"Engine adapter {ctx.engine_adapter.dialect} doesn't support transactions" + ) table = ctx.table("test_table") input_data = pd.DataFrame( @@ -1943,7 +2042,9 @@ def test_transaction(ctx: TestContext): with ctx.engine_adapter.transaction(): ctx.engine_adapter.create_table(table, ctx.columns_to_types) ctx.engine_adapter.insert_append( - table, ctx.input_data(input_data, ctx.columns_to_types), ctx.columns_to_types + table, + ctx.input_data(input_data, ctx.columns_to_types), + ctx.columns_to_types, ) ctx.compare_with_current(table, input_data) with ctx.engine_adapter.transaction(): @@ -1953,10 +2054,13 @@ def test_transaction(ctx: TestContext): @pytest.mark.parametrize( - "virtual_environment_mode", [VirtualEnvironmentMode.FULL, VirtualEnvironmentMode.DEV_ONLY] + "virtual_environment_mode", + [VirtualEnvironmentMode.FULL, VirtualEnvironmentMode.DEV_ONLY], ) def test_sushi( - ctx: TestContext, tmp_path: pathlib.Path, virtual_environment_mode: VirtualEnvironmentMode + ctx: TestContext, + tmp_path: pathlib.Path, + virtual_environment_mode: VirtualEnvironmentMode, ): if ctx.mark == "athena_hive": pytest.skip( @@ -2011,7 +2115,9 @@ def _mutate_config(gateway: str, config: Config) -> None: ] config.virtual_environment_mode = virtual_environment_mode - context = ctx.create_context(_mutate_config, path=tmp_path, ephemeral_state_connection=False) + context = ctx.create_context( + _mutate_config, path=tmp_path, ephemeral_state_connection=False + ) end = now() start = to_date(end - timedelta(days=7)) @@ -2021,9 +2127,9 @@ def _mutate_config(gateway: str, config: Config) -> None: # spaces in column names. Other engines error if it is set in the model definition, # so we set it here. if ctx.dialect == "databricks": - cust_rev_by_day_key = [key for key in context._models if "customer_revenue_by_day" in key][ - 0 - ] + cust_rev_by_day_key = [ + key for key in context._models if "customer_revenue_by_day" in key + ][0] cust_rev_by_day_model_tbl_props = context._models[cust_rev_by_day_key].copy( update={ @@ -2121,7 +2227,9 @@ def _mutate_config(gateway: str, config: Config) -> None: auto_apply=True, ) - data_validator = SushiDataValidator.from_context(context, sushi_schema_name=sushi_test_schema) + data_validator = SushiDataValidator.from_context( + context, sushi_schema_name=sushi_test_schema + ) data_validator.validate( f"{sushi_test_schema}.customer_revenue_lifetime", start, @@ -2199,7 +2307,9 @@ def validate_comments( if not model_name in layer_models: continue layer_table_name = layer_models[model_name]["table_name"] - table_kind = "VIEW" if layer_models[model_name]["is_view"] else "BASE TABLE" + table_kind = ( + "VIEW" if layer_models[model_name]["is_view"] else "BASE TABLE" + ) # is this model in a physical layer or PROD environment? is_physical_or_prod = is_physical_layer or ( @@ -2251,9 +2361,16 @@ def validate_comments( table_kind=table_kind, snowflake_capitalize_ids=False, ) - for column_name, expected_col_comment in expected_col_comments.items(): - expected_col_comment = expected_col_comments.get(column_name, None) - actual_col_comment = actual_col_comments.get(column_name, None) + for ( + column_name, + expected_col_comment, + ) in expected_col_comments.items(): + expected_col_comment = expected_col_comments.get( + column_name, None + ) + actual_col_comment = actual_col_comments.get( + column_name, None + ) assert expected_col_comment == actual_col_comment return None @@ -2276,11 +2393,15 @@ def validate_no_comments( if x.name.endswith(table_name_suffix) } if not check_temp_tables: - layer_models = {k: v for k, v in layer_models.items() if not k.endswith("__dev")} + layer_models = { + k: v for k, v in layer_models.items() if not k.endswith("__dev") + } for model_name, comment in comments.items(): layer_table_name = layer_models[model_name]["table_name"] - table_kind = "VIEW" if layer_models[model_name]["is_view"] else "BASE TABLE" + table_kind = ( + "VIEW" if layer_models[model_name]["is_view"] else "BASE TABLE" + ) actual_tbl_comment = ctx.get_table_comment( schema_name, @@ -2309,12 +2430,18 @@ def validate_no_comments( snowflake_capitalize_ids=False, ) for column_name in expected_col_comments: - actual_col_comment = actual_col_comments.get(column_name, None) - assert actual_col_comment is None or actual_col_comment == "" + actual_col_comment = actual_col_comments.get( + column_name, None + ) + assert ( + actual_col_comment is None or actual_col_comment == "" + ) return None - validate_comments(f"sqlmesh__{sushi_test_schema}", prod_schema_name=sushi_test_schema) + validate_comments( + f"sqlmesh__{sushi_test_schema}", prod_schema_name=sushi_test_schema + ) # confirm view layer comments are not registered in non-PROD environment env_name = "test_prod" @@ -2394,7 +2521,8 @@ def _normalize_snowflake(name: str, prefix_regex: str = "(sqlmesh__)(.*)"): return name.upper() object_names = { - k: [_normalize_snowflake(name) for name in v] for k, v in object_names.items() + k: [_normalize_snowflake(name) for name in v] + for k, v in object_names.items() } init_example_project(tmp_path, ctx.engine_type, schema_name=schema_name) @@ -2402,12 +2530,16 @@ def _normalize_snowflake(name: str, prefix_regex: str = "(sqlmesh__)(.*)"): def _mutate_config(gateway: str, config: Config): # ensure default dialect comes from init_example_project and not ~/.sqlmesh/config.yaml if config.model_defaults.dialect != ctx.dialect: - config.model_defaults = config.model_defaults.copy(update={"dialect": ctx.dialect}) + config.model_defaults = config.model_defaults.copy( + update={"dialect": ctx.dialect} + ) # Ensure the state schema is unique to this test (since we deliberately use the warehouse as the state connection) config.gateways[gateway].state_schema = state_schema - context = ctx.create_context(_mutate_config, path=tmp_path, ephemeral_state_connection=False) + context = ctx.create_context( + _mutate_config, path=tmp_path, ephemeral_state_connection=False + ) if ctx.default_table_format: # if the default table format is explicitly set, ensure its being used @@ -2434,9 +2566,9 @@ def capture_execution_stats( auto_restatement_triggers=None, ): if execution_stats is not None: - actual_execution_stats[snapshot.model.name.replace(f"{schema_name}.", "")] = ( - execution_stats - ) + actual_execution_stats[ + snapshot.model.name.replace(f"{schema_name}.", "") + ] = execution_stats # apply prod plan with patch.object( @@ -2447,18 +2579,28 @@ def capture_execution_stats( prod_schema_results = ctx.get_metadata_results(object_names["view_schema"][0]) assert sorted(prod_schema_results.views) == object_names["views"] assert len(prod_schema_results.materialized_views) == 0 - assert len(prod_schema_results.tables) == len(prod_schema_results.non_temp_tables) == 0 + assert ( + len(prod_schema_results.tables) == len(prod_schema_results.non_temp_tables) == 0 + ) - physical_layer_results = ctx.get_metadata_results(object_names["physical_schema"][0]) + physical_layer_results = ctx.get_metadata_results( + object_names["physical_schema"][0] + ) assert len(physical_layer_results.views) == 0 assert len(physical_layer_results.materialized_views) == 0 - assert len(physical_layer_results.tables) == len(physical_layer_results.non_temp_tables) == 3 + assert ( + len(physical_layer_results.tables) + == len(physical_layer_results.non_temp_tables) + == 3 + ) if ctx.engine_adapter.SUPPORTS_QUERY_EXECUTION_TRACKING: assert actual_execution_stats["incremental_model"].total_rows_processed == 7 # snowflake and redshift don't track rows for CTAS assert actual_execution_stats["full_model"].total_rows_processed == ( - None if ctx.mark.startswith("snowflake") or ctx.mark.startswith("redshift") else 3 + None + if ctx.mark.startswith("snowflake") or ctx.mark.startswith("redshift") + else 3 ) assert actual_execution_stats["seed_model"].total_rows_processed == ( None if ctx.mark.startswith("snowflake") else 7 @@ -2473,7 +2615,9 @@ def capture_execution_stats( if not ctx.is_remote: actual_execution_stats = {} with patch.object( - context.console, "update_snapshot_evaluation_progress", capture_execution_stats + context.console, + "update_snapshot_evaluation_progress", + capture_execution_stats, ): with time_machine.travel(date.today() + timedelta(days=1)): context.run() @@ -2501,7 +2645,9 @@ def capture_execution_stats( dev_schema_results = ctx.get_metadata_results(schema_name) assert sorted(dev_schema_results.views) == object_names["views"] assert len(dev_schema_results.materialized_views) == 0 - assert len(dev_schema_results.tables) == len(dev_schema_results.non_temp_tables) == 0 + assert ( + len(dev_schema_results.tables) == len(dev_schema_results.non_temp_tables) == 0 + ) # register the schemas to be cleaned up for schema in [ @@ -2541,8 +2687,7 @@ def test_dialects(ctx: TestContext): c = '"C"' d = '"D"' - q = parse_one( - f""" + q = parse_one(f""" WITH "a" AS (SELECT 1 w), "B" AS (SELECT 1 x), @@ -2554,10 +2699,11 @@ def test_dialects(ctx: TestContext): CROSS JOIN {b} CROSS JOIN {c} CROSS JOIN {d} - """ - ) + """) df = ctx.engine_adapter.fetchdf(q) - expected_columns = ["W", "X", "Y", "Z"] if ctx.dialect == "snowflake" else ["w", "x", "y", "z"] + expected_columns = ( + ["W", "X", "Y", "Z"] if ctx.dialect == "snowflake" else ["w", "x", "y", "z"] + ) pd.testing.assert_frame_equal( df, pd.DataFrame([[1, 1, 1, 1]], columns=expected_columns), check_dtype=False ) @@ -2633,7 +2779,9 @@ def test_to_time_column( ctx: TestContext, time_column, time_column_type, time_column_format, result ): # TODO: can this be cleaned up after recent sqlglot updates? - if ctx.dialect == "clickhouse" and time_column_type.is_type(exp.DataType.Type.TIMESTAMPTZ): + if ctx.dialect == "clickhouse" and time_column_type.is_type( + exp.DataType.Type.TIMESTAMPTZ + ): # Clickhouse does not have natively timezone-aware types and does not accept timestrings # with UTC offset "+XX:XX". Therefore, we remove the timezone offset and set a timezone- # specific data type to validate what is returned. @@ -2641,7 +2789,9 @@ def test_to_time_column( time_column = re.match(r"^(.*?)\+", time_column).group(1) time_column_type = exp.DataType.build("TIMESTAMP('UTC')", dialect="clickhouse") - time_column = to_time_column(time_column, time_column_type, ctx.dialect, time_column_format) + time_column = to_time_column( + time_column, time_column_type, ctx.dialect, time_column_format + ) df = ctx.engine_adapter.fetchdf(exp.select(time_column).as_("the_col")) expected = result.get(ctx.dialect, result.get("default")) col_name = "THE_COL" if ctx.dialect == "snowflake" else "the_col" @@ -2727,7 +2877,9 @@ def _mutate_config(current_gateway_name: str, config: Config): assert "test_model" in results.views actual_df = ( - ctx.get_current_data(test_model.fqn).sort_values(by="event_date").reset_index(drop=True) + ctx.get_current_data(test_model.fqn) + .sort_values(by="event_date") + .reset_index(drop=True) ) actual_df["event_date"] = actual_df["event_date"].astype(str) assert actual_df.count()[0] == 3 @@ -2772,8 +2924,18 @@ def _mutate_config(current_gateway_name: str, config: Config): [ [1, "item_a", 100, "2020-01-01"], [2, "item_b", 200, "2020-01-01"], - [1, "item_a_changed", 150, "2020-01-02"], # Same item_id, different name and value - [2, "item_b_changed", 250, "2020-01-02"], # Same item_id, different name and value + [ + 1, + "item_a_changed", + 150, + "2020-01-02", + ], # Same item_id, different name and value + [ + 2, + "item_b_changed", + 250, + "2020-01-02", + ], # Same item_id, different name and value [3, "item_c", 300, "2020-01-02"], # New item on day 2 ], columns=["item_id", "name", "value", "event_date"], @@ -2840,7 +3002,9 @@ def _mutate_config(current_gateway_name: str, config: Config): ) actual_df = ( - ctx.get_current_data(test_model.fqn).sort_values(by="item_id").reset_index(drop=True) + ctx.get_current_data(test_model.fqn) + .sort_values(by="item_id") + .reset_index(drop=True) ) # Expected results after batch processing: @@ -2850,8 +3014,18 @@ def _mutate_config(current_gateway_name: str, config: Config): expected_df = ( pd.DataFrame( [ - [1, "item_a", 150, "2020-01-02"], # name from day 1, value and date from day 2 - [2, "item_b", 250, "2020-01-02"], # name from day 1, value and date from day 2 + [ + 1, + "item_a", + 150, + "2020-01-02", + ], # name from day 1, value and date from day 2 + [ + 2, + "item_b", + 250, + "2020-01-02", + ], # name from day 1, value and date from day 2 [3, "item_c", 300, "2020-01-02"], # new item from day 2 ], columns=["item_id", "name", "value", "event_date"], @@ -2893,15 +3067,15 @@ def test_managed_model_upstream_forward_only(ctx: TestContext): pytest.skip("This test only runs for engines that support managed models") def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: - plan: Plan = sqlmesh_context.plan(auto_apply=True, no_prompts=True, environment=environment) + plan: Plan = sqlmesh_context.plan( + auto_apply=True, no_prompts=True, environment=environment + ) return PlanResults.create(plan, ctx, schema) context = ctx.create_context() schema = ctx.add_test_suffix(TEST_SCHEMA) - model_a = load_sql_based_model( - d.parse( # type: ignore - f""" + model_a = load_sql_based_model(d.parse(f""" MODEL ( name {schema}.upstream_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -2911,13 +3085,9 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: ); SELECT 1 as id, 'foo' as name, current_timestamp as ts; - """ - ) - ) + """)) # type: ignore - model_b = load_sql_based_model( - d.parse( # type: ignore - f""" + model_b = load_sql_based_model(d.parse(f""" MODEL ( name {schema}.managed_model, kind MANAGED, @@ -2927,18 +3097,20 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: ); SELECT * from {schema}.upstream_model; - """ - ) - ) + """)) # type: ignore context.upsert_model(model_a) context.upsert_model(model_b) plan_1 = _run_plan(context) - assert plan_1.snapshot_for(model_a).change_category == SnapshotChangeCategory.BREAKING + assert ( + plan_1.snapshot_for(model_a).change_category == SnapshotChangeCategory.BREAKING + ) assert not plan_1.snapshot_for(model_a).is_forward_only - assert plan_1.snapshot_for(model_b).change_category == SnapshotChangeCategory.BREAKING + assert ( + plan_1.snapshot_for(model_b).change_category == SnapshotChangeCategory.BREAKING + ) assert not plan_1.snapshot_for(model_b).is_forward_only # so far so good, model_a should exist as a normal table, model b should be a managed table and the prod views should exist @@ -2954,15 +3126,16 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: ) # because its a managed table assert len(plan_1.internal_schema_metadata.managed_tables) == 1 - assert plan_1.table_name_for(model_b) in plan_1.internal_schema_metadata.managed_tables assert ( - plan_1.dev_table_name_for(model_b) not in plan_1.internal_schema_metadata.managed_tables + plan_1.table_name_for(model_b) in plan_1.internal_schema_metadata.managed_tables + ) + assert ( + plan_1.dev_table_name_for(model_b) + not in plan_1.internal_schema_metadata.managed_tables ) # the dev table should not be created as managed # Let's modify model A with a breaking change and plan it against a dev environment. This should trigger a forward-only plan - new_model_a = load_sql_based_model( - d.parse( # type: ignore - f""" + new_model_a = load_sql_based_model(d.parse(f""" MODEL ( name {schema}.upstream_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -2972,9 +3145,7 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: ); SELECT 1 as id, 'foo' as name, 'bar' as extra, current_timestamp as ts; - """ - ) - ) + """)) # type: ignore context.upsert_model(new_model_a) # apply plan to dev environment @@ -2982,9 +3153,15 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: assert plan_2.plan.has_changes assert len(plan_2.plan.modified_snapshots) == 2 - assert plan_2.snapshot_for(new_model_a).change_category == SnapshotChangeCategory.NON_BREAKING + assert ( + plan_2.snapshot_for(new_model_a).change_category + == SnapshotChangeCategory.NON_BREAKING + ) assert plan_2.snapshot_for(new_model_a).is_forward_only - assert plan_2.snapshot_for(model_b).change_category == SnapshotChangeCategory.NON_BREAKING + assert ( + plan_2.snapshot_for(model_b).change_category + == SnapshotChangeCategory.NON_BREAKING + ) assert not plan_2.snapshot_for(model_b).is_forward_only # verify that the new snapshots were created correctly @@ -3007,15 +3184,16 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: assert ( plan_2.table_name_for(model_b) not in plan_2.internal_schema_metadata.tables ) # the new main table is not actually created, because it was triggered by a forward-only change. downstream models use the dev table - assert plan_2.table_name_for(model_b) not in plan_2.internal_schema_metadata.managed_tables + assert ( + plan_2.table_name_for(model_b) + not in plan_2.internal_schema_metadata.managed_tables + ) assert ( plan_2.dev_table_name_for(model_b) in plan_2.internal_schema_metadata.tables ) # dev tables are always regular tables for managed models # modify model B, still in the dev environment - new_model_b = load_sql_based_model( - d.parse( # type: ignore - f""" + new_model_b = load_sql_based_model(d.parse(f""" MODEL ( name {schema}.managed_model, kind MANAGED, @@ -3025,9 +3203,7 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: ); SELECT *, 'modified' as extra_b from {schema}.upstream_model; - """ - ) - ) + """)) # type: ignore context.upsert_model(new_model_b) plan_3 = _run_plan(context, "dev") @@ -3035,7 +3211,8 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: assert plan_3.plan.has_changes assert len(plan_3.plan.modified_snapshots) == 1 assert ( - plan_3.modified_snapshot_for(model_b).change_category == SnapshotChangeCategory.NON_BREAKING + plan_3.modified_snapshot_for(model_b).change_category + == SnapshotChangeCategory.NON_BREAKING ) # model A should be unchanged @@ -3050,16 +3227,23 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: ) # still using the dev table, no main table created assert plan_3.dev_table_name_for(model_b) in plan_3.internal_schema_metadata.tables assert ( - plan_3.table_name_for(model_b) not in plan_3.internal_schema_metadata.managed_tables + plan_3.table_name_for(model_b) + not in plan_3.internal_schema_metadata.managed_tables ) # still not a managed table # apply plan to prod plan_4 = _run_plan(context) assert plan_4.plan.has_changes - assert plan_4.snapshot_for(model_a).change_category == SnapshotChangeCategory.NON_BREAKING + assert ( + plan_4.snapshot_for(model_a).change_category + == SnapshotChangeCategory.NON_BREAKING + ) assert plan_4.snapshot_for(model_a).is_forward_only - assert plan_4.snapshot_for(model_b).change_category == SnapshotChangeCategory.NON_BREAKING + assert ( + plan_4.snapshot_for(model_b).change_category + == SnapshotChangeCategory.NON_BREAKING + ) assert not plan_4.snapshot_for(model_b).is_forward_only # verify the Model B table is created as a managed table in prod @@ -3069,7 +3253,9 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: assert ( plan_4.table_name_for(model_b) not in plan_4.internal_schema_metadata.tables ) # however, it should be a managed table, not a normal table - assert plan_4.table_name_for(model_b) in plan_4.internal_schema_metadata.managed_tables + assert ( + plan_4.table_name_for(model_b) in plan_4.internal_schema_metadata.managed_tables + ) @pytest.mark.parametrize( @@ -3078,7 +3264,11 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: (DATA_TYPE.BOOLEAN, (True, False, None), ("1", "0", None)), ( DATA_TYPE.DATE, - (datetime(2023, 1, 1).date(), datetime(2024, 12, 15, 5, 30, 0).date(), None), + ( + datetime(2023, 1, 1).date(), + datetime(2024, 12, 15, 5, 30, 0).date(), + None, + ), ("2023-01-01", "2024-12-15", None), ), ( @@ -3115,7 +3305,9 @@ def _run_plan(sqlmesh_context: Context, environment: str = None) -> PlanResults: DATA_TYPE.TIMESTAMPTZ, ( pytz.timezone("America/Los_Angeles").localize(datetime(2023, 1, 1)), - pytz.timezone("Europe/Athens").localize(datetime(2023, 1, 1, 13, 14, 15)), + pytz.timezone("Europe/Athens").localize( + datetime(2023, 1, 1, 13, 14, 15) + ), pytz.timezone("Pacific/Auckland").localize( datetime(2023, 1, 1, 13, 14, 15, 123456) ), @@ -3159,7 +3351,9 @@ def test_value_normalization( column_type, expressions=[ exp.DataTypeParam( - this=exp.Literal.number(ctx.engine_adapter.MAX_TIMESTAMP_PRECISION) + this=exp.Literal.number( + ctx.engine_adapter.MAX_TIMESTAMP_PRECISION + ) ) ], ) @@ -3184,7 +3378,9 @@ def test_value_normalization( ctx.engine_adapter.create_table( table_name=test_table, target_columns_to_types=columns_to_types_normalized ) - data_query = next(select_from_values(input_data_with_idx, columns_to_types_normalized)) + data_query = next( + select_from_values(input_data_with_idx, columns_to_types_normalized) + ) ctx.engine_adapter.insert_append( table_name=test_table, query_or_df=data_query, @@ -3213,7 +3409,9 @@ def truncate_timestamp(ts: str, precision: int) -> str: digits_to_truncate = 6 - precision return ts[:-digits_to_truncate] if digits_to_truncate > 0 else ts - if full_column_type.is_type(DATA_TYPE.DATETIME, DATA_TYPE.TIMESTAMP, DATA_TYPE.TIMESTAMPTZ): + if full_column_type.is_type( + DATA_TYPE.DATETIME, DATA_TYPE.TIMESTAMP, DATA_TYPE.TIMESTAMPTZ + ): # truncate our expected results to the engine precision expected_results = tuple( truncate_timestamp(e, ctx.engine_adapter.MAX_TIMESTAMP_PRECISION) @@ -3226,7 +3424,9 @@ def truncate_timestamp(ts: str, precision: int) -> str: def test_table_diff_grain_check_single_key(ctx: TestContext): if not isinstance(ctx.engine_adapter, RowDiffMixin): - pytest.skip("table_diff tests are only relevant for engines with row diffing implemented") + pytest.skip( + "table_diff tests are only relevant for engines with row diffing implemented" + ) src_table = ctx.table("source") target_table = ctx.table("target") @@ -3255,10 +3455,14 @@ def test_table_diff_grain_check_single_key(ctx: TestContext): ] ctx.engine_adapter.replace_query( - src_table, pd.DataFrame(src_data, columns=columns_to_types.keys()), columns_to_types + src_table, + pd.DataFrame(src_data, columns=columns_to_types.keys()), + columns_to_types, ) ctx.engine_adapter.replace_query( - target_table, pd.DataFrame(target_data, columns=columns_to_types.keys()), columns_to_types + target_table, + pd.DataFrame(target_data, columns=columns_to_types.keys()), + columns_to_types, ) table_diff = TableDiff( @@ -3290,7 +3494,9 @@ def test_table_diff_grain_check_single_key(ctx: TestContext): def test_table_diff_grain_check_multiple_keys(ctx: TestContext): if not isinstance(ctx.engine_adapter, RowDiffMixin): - pytest.skip("table_diff tests are only relevant for engines with row diffing implemented") + pytest.skip( + "table_diff tests are only relevant for engines with row diffing implemented" + ) src_table = ctx.table("source") target_table = ctx.table("target") @@ -3317,10 +3523,14 @@ def test_table_diff_grain_check_multiple_keys(ctx: TestContext): target_data = src_data + [(1, 6, 1), (1, 5, 3), (None, 2, 3)] ctx.engine_adapter.insert_append( - src_table, next(select_from_values(src_data, columns_to_types)), columns_to_types + src_table, + next(select_from_values(src_data, columns_to_types)), + columns_to_types, ) ctx.engine_adapter.insert_append( - target_table, next(select_from_values(target_data, columns_to_types)), columns_to_types + target_table, + next(select_from_values(target_data, columns_to_types)), + columns_to_types, ) table_diff = TableDiff( @@ -3350,7 +3560,9 @@ def test_table_diff_grain_check_multiple_keys(ctx: TestContext): def test_table_diff_arbitrary_condition(ctx: TestContext): if not isinstance(ctx.engine_adapter, RowDiffMixin): - pytest.skip("table_diff tests are only relevant for engines with row diffing implemented") + pytest.skip( + "table_diff tests are only relevant for engines with row diffing implemented" + ) src_table = ctx.table("source") target_table = ctx.table("target") @@ -3379,7 +3591,9 @@ def test_table_diff_arbitrary_condition(ctx: TestContext): target_data = src_data + [(4, "four", datetime(2024, 2, 1, 8, 13, 14))] ctx.engine_adapter.replace_query( - src_table, pd.DataFrame(src_data, columns=columns_to_types_src.keys()), columns_to_types_src + src_table, + pd.DataFrame(src_data, columns=columns_to_types_src.keys()), + columns_to_types_src, ) ctx.engine_adapter.replace_query( target_table, @@ -3392,7 +3606,9 @@ def test_table_diff_arbitrary_condition(ctx: TestContext): source=exp.table_name(src_table), target=exp.table_name(target_table), on=parse_one('"s"."id" = "t"."item_id"', into=exp.Condition), - where=parse_one("to_char(\"ts\", 'YYYY') = '2024'", dialect="postgres", into=exp.Condition), + where=parse_one( + "to_char(\"ts\", 'YYYY') = '2024'", dialect="postgres", into=exp.Condition + ), ) row_diff = table_diff.row_diff() @@ -3417,7 +3633,9 @@ def test_table_diff_arbitrary_condition(ctx: TestContext): def test_table_diff_identical_dataset(ctx: TestContext): if not isinstance(ctx.engine_adapter, RowDiffMixin): - pytest.skip("table_diff tests are only relevant for engines with row diffing implemented") + pytest.skip( + "table_diff tests are only relevant for engines with row diffing implemented" + ) src_table = ctx.table("source") target_table = ctx.table("target") @@ -3442,10 +3660,14 @@ def test_table_diff_identical_dataset(ctx: TestContext): target_data = src_data ctx.engine_adapter.insert_append( - src_table, next(select_from_values(src_data, columns_to_types)), columns_to_types + src_table, + next(select_from_values(src_data, columns_to_types)), + columns_to_types, ) ctx.engine_adapter.insert_append( - target_table, next(select_from_values(target_data, columns_to_types)), columns_to_types + target_table, + next(select_from_values(target_data, columns_to_types)), + columns_to_types, ) table_diff = TableDiff( @@ -3492,7 +3714,8 @@ def _use_warehouse_as_state_connection(gateway_name: str, config: Config): config.gateways[gateway_name].state_schema = test_schema sqlmesh_context = ctx.create_context( - config_mutator=_use_warehouse_as_state_connection, ephemeral_state_connection=False + config_mutator=_use_warehouse_as_state_connection, + ephemeral_state_connection=False, ) assert sqlmesh_context.config.get_state_schema(ctx.gateway) == test_schema @@ -3613,7 +3836,11 @@ def execute( .replace("DIALECT", ctx.dialect) .replace( "TABLE_FORMAT", - f"table_format='{ctx.default_table_format}'" if ctx.default_table_format else "", + ( + f"table_format='{ctx.default_table_format}'" + if ctx.default_table_format + else "" + ), ) ) ) @@ -3634,7 +3861,9 @@ def execute( # This test uses the dialect= on the model. # For dialect=snowflake, this means that the identifiers are all normalized to uppercase by default expected_result = ( - {"id": 1, "name": "foo"} if ctx.dialect != "snowflake" else {"ID": 1, "NAME": "foo"} + {"id": 1, "name": "foo"} + if ctx.dialect != "snowflake" + else {"ID": 1, "NAME": "foo"} ) assert df.iloc[0].to_dict() == expected_result @@ -3642,7 +3871,9 @@ def execute( def test_identifier_length_limit(ctx: TestContext): adapter = ctx.engine_adapter if adapter.MAX_IDENTIFIER_LENGTH is None: - pytest.skip(f"Engine {adapter.dialect} does not have identifier length limits set.") + pytest.skip( + f"Engine {adapter.dialect} does not have identifier length limits set." + ) long_table_name = "a" * (adapter.MAX_IDENTIFIER_LENGTH + 1) @@ -3664,7 +3895,9 @@ def test_identifier_length_limit(ctx: TestContext): ) @pytest.mark.xdist_group("serial") def test_janitor( - ctx: TestContext, tmp_path: pathlib.Path, environment_suffix_target: EnvironmentSuffixTarget + ctx: TestContext, + tmp_path: pathlib.Path, + environment_suffix_target: EnvironmentSuffixTarget, ): if ( environment_suffix_target == EnvironmentSuffixTarget.CATALOG @@ -3812,7 +4045,9 @@ def test_materialized_view_evaluation(ctx: TestContext): dialect = ctx.dialect if not adapter.SUPPORTS_MATERIALIZED_VIEWS: - pytest.skip(f"Skipping engine {dialect} as it does not support materialized views") + pytest.skip( + f"Skipping engine {dialect} as it does not support materialized views" + ) elif dialect in ("snowflake", "databricks"): pytest.skip(f"Skipping {dialect} as they're not enabled on standard accounts") elif dialect == "starrocks": @@ -3827,29 +4062,17 @@ def test_materialized_view_evaluation(ctx: TestContext): sqlmesh = ctx.create_context() - sqlmesh.upsert_model( - load_sql_based_model( - d.parse( - f""" + sqlmesh.upsert_model(load_sql_based_model(d.parse(f""" MODEL (name {model_name}, kind FULL); SELECT 1 AS col - """ - ) - ) - ) + """))) - sqlmesh.upsert_model( - load_sql_based_model( - d.parse( - f""" + sqlmesh.upsert_model(load_sql_based_model(d.parse(f""" MODEL (name {mview_name}, kind VIEW (materialized true)); SELECT * FROM {model_name} - """ - ) - ) - ) + """))) def _assert_mview_value(value: int): df = adapter.fetchdf(f"SELECT * FROM {mview_name.sql(dialect=dialect)}") @@ -3862,7 +4085,9 @@ def _assert_mview_value(value: int): # Case 2: Ensure that we can change the underlying table and the materialized view is recreated sqlmesh.upsert_model( - load_sql_based_model(d.parse(f"""MODEL (name {model_name}, kind FULL); SELECT 2 AS col""")) + load_sql_based_model( + d.parse(f"""MODEL (name {model_name}, kind FULL); SELECT 2 AS col""") + ) ) logger = logging.getLogger("sqlmesh.core.snapshot.evaluator") @@ -3870,7 +4095,9 @@ def _assert_mview_value(value: int): with mock.patch.object(logger, "info") as mock_logger: sqlmesh.plan(auto_apply=True, no_prompts=True) - assert any("Replacing view" in call[0][0] for call in mock_logger.call_args_list) + assert any( + "Replacing view" in call[0][0] for call in mock_logger.call_args_list + ) _assert_mview_value(value=2) @@ -3880,7 +4107,9 @@ def test_unicode_characters(ctx: TestContext, tmp_path: Path): # at the time of writing this is Spark/Trino and they do this for compatibility reasons. # I also think Spark may not support unicode in general but that would need to be verified. if not ctx.engine_adapter.QUOTE_IDENTIFIERS_IN_VIEWS: - pytest.skip("Skipping as these engines have issues with unicode characters in model names") + pytest.skip( + "Skipping as these engines have issues with unicode characters in model names" + ) model_name = "客户数据" table = ctx.table(model_name).sql(dialect=ctx.dialect) @@ -4103,7 +4332,10 @@ def test_grants_plan(ctx: TestContext, tmp_path: Path): assert set(final_grants.get(select_privilege, [])) == set( expected_final_grants[select_privilege] ) - assert final_grants.get(insert_privilege, []) == expected_final_grants[insert_privilege] + assert ( + final_grants.get(insert_privilege, []) + == expected_final_grants[insert_privilege] + ) # Virtual layer should also have the updated grants updated_virtual_grants = ctx.engine_adapter._get_current_grants_config( diff --git a/tests/core/engine_adapter/integration/test_integration_athena.py b/tests/core/engine_adapter/integration/test_integration_athena.py index 9d23af206e..707b82f6d0 100644 --- a/tests/core/engine_adapter/integration/test_integration_athena.py +++ b/tests/core/engine_adapter/integration/test_integration_athena.py @@ -1,20 +1,20 @@ +import dataclasses +import datetime import typing as t + +import pandas as pd # noqa: TID253 import pytest from pytest import FixtureRequest -import pandas as pd # noqa: TID253 -import datetime +from sqlglot import exp + from sqlmesh.core.engine_adapter import AthenaEngineAdapter from sqlmesh.utils.aws import parse_s3_uri +from sqlmesh.utils.date import TimeLike, to_ds, to_ts from sqlmesh.utils.pandas import columns_to_types_from_df -from sqlmesh.utils.date import to_ds, to_ts, TimeLike -import dataclasses -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) -from sqlglot import exp +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) # The tests in this file dont need to be called twice, so we create a single instance of Athena ENGINE_ATHENA = dataclasses.replace(ENGINES_BY_NAME["athena"], catalog_types=None) @@ -24,7 +24,9 @@ @pytest.fixture(params=list(generate_pytest_params(ENGINE_ATHENA))) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -40,15 +42,21 @@ def s3(engine_adapter: AthenaEngineAdapter) -> t.Any: return engine_adapter._s3_client -def s3_list_objects(s3: t.Any, location: str, **list_objects_kwargs: t.Any) -> t.List[str]: +def s3_list_objects( + s3: t.Any, location: str, **list_objects_kwargs: t.Any +) -> t.List[str]: bucket, prefix = parse_s3_uri(location) lst = [] - for page in s3.get_paginator("list_objects_v2").paginate(Bucket=bucket, Prefix=prefix): + for page in s3.get_paginator("list_objects_v2").paginate( + Bucket=bucket, Prefix=prefix + ): lst.extend([o["Key"] for o in page.get("Contents", [])]) return lst -def test_clear_partition_data(ctx: TestContext, engine_adapter: AthenaEngineAdapter, s3: t.Any): +def test_clear_partition_data( + ctx: TestContext, engine_adapter: AthenaEngineAdapter, s3: t.Any +): base_uri = engine_adapter.s3_warehouse_location_or_raise assert len(s3_list_objects(s3, base_uri)) == 0 @@ -68,8 +76,7 @@ def test_clear_partition_data(ctx: TestContext, engine_adapter: AthenaEngineAdap query_or_df=base_data, ) - sqlmesh_context, model = ctx.upsert_sql_model( - f""" + sqlmesh_context, model = ctx.upsert_sql_model(f""" MODEL ( name {test_table}, kind INCREMENTAL_BY_TIME_RANGE ( @@ -82,8 +89,7 @@ def test_clear_partition_data(ctx: TestContext, engine_adapter: AthenaEngineAdap id, ts, (ts::date)::varchar as ds FROM {src_table} WHERE ts BETWEEN @start_dt AND @end_dt - """ - ) + """) plan = sqlmesh_context.plan(no_prompts=True, auto_apply=True) assert len(plan.snapshots) == 1 @@ -100,7 +106,11 @@ def test_clear_partition_data(ctx: TestContext, engine_adapter: AthenaEngineAdap test_table_physical_name = exp.to_table(test_table_snapshot.table_name()) partitions = engine_adapter._list_partitions(test_table_physical_name, where=None) assert len(partitions) == 3 - assert [p[0] for p in partitions] == [["2023-01-01"], ["2023-01-02"], ["2023-01-03"]] + assert [p[0] for p in partitions] == [ + ["2023-01-01"], + ["2023-01-02"], + ["2023-01-03"], + ] assert engine_adapter.fetchone(f"select count(*) from {test_table}")[0] == 3 # type: ignore @@ -151,8 +161,7 @@ def test_clear_partition_data_multiple_columns( query_or_df=base_data, ) - sqlmesh_context, model = ctx.upsert_sql_model( - f""" + sqlmesh_context, model = ctx.upsert_sql_model(f""" MODEL ( name {test_table}, kind INCREMENTAL_BY_TIME_RANGE ( @@ -166,8 +175,7 @@ def test_clear_partition_data_multiple_columns( id, ts, (ts::date)::varchar as ds, system FROM {src_table} WHERE ts BETWEEN @start_dt AND @end_dt - """ - ) + """) plan = sqlmesh_context.plan(no_prompts=True, auto_apply=True) assert len(plan.snapshots) == 1 @@ -222,7 +230,9 @@ def _match_partition(location_list: t.List[str], match: str): assert engine_adapter.fetchone(f"select count(*) from {test_table}")[0] == 4 # type: ignore -def test_hive_truncate_table(ctx: TestContext, engine_adapter: AthenaEngineAdapter, s3: t.Any): +def test_hive_truncate_table( + ctx: TestContext, engine_adapter: AthenaEngineAdapter, s3: t.Any +): base_uri = engine_adapter.s3_warehouse_location_or_raise table_1 = ctx.table("table_one") @@ -269,7 +279,9 @@ def test_hive_truncate_table(ctx: TestContext, engine_adapter: AthenaEngineAdapt engine_adapter._truncate_table(table_1) -def test_hive_drop_table_removes_data(ctx: TestContext, engine_adapter: AthenaEngineAdapter): +def test_hive_drop_table_removes_data( + ctx: TestContext, engine_adapter: AthenaEngineAdapter +): # check no exception with dropping a table that doesnt exist engine_adapter.drop_table("nonexist") @@ -287,7 +299,9 @@ def test_hive_drop_table_removes_data(ctx: TestContext, engine_adapter: AthenaEn table_name=seed_table, target_columns_to_types=columns_to_types, exists=False ) engine_adapter.insert_append( - table_name=seed_table, query_or_df=data, target_columns_to_types=columns_to_types + table_name=seed_table, + query_or_df=data, + target_columns_to_types=columns_to_types, ) assert engine_adapter.fetchone(f"select count(*) from {seed_table}")[0] == 1 # type: ignore @@ -300,7 +314,9 @@ def test_hive_drop_table_removes_data(ctx: TestContext, engine_adapter: AthenaEn assert engine_adapter.fetchone(f"select count(*) from {seed_table}")[0] == 0 # type: ignore -def test_hive_replace_query_same_schema(ctx: TestContext, engine_adapter: AthenaEngineAdapter): +def test_hive_replace_query_same_schema( + ctx: TestContext, engine_adapter: AthenaEngineAdapter +): seed_table = ctx.table("seed") data = pd.DataFrame( @@ -323,7 +339,9 @@ def test_hive_replace_query_same_schema(ctx: TestContext, engine_adapter: Athena assert engine_adapter.fetchone(f"select count(*) from {seed_table}")[0] == 3 # type: ignore -def test_hive_replace_query_new_schema(ctx: TestContext, engine_adapter: AthenaEngineAdapter): +def test_hive_replace_query_new_schema( + ctx: TestContext, engine_adapter: AthenaEngineAdapter +): seed_table = ctx.table("seed") orig_data = pd.DataFrame( @@ -341,7 +359,9 @@ def test_hive_replace_query_new_schema(ctx: TestContext, engine_adapter: AthenaE engine_adapter.replace_query(table_name=seed_table, query_or_df=orig_data) - assert engine_adapter.fetchall(f"select id, name from {seed_table} order by id") == [ + assert engine_adapter.fetchall( + f"select id, name from {seed_table} order by id" + ) == [ (1, "one"), (2, "two"), ] @@ -378,7 +398,9 @@ def test_insert_overwrite_by_time_partition_date_type( ), # note: columns_to_types_from_df() would infer this as TEXT but we need a DATE type } - def time_formatter(time: TimeLike, _: t.Optional[t.Dict[str, exp.DataType]]) -> exp.Expr: + def time_formatter( + time: TimeLike, _: t.Optional[t.Dict[str, exp.DataType]] + ) -> exp.Expr: return exp.cast(exp.Literal.string(to_ds(time)), "date") engine_adapter.create_table( @@ -400,7 +422,10 @@ def time_formatter(time: TimeLike, _: t.Optional[t.Dict[str, exp.DataType]]) -> new_data = pd.DataFrame( [ - {"id": 4, "date": datetime.date(2023, 1, 3)}, # replaces the old entry for 2023-01-03 + { + "id": 4, + "date": datetime.date(2023, 1, 3), + }, # replaces the old entry for 2023-01-03 {"id": 5, "date": datetime.date(2023, 1, 4)}, ] ) @@ -440,7 +465,9 @@ def test_insert_overwrite_by_time_partition_datetime_type( ), # note: columns_to_types_from_df() would infer this as TEXT but we need a DATETIME type } - def time_formatter(time: TimeLike, _: t.Optional[t.Dict[str, exp.DataType]]) -> exp.Expr: + def time_formatter( + time: TimeLike, _: t.Optional[t.Dict[str, exp.DataType]] + ) -> exp.Expr: return exp.cast(exp.Literal.string(to_ts(time)), "datetime") engine_adapter.create_table( @@ -507,8 +534,7 @@ def test_scd_type_2_iceberg_timestamps( query_or_df=base_data, ) - sqlmesh_context, model = ctx.upsert_sql_model( - f""" + sqlmesh_context, model = ctx.upsert_sql_model(f""" MODEL ( name {scd_model_table}, kind SCD_TYPE_2_BY_TIME ( @@ -524,8 +550,7 @@ def test_scd_type_2_iceberg_timestamps( SELECT id, ts::timestamp(6) as ts FROM {src_table}; - """ - ) + """) assert model.table_format == "iceberg" @@ -544,4 +569,6 @@ def test_scd_type_2_iceberg_timestamps( if k in {"ts", "valid_from", "valid_to"} ] assert len(timestamp_columns) == 3 - assert all([v.sql(dialect="athena").lower() == "timestamp(6)" for v in timestamp_columns]) + assert all( + [v.sql(dialect="athena").lower() == "timestamp(6)" for v in timestamp_columns] + ) diff --git a/tests/core/engine_adapter/integration/test_integration_bigquery.py b/tests/core/engine_adapter/integration/test_integration_bigquery.py index ce31c255ad..c98ccbc09c 100644 --- a/tests/core/engine_adapter/integration/test_integration_bigquery.py +++ b/tests/core/engine_adapter/integration/test_integration_bigquery.py @@ -1,39 +1,37 @@ import typing as t +from pathlib import Path from unittest import mock import pytest import time_machine -from pathlib import Path +from pytest import FixtureRequest from sqlglot import exp -from sqlglot.optimizer.qualify_columns import quote_identifiers from sqlglot.helper import seq_get +from sqlglot.optimizer.qualify_columns import quote_identifiers + +import sqlmesh.core.dialect as d from sqlmesh.cli.project_init import ProjectTemplate, init_example_project from sqlmesh.core.config import Config from sqlmesh.core.engine_adapter import BigQueryEngineAdapter from sqlmesh.core.engine_adapter.mixins import ( - TableAlterDropClusterKeyOperation, - TableAlterChangeClusterKeyOperation, -) + TableAlterChangeClusterKeyOperation, TableAlterDropClusterKeyOperation) from sqlmesh.core.engine_adapter.shared import DataObject -import sqlmesh.core.dialect as d from sqlmesh.core.model import SqlModel, load_sql_based_model -from sqlmesh.core.plan import Plan, BuiltInPlanEvaluator +from sqlmesh.core.plan import BuiltInPlanEvaluator, Plan from sqlmesh.core.table_diff import TableDiff from sqlmesh.utils import CorrelationId -from tests.core.engine_adapter.integration import TestContext -from pytest import FixtureRequest -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) @pytest.fixture(params=list(generate_pytest_params(ENGINES_BY_NAME["bigquery"]))) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -62,9 +60,12 @@ def test_get_alter_expressions_includes_clustering( ) metadata = engine_adapter.get_data_objects( - normal_table.db, {clustered_table.name, clustered_differently_table.name, normal_table.name} + normal_table.db, + {clustered_table.name, clustered_differently_table.name, normal_table.name}, + ) + clustered_table_metadata = next( + md for md in metadata if md.name == clustered_table.name ) - clustered_table_metadata = next(md for md in metadata if md.name == clustered_table.name) clustered_differently_table_metadata = next( md for md in metadata if md.name == clustered_differently_table.name ) @@ -75,17 +76,23 @@ def test_get_alter_expressions_includes_clustering( assert normal_table_metadata.clustering_key is None assert len(engine_adapter.get_alter_operations(normal_table, normal_table)) == 0 - assert len(engine_adapter.get_alter_operations(clustered_table, clustered_table)) == 0 + assert ( + len(engine_adapter.get_alter_operations(clustered_table, clustered_table)) == 0 + ) # alter table drop clustered - clustered_to_normal = engine_adapter.get_alter_operations(clustered_table, normal_table) + clustered_to_normal = engine_adapter.get_alter_operations( + clustered_table, normal_table + ) assert len(clustered_to_normal) == 1 assert isinstance(clustered_to_normal[0], TableAlterDropClusterKeyOperation) assert clustered_to_normal[0].target_table == clustered_table assert not hasattr(clustered_to_normal[0], "clustering_key") # alter table add clustered - normal_to_clustered = engine_adapter.get_alter_operations(normal_table, clustered_table) + normal_to_clustered = engine_adapter.get_alter_operations( + normal_table, clustered_table + ) assert len(normal_to_clustered) == 1 operation = normal_to_clustered[0] assert isinstance(operation, TableAlterChangeClusterKeyOperation) @@ -124,9 +131,7 @@ def _create_model(**kwargs: t.Any) -> SqlModel: extra_props = "\n".join([f"{k} {v}," for k, v in kwargs.items()]) return t.cast( SqlModel, - load_sql_based_model( - d.parse( - f""" + load_sql_based_model(d.parse(f""" MODEL ( name {model_name.sql(dialect="bigquery")}, kind INCREMENTAL_BY_TIME_RANGE ( @@ -139,13 +144,13 @@ def _create_model(**kwargs: t.Any) -> SqlModel: ); select 1 as ID, current_date() as partitiondate - """ - ) - ), + """)), ) def _get_data_object(table: exp.Table) -> DataObject: - data_object = seq_get(engine_adapter.get_data_objects(table.db, {table.name}), 0) + data_object = seq_get( + engine_adapter.get_data_objects(table.db, {table.name}), 0 + ) if not data_object: raise ValueError(f"Expected metadata for {table}") return data_object @@ -219,10 +224,11 @@ def test_information_schema_view_external_model(ctx: TestContext, tmp_path: Path model_name = ctx.table("test") dependency = f"`{'.'.join(part.name for part in information_schema_tables.parts)}`" - init_example_project(tmp_path, engine_type="bigquery", template=ProjectTemplate.EMPTY) + init_example_project( + tmp_path, engine_type="bigquery", template=ProjectTemplate.EMPTY + ) with open(tmp_path / "models" / "test.sql", "w", encoding="utf-8") as f: - f.write( - f""" + f.write(f""" MODEL ( name {model_name.sql("bigquery")}, kind FULL, @@ -230,8 +236,7 @@ def test_information_schema_view_external_model(ctx: TestContext, tmp_path: Path ); SELECT * FROM {dependency} AS tables - """ - ) + """) def _mutate_config(_: str, config: Config) -> None: config.model_defaults.dialect = "bigquery" @@ -240,7 +245,9 @@ def _mutate_config(_: str, config: Config) -> None: sqlmesh.create_external_models() sqlmesh.load() - actual_columns_to_types = sqlmesh.get_model(information_schema_tables.sql()).columns_to_types + actual_columns_to_types = sqlmesh.get_model( + information_schema_tables.sql() + ).columns_to_types expected_columns_to_types = { "table_catalog": exp.DataType.build("TEXT"), "table_schema": exp.DataType.build("TEXT"), @@ -377,19 +384,22 @@ def test_get_bq_schema(ctx: TestContext, engine_adapter: BigQueryEngineAdapter): SchemaField(name="address", field_type="STRING", mode="NULLABLE"), ], ) - assert bg_schema[2] == SchemaField(name="tags", field_type="STRING", mode="REPEATED") - assert bg_schema[3] == SchemaField(name="score", field_type="NUMERIC", mode="NULLABLE") - assert bg_schema[4] == SchemaField(name="created_at", field_type="DATETIME", mode="NULLABLE") + assert bg_schema[2] == SchemaField( + name="tags", field_type="STRING", mode="REPEATED" + ) + assert bg_schema[3] == SchemaField( + name="score", field_type="NUMERIC", mode="NULLABLE" + ) + assert bg_schema[4] == SchemaField( + name="created_at", field_type="DATETIME", mode="NULLABLE" + ) def test_column_types(ctx: TestContext): model_name = ctx.table("test") sqlmesh = ctx.create_context() - sqlmesh.upsert_model( - load_sql_based_model( - d.parse( - f""" + sqlmesh.upsert_model(load_sql_based_model(d.parse(f""" MODEL ( name {model_name}, ); @@ -398,10 +408,7 @@ def test_column_types(ctx: TestContext): RANGE('01-01-1900'::DATE, '01-01-1902'::DATE) AS col1, JSON '{{"id": 10}}' AS col2, STRUCT([PARSE_JSON('{{"id": 10}}')] AS arr) AS col3; - """ - ) - ) - ) + """))) sqlmesh.plan(auto_apply=True, no_prompts=True) @@ -449,7 +456,9 @@ def test_plan_correlation_id_in_job_labels(ctx: TestContext): sqlmesh = ctx.create_context() sqlmesh.upsert_model( - load_sql_based_model(d.parse(f"MODEL (name {model_name}, kind FULL); SELECT 1 AS col")) + load_sql_based_model( + d.parse(f"MODEL (name {model_name}, kind FULL); SELECT 1 AS col") + ) ) # Create a plan evaluator and a plan to evaluate @@ -480,10 +489,7 @@ def test_run_correlation_id_in_job_labels(ctx: TestContext): model_name = ctx.table("run_test") sqlmesh = ctx.create_context() - sqlmesh.upsert_model( - load_sql_based_model( - d.parse( - f""" + sqlmesh.upsert_model(load_sql_based_model(d.parse(f""" MODEL ( name {model_name}, kind INCREMENTAL_BY_TIME_RANGE ( @@ -493,10 +499,7 @@ def test_run_correlation_id_in_job_labels(ctx: TestContext): start '2023-01-07' ); SELECT 1 AS col, '2023-01-07' AS event_ts -""" - ) - ) - ) +"""))) sqlmesh.plan(auto_apply=True, no_prompts=True) captured_evaluators: t.List = [] @@ -508,7 +511,9 @@ def scheduler_wrapper( ): if snapshot_evaluator is not None: captured_evaluators.append(snapshot_evaluator) - return original_scheduler(environment=environment, snapshot_evaluator=snapshot_evaluator) + return original_scheduler( + environment=environment, snapshot_evaluator=snapshot_evaluator + ) with time_machine.travel("2023-01-09 00:00:00 UTC"): with mock.patch( diff --git a/tests/core/engine_adapter/integration/test_integration_clickhouse.py b/tests/core/engine_adapter/integration/test_integration_clickhouse.py index 4420acec71..919ffad221 100644 --- a/tests/core/engine_adapter/integration/test_integration_clickhouse.py +++ b/tests/core/engine_adapter/integration/test_integration_clickhouse.py @@ -1,28 +1,30 @@ import typing as t + +import pandas as pd # noqa: TID253 import pytest from pytest import FixtureRequest -from tests.core.engine_adapter.integration import TestContext -from sqlmesh.core.engine_adapter.clickhouse import ClickhouseEngineAdapter -import pandas as pd # noqa: TID253 from sqlglot import exp, parse_one -from sqlmesh.core.snapshot import SnapshotChangeCategory -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) +from sqlmesh.core.engine_adapter.clickhouse import ClickhouseEngineAdapter +from sqlmesh.core.snapshot import SnapshotChangeCategory +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) @pytest.fixture( params=list( - generate_pytest_params([ENGINES_BY_NAME["clickhouse"], ENGINES_BY_NAME["clickhouse_cloud"]]) + generate_pytest_params( + [ENGINES_BY_NAME["clickhouse"], ENGINES_BY_NAME["clickhouse_cloud"]] + ) ) ) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -64,7 +66,9 @@ def _create_table_and_insert_existing_data( "ds": exp.DataType.build("Date", "clickhouse"), }, table_name: str = "data_existing", - partitioned_by: t.Optional[t.List[exp.Expr]] = [parse_one("toMonth(ds)", dialect="clickhouse")], + partitioned_by: t.Optional[t.List[exp.Expr]] = [ + parse_one("toMonth(ds)", dialect="clickhouse") + ], ) -> exp.Table: existing_data = existing_data existing_table_name: exp.Table = ctx.table(table_name) @@ -114,7 +118,9 @@ def test_insert_overwrite_by_condition_replace_partitioned(ctx: TestContext): def test_insert_overwrite_by_condition_replace(ctx: TestContext): - existing_table_name = _create_table_and_insert_existing_data(ctx, partitioned_by=None) + existing_table_name = _create_table_and_insert_existing_data( + ctx, partitioned_by=None + ) # new data to insert insert_data = pd.DataFrame( @@ -223,7 +229,10 @@ def test_insert_overwrite_by_condition_where_compound_partitioned(ctx: TestConte ] ), columns_to_types=compound_columns_to_types, - partitioned_by=[parse_one("toMonth(ds)", dialect=ctx.dialect), exp.column("city")], + partitioned_by=[ + parse_one("toMonth(ds)", dialect=ctx.dialect), + exp.column("city"), + ], ) # new data to insert @@ -280,7 +289,9 @@ def test_insert_overwrite_by_condition_by_key(ctx: TestContext): key_exp = key[0] # data currently in target table - existing_table_name = _create_table_and_insert_existing_data(ctx, partitioned_by=None) + existing_table_name = _create_table_and_insert_existing_data( + ctx, partitioned_by=None + ) # new data to insert insert_data = pd.DataFrame( @@ -500,8 +511,7 @@ def test_inc_by_time_auto_partition_string(ctx: TestContext): partitioned_by=None, ) - sqlmesh_context, model = ctx.upsert_sql_model( - f""" + sqlmesh_context, model = ctx.upsert_sql_model(f""" MODEL ( name test.inc_by_time_no_partition, kind INCREMENTAL_BY_TIME_RANGE ( @@ -516,8 +526,7 @@ def test_inc_by_time_auto_partition_string(ctx: TestContext): ds::String FROM {existing_table_name.sql()} WHERE ds BETWEEN @start_ds AND @end_ds - """ - ) + """) plan = sqlmesh_context.plan(no_prompts=True, auto_apply=True) diff --git a/tests/core/engine_adapter/integration/test_integration_duckdb.py b/tests/core/engine_adapter/integration/test_integration_duckdb.py index a53c559a55..b8f715e058 100644 --- a/tests/core/engine_adapter/integration/test_integration_duckdb.py +++ b/tests/core/engine_adapter/integration/test_integration_duckdb.py @@ -1,10 +1,11 @@ +import random import typing as t +from concurrent.futures import ThreadPoolExecutor, as_completed +from pathlib import Path +from threading import Thread, current_thread + import pytest -from threading import current_thread, Thread -import random from sqlglot import exp -from pathlib import Path -from concurrent.futures import ThreadPoolExecutor, as_completed from sqlmesh.core.config.connection import DuckDBConnectionConfig from sqlmesh.utils.connection_pool import ThreadLocalSharedConnectionPool @@ -37,7 +38,9 @@ def test_multithread_concurrency(tmp_path: Path, database: t.Optional[str]): def write_from_thread(): thread_name = str(current_thread().name) query = exp.insert( - exp.values([(exp.Literal.string(thread_name),)]), "tbl", columns=["thread_name"] + exp.values([(exp.Literal.string(thread_name),)]), + "tbl", + columns=["thread_name"], ) adapter.execute(query) adapter.execute(f"CREATE TABLE thread_{thread_name} (id int)") @@ -82,7 +85,14 @@ def test_secret_registration_from_multiple_connections(tmp_path: Path): config = DuckDBConnectionConfig( database=database, concurrent_tasks=2, - secrets={"s3": {"type": "s3", "region": "us-east-1", "key_id": "foo", "secret": "bar"}}, + secrets={ + "s3": { + "type": "s3", + "region": "us-east-1", + "key_id": "foo", + "secret": "bar", + } + }, ) adapter = config.create_engine_adapter() diff --git a/tests/core/engine_adapter/integration/test_integration_fabric.py b/tests/core/engine_adapter/integration/test_integration_fabric.py index 41f399b3b8..3157bd5ddd 100644 --- a/tests/core/engine_adapter/integration/test_integration_fabric.py +++ b/tests/core/engine_adapter/integration/test_integration_fabric.py @@ -1,27 +1,29 @@ -import typing as t -import threading import queue +import threading +import typing as t +from concurrent.futures import ThreadPoolExecutor + import pytest from pytest import FixtureRequest + from sqlmesh.core.engine_adapter import FabricEngineAdapter from sqlmesh.utils.connection_pool import ThreadLocalConnectionPool -from tests.core.engine_adapter.integration import TestContext -from concurrent.futures import ThreadPoolExecutor - -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) @pytest.fixture( - params=list(generate_pytest_params(ENGINES_BY_NAME["fabric"], show_variant_in_test_id=False)) + params=list( + generate_pytest_params(ENGINES_BY_NAME["fabric"], show_variant_in_test_id=False) + ) ) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -92,7 +94,9 @@ def _set_and_return_catalog_in_another_thread( with ThreadPoolExecutor() as executor: lock.acquire() # we have the lock, thread will be blocked until we release it - future = executor.submit(_set_and_return_catalog_in_another_thread, q, engine_adapter) + future = executor.submit( + _set_and_return_catalog_in_another_thread, q, engine_adapter + ) assert q.get() == "thread_started" assert not future.done() diff --git a/tests/core/engine_adapter/integration/test_integration_postgres.py b/tests/core/engine_adapter/integration/test_integration_postgres.py index f236fdebce..fdab60e694 100644 --- a/tests/core/engine_adapter/integration/test_integration_postgres.py +++ b/tests/core/engine_adapter/integration/test_integration_postgres.py @@ -1,28 +1,27 @@ import typing as t +import uuid from contextlib import contextmanager +from datetime import timedelta +from pathlib import Path + import pytest +import time_machine from pytest import FixtureRequest -from pathlib import Path -from sqlmesh.core.engine_adapter import PostgresEngineAdapter +from sqlglot import exp + from sqlmesh.core.config import Config, DuckDBConnectionConfig from sqlmesh.core.config.common import VirtualEnvironmentMode -from tests.core.engine_adapter.integration import TestContext -import time_machine -from datetime import timedelta -from sqlmesh.utils.date import to_ds -from sqlglot import exp from sqlmesh.core.context import Context -from sqlmesh.core.state_sync import CachingStateSync, EngineAdapterStateSync +from sqlmesh.core.engine_adapter import PostgresEngineAdapter from sqlmesh.core.snapshot.definition import SnapshotId +from sqlmesh.core.state_sync import CachingStateSync, EngineAdapterStateSync from sqlmesh.utils import random_id - -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, - TEST_SCHEMA, -) +from sqlmesh.utils.date import to_ds +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + TEST_SCHEMA, + IntegrationTestEngine, + TestContext, + generate_pytest_params) def _cleanup_user(engine_adapter: PostgresEngineAdapter, user_name: str) -> None: @@ -49,17 +48,17 @@ def create_users( try: for role_name in role_names: - user_name = f"test_{role_name}" - _cleanup_user(engine_adapter, user_name) + random_suffix = uuid.uuid4().hex[:6] + user_name = f"test_{role_name}_{random_suffix}" - for role_name in role_names: - user_name = f"test_{role_name}" password = random_id() - engine_adapter.execute(f"CREATE USER \"{user_name}\" WITH PASSWORD '{password}'") + engine_adapter.execute( + f"CREATE USER \"{user_name}\" WITH PASSWORD '{password}'" + ) engine_adapter.execute(f'GRANT USAGE ON SCHEMA public TO "{user_name}"') + created_users.append(user_name) roles[role_name] = {"username": user_name, "password": password} - yield roles finally: @@ -109,7 +108,9 @@ def engine_adapter_for_role( @pytest.fixture(params=list(generate_pytest_params(ENGINES_BY_NAME["postgres"]))) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -191,7 +192,9 @@ def _mutate_config(gateway: str, config: Config): with time_machine.travel("2020-01-01 00:00:00"): sqlmesh = ctx.create_context( - path=tmp_path, config_mutator=_mutate_config, ephemeral_state_connection=False + path=tmp_path, + config_mutator=_mutate_config, + ephemeral_state_connection=False, ) sqlmesh.plan(auto_apply=True) @@ -223,7 +226,9 @@ def _mutate_config(gateway: str, config: Config): """) sqlmesh = ctx.create_context( - path=tmp_path, config_mutator=_mutate_config, ephemeral_state_connection=False + path=tmp_path, + config_mutator=_mutate_config, + ephemeral_state_connection=False, ) sqlmesh.plan(environment="dev", auto_apply=True) @@ -238,12 +243,16 @@ def _mutate_config(gateway: str, config: Config): assert len(sqlmesh.snapshots) == 2 # these expire 1 day later than what's in prod - model_a_snapshot = next(s for n, s in sqlmesh.snapshots.items() if "model_a" in n) + model_a_snapshot = next( + s for n, s in sqlmesh.snapshots.items() if "model_a" in n + ) assert timedelta(milliseconds=model_a_snapshot.ttl_ms) == timedelta(weeks=1) assert to_ds(model_a_snapshot.updated_ts) == "2020-01-02" assert to_ds(model_a_snapshot.expiration_ts) == "2020-01-09" - model_b_snapshot = next(s for n, s in sqlmesh.snapshots.items() if "model_b" in n) + model_b_snapshot = next( + s for n, s in sqlmesh.snapshots.items() if "model_b" in n + ) assert timedelta(milliseconds=model_b_snapshot.ttl_ms) == timedelta(weeks=1) assert to_ds(model_b_snapshot.updated_ts) == "2020-01-02" assert to_ds(model_b_snapshot.expiration_ts) == "2020-01-09" @@ -277,7 +286,9 @@ def _mutate_config(gateway: str, config: Config): """) sqlmesh = ctx.create_context( - path=tmp_path, config_mutator=_mutate_config, ephemeral_state_connection=False + path=tmp_path, + config_mutator=_mutate_config, + ephemeral_state_connection=False, ) # need run=True to prevent a "start date is greater than end date" error # since dev cant exceed what is in prod, and prod has no cadence runs, @@ -296,26 +307,38 @@ def _mutate_config(gateway: str, config: Config): assert len(sqlmesh.snapshots) == 4 # model a expiry should not have changed - model_a_snapshot = next(s for n, s in sqlmesh.snapshots.items() if "model_a" in n) + model_a_snapshot = next( + s for n, s in sqlmesh.snapshots.items() if "model_a" in n + ) assert timedelta(milliseconds=model_a_snapshot.ttl_ms) == timedelta(weeks=1) assert to_ds(model_a_snapshot.updated_ts) == "2020-01-02" assert to_ds(model_a_snapshot.expiration_ts) == "2020-01-09" # model b should now expire well after model a - model_b_snapshot = next(s for n, s in sqlmesh.snapshots.items() if "model_b" in n) + model_b_snapshot = next( + s for n, s in sqlmesh.snapshots.items() if "model_b" in n + ) assert timedelta(milliseconds=model_b_snapshot.ttl_ms) == timedelta(weeks=1) assert to_ds(model_b_snapshot.updated_ts) == "2020-01-05" assert to_ds(model_b_snapshot.expiration_ts) == "2020-01-12" # model c should expire at the same time as model b - model_c_snapshot = next(s for n, s in sqlmesh.snapshots.items() if "model_c" in n) + model_c_snapshot = next( + s for n, s in sqlmesh.snapshots.items() if "model_c" in n + ) assert to_ds(model_c_snapshot.updated_ts) == to_ds(model_b_snapshot.updated_ts) - assert to_ds(model_c_snapshot.expiration_ts) == to_ds(model_b_snapshot.expiration_ts) + assert to_ds(model_c_snapshot.expiration_ts) == to_ds( + model_b_snapshot.expiration_ts + ) # model d should expire at the same time as model b - model_d_snapshot = next(s for n, s in sqlmesh.snapshots.items() if "model_d" in n) + model_d_snapshot = next( + s for n, s in sqlmesh.snapshots.items() if "model_d" in n + ) assert to_ds(model_d_snapshot.updated_ts) == to_ds(model_b_snapshot.updated_ts) - assert to_ds(model_d_snapshot.expiration_ts) == to_ds(model_b_snapshot.expiration_ts) + assert to_ds(model_d_snapshot.expiration_ts) == to_ds( + model_b_snapshot.expiration_ts + ) # move forward to date where after model a has expired but before model b has expired # invalidate dev to trigger cleanups @@ -325,7 +348,9 @@ def _mutate_config(gateway: str, config: Config): # - table model d is a not a view, so even though its parent view model b got dropped, it doesnt need to be dropped with time_machine.travel("2020-01-10 00:00:00"): sqlmesh = ctx.create_context( - path=tmp_path, config_mutator=_mutate_config, ephemeral_state_connection=False + path=tmp_path, + config_mutator=_mutate_config, + ephemeral_state_connection=False, ) before_snapshot_ids = _all_snapshot_ids(sqlmesh) @@ -460,8 +485,7 @@ def test_grants_plan_full_refresh_model_via_replace( ): with create_users(engine_adapter, "reader") as roles: (tmp_path / "models").mkdir(exist_ok=True) - (tmp_path / "models" / "full_refresh_model.sql").write_text( - f""" + (tmp_path / "models" / "full_refresh_model.sql").write_text(f""" MODEL ( name test_schema.full_refresh_model, kind FULL, @@ -471,8 +495,7 @@ def test_grants_plan_full_refresh_model_via_replace( grants_target_layer 'all' ); SELECT 1 as id, 'test_data' as status - """ - ) + """) context = ctx.create_context(path=tmp_path) @@ -527,7 +550,11 @@ def test_grants_plan_incremental_model( context = ctx.create_context(path=tmp_path) plan_result = context.plan( - "dev", start="2020-01-01", end="2020-01-01", auto_apply=True, no_prompts=True + "dev", + start="2020-01-01", + end="2020-01-01", + auto_apply=True, + no_prompts=True, ) assert len(plan_result.new_snapshots) == 1 @@ -553,8 +580,7 @@ def test_grants_plan_clone_environment( ): with create_users(engine_adapter, "reader") as roles: (tmp_path / "models").mkdir(exist_ok=True) - (tmp_path / "models" / "clone_model.sql").write_text( - f""" + (tmp_path / "models" / "clone_model.sql").write_text(f""" MODEL ( name test_schema.clone_model, kind FULL, @@ -565,8 +591,7 @@ def test_grants_plan_clone_environment( ); SELECT 1 as id, 'data' as value - """ - ) + """) context = ctx.create_context(path=tmp_path) prod_plan_result = context.plan("prod", auto_apply=True, no_prompts=True) @@ -713,13 +738,17 @@ def test_grants_metadata_only_changes( updated_physical_grants = engine_adapter._get_current_grants_config( exp.to_table(physical_table_name, dialect=engine_adapter.dialect) ) - assert set(updated_physical_grants.get("SELECT", [])) == set(expected_grants["SELECT"]) + assert set(updated_physical_grants.get("SELECT", [])) == set( + expected_grants["SELECT"] + ) assert updated_physical_grants.get("INSERT", []) == expected_grants["INSERT"] updated_virtual_grants = engine_adapter._get_current_grants_config( exp.to_table(virtual_view_name, dialect=engine_adapter.dialect) ) - assert set(updated_virtual_grants.get("SELECT", [])) == set(expected_grants["SELECT"]) + assert set(updated_virtual_grants.get("SELECT", [])) == set( + expected_grants["SELECT"] + ) assert updated_virtual_grants.get("INSERT", []) == expected_grants["INSERT"] @@ -748,9 +777,7 @@ def test_grants_target_layer_with_vde_dev_only( (tmp_path / "models").mkdir(exist_ok=True) if model_kind == "VIEW": - grants_config = ( - f"'SELECT' = ['{roles['reader']['username']}', '{roles['writer']['username']}']" - ) + grants_config = f"'SELECT' = ['{roles['reader']['username']}', '{roles['writer']['username']}']" else: grants_config = f""" 'SELECT' = ['{roles["reader"]["username"]}', '{roles["writer"]["username"]}'], @@ -769,7 +796,9 @@ def test_grants_target_layer_with_vde_dev_only( SELECT 1 as id, '{grants_target_layer}_{model_kind}' as test_type """ ( - tmp_path / "models" / f"vde_model_{grants_target_layer}_{model_kind.lower()}.sql" + tmp_path + / "models" + / f"vde_model_{grants_target_layer}_{model_kind.lower()}.sql" ).write_text(model_def) context = ctx.create_context(path=tmp_path, config_mutator=_vde_dev_only_config) @@ -983,7 +1012,9 @@ def test_grants_target_layer_plan_env_with_vde_dev_only( ) assert roles["grantee"]["username"] in grants.get("SELECT", []) else: - context.plan(environment, auto_apply=True, no_prompts=True, include_unmodified=True) + context.plan( + environment, auto_apply=True, no_prompts=True, include_unmodified=True + ) virtual_view = f"test_schema__{environment}.vde_layer_model" assert context.engine_adapter.table_exists(virtual_view) virtual_grants = engine_adapter._get_current_grants_config( @@ -1007,21 +1038,31 @@ def test_grants_target_layer_plan_env_with_vde_dev_only( for physical_table in physical_tables: physical_table_name = f"sqlmesh__test_schema.{physical_table.name}" physical_grants = engine_adapter._get_current_grants_config( - exp.to_table(physical_table_name, dialect=engine_adapter.dialect) + exp.to_table( + physical_table_name, dialect=engine_adapter.dialect + ) + ) + assert roles["grantee"]["username"] not in physical_grants.get( + "SELECT", [] ) - assert roles["grantee"]["username"] not in physical_grants.get("SELECT", []) elif grants_target_layer == "physical": # Virtual layer should not have grants, physical should - assert roles["grantee"]["username"] not in virtual_grants.get("SELECT", []) + assert roles["grantee"]["username"] not in virtual_grants.get( + "SELECT", [] + ) assert len(physical_tables) > 0 for physical_table in physical_tables: physical_table_name = f"sqlmesh__test_schema.{physical_table.name}" physical_grants = engine_adapter._get_current_grants_config( - exp.to_table(physical_table_name, dialect=engine_adapter.dialect) + exp.to_table( + physical_table_name, dialect=engine_adapter.dialect + ) + ) + assert roles["grantee"]["username"] in physical_grants.get( + "SELECT", [] ) - assert roles["grantee"]["username"] in physical_grants.get("SELECT", []) else: # grants_target_layer == "all" # Both layers should have grants @@ -1030,9 +1071,13 @@ def test_grants_target_layer_plan_env_with_vde_dev_only( for physical_table in physical_tables: physical_table_name = f"sqlmesh__test_schema.{physical_table.name}" physical_grants = engine_adapter._get_current_grants_config( - exp.to_table(physical_table_name, dialect=engine_adapter.dialect) + exp.to_table( + physical_table_name, dialect=engine_adapter.dialect + ) + ) + assert roles["grantee"]["username"] in physical_grants.get( + "SELECT", [] ) - assert roles["grantee"]["username"] in physical_grants.get("SELECT", []) @pytest.mark.parametrize( @@ -1069,15 +1114,17 @@ def test_grants_plan_scd_type_2_models( context = ctx.create_context(path=tmp_path) plan_result = context.plan( - "dev", start="2023-01-01", end="2023-01-01", auto_apply=True, no_prompts=True + "dev", + start="2023-01-01", + end="2023-01-01", + auto_apply=True, + no_prompts=True, ) assert len(plan_result.new_snapshots) == 1 current_snapshot = plan_result.new_snapshots[0] fingerprint_version = current_snapshot.fingerprint.to_version() - physical_table_name = ( - f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint_version}__dev" - ) + physical_table_name = f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint_version}__dev" physical_grants = engine_adapter._get_current_grants_config( exp.to_table(physical_table_name, dialect=engine_adapter.dialect) ) @@ -1107,12 +1154,20 @@ def test_grants_plan_scd_type_2_models( (tmp_path / "models" / f"{model_name}.sql").write_text(updated_model_definition) context.load() - context.plan("dev", start="2023-01-02", end="2023-01-02", auto_apply=True, no_prompts=True) + context.plan( + "dev", + start="2023-01-02", + end="2023-01-02", + auto_apply=True, + no_prompts=True, + ) snapshot = context.get_snapshot(f"test_schema.{model_name}") assert snapshot fingerprint = snapshot.fingerprint.to_version() - table_name = f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint}__dev" + table_name = ( + f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint}__dev" + ) data_change_grants = engine_adapter._get_current_grants_config( exp.to_table(table_name, dialect=engine_adapter.dialect) ) @@ -1133,19 +1188,32 @@ def test_grants_plan_scd_type_2_models( ); SELECT 1 as id, 'grant_changed_data' as name, CURRENT_TIMESTAMP as updated_at """ - (tmp_path / "models" / f"{model_name}.sql").write_text(grant_change_model_definition) + (tmp_path / "models" / f"{model_name}.sql").write_text( + grant_change_model_definition + ) context.load() - context.plan("dev", start="2023-01-03", end="2023-01-03", auto_apply=True, no_prompts=True) + context.plan( + "dev", + start="2023-01-03", + end="2023-01-03", + auto_apply=True, + no_prompts=True, + ) snapshot = context.get_snapshot(f"test_schema.{model_name}") assert snapshot fingerprint = snapshot.fingerprint.to_version() - table_name = f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint}__dev" + table_name = ( + f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint}__dev" + ) final_grants = engine_adapter._get_current_grants_config( exp.to_table(table_name, dialect=engine_adapter.dialect) ) - expected_select_users = {roles["reader"]["username"], roles["analyst"]["username"]} + expected_select_users = { + roles["reader"]["username"], + roles["analyst"]["username"], + } assert set(final_grants.get("SELECT", [])) == expected_select_users assert final_grants.get("INSERT", []) == [roles["writer"]["username"]] assert final_grants.get("UPDATE", []) == [roles["analyst"]["username"]] @@ -1215,9 +1283,7 @@ def test_grants_plan_scd_type_2_with_vde_dev_only( snapshot = context.get_snapshot(f"test_schema.{model_name}") assert snapshot fingerprint_version = snapshot.fingerprint.to_version() - dev_physical_table_name = ( - f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint_version}__dev" - ) + dev_physical_table_name = f"sqlmesh__test_schema.test_schema__{model_name}__{fingerprint_version}__dev" dev_physical_grants = engine_adapter._get_current_grants_config( exp.to_table(dev_physical_table_name, dialect=engine_adapter.dialect) diff --git a/tests/core/engine_adapter/integration/test_integration_redshift.py b/tests/core/engine_adapter/integration/test_integration_redshift.py index be5a47e714..9b7117dbb8 100644 --- a/tests/core/engine_adapter/integration/test_integration_redshift.py +++ b/tests/core/engine_adapter/integration/test_integration_redshift.py @@ -1,22 +1,22 @@ import typing as t + import pytest from pytest import FixtureRequest -from tests.core.engine_adapter.integration import TestContext -from sqlmesh.core.engine_adapter.redshift import RedshiftEngineAdapter from sqlglot import exp -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) +from sqlmesh.core.engine_adapter.redshift import RedshiftEngineAdapter +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) @pytest.fixture(params=list(generate_pytest_params(ENGINES_BY_NAME["redshift"]))) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -43,16 +43,28 @@ def test_columns(ctx: TestContext): sql += ( ", ".join( f"{col.replace(' ', '_')}10 {col}(10)" - for col in [*col_strings["char"], *col_strings["varchar"], *col_strings["varbinary"]] + for col in [ + *col_strings["char"], + *col_strings["varchar"], + *col_strings["varbinary"], + ] ) + ", " ) # bare types that should have their default lengths of 1 added by columns() - sql += ", ".join(f"{col.replace(' ', '_')}1 {col}" for col in col_strings["char"]) + ", " + sql += ( + ", ".join(f"{col.replace(' ', '_')}1 {col}" for col in col_strings["char"]) + + ", " + ) # bare types that should have their default lengths of 256 added by columns() - sql += ", ".join(f"{col.replace(' ', '_')}256 {col}" for col in col_strings["varchar"]) + ", " sql += ( - ", ".join(f"{col.replace(' ', '_')}172 {col}(17, 2)" for col in col_strings["decimal"]) + ", ".join(f"{col.replace(' ', '_')}256 {col}" for col in col_strings["varchar"]) + + ", " + ) + sql += ( + ", ".join( + f"{col.replace(' ', '_')}172 {col}(17, 2)" for col in col_strings["decimal"] + ) + ")" ) @@ -61,24 +73,36 @@ def test_columns(ctx: TestContext): # columns to types cols_to_types = { - f"{col.replace(' ', '_')}10": exp.DataType.build(f"{col}(10)", dialect=ctx.dialect) - for col in [*col_strings["char"], *col_strings["varchar"], *col_strings["varbinary"]] + f"{col.replace(' ', '_')}10": exp.DataType.build( + f"{col}(10)", dialect=ctx.dialect + ) + for col in [ + *col_strings["char"], + *col_strings["varchar"], + *col_strings["varbinary"], + ] } cols_to_types.update( { - f"{col.replace(' ', '_')}1": exp.DataType.build(f"{col}(1)", dialect=ctx.dialect) + f"{col.replace(' ', '_')}1": exp.DataType.build( + f"{col}(1)", dialect=ctx.dialect + ) for col in col_strings["char"] } ) cols_to_types.update( { - f"{col.replace(' ', '_')}256": exp.DataType.build(f"{col}(256)", dialect=ctx.dialect) + f"{col.replace(' ', '_')}256": exp.DataType.build( + f"{col}(256)", dialect=ctx.dialect + ) for col in col_strings["varchar"] } ) cols_to_types.update( { - f"{col.replace(' ', '_')}172": exp.DataType.build(f"{col}(17, 2)", dialect=ctx.dialect) + f"{col.replace(' ', '_')}172": exp.DataType.build( + f"{col}(17, 2)", dialect=ctx.dialect + ) for col in col_strings["decimal"] } ) @@ -95,13 +119,22 @@ def test_columns(ctx: TestContext): for col in ctx.engine_adapter._default_precision_to_max( # type: ignore {k: columns[k] for k in max_cols} ).values() - ] == ["CHAR(max)", "CHAR(max)", "CHAR(max)", "VARCHAR(max)", "VARCHAR(max)", "VARCHAR(max)"] + ] == [ + "CHAR(max)", + "CHAR(max)", + "CHAR(max)", + "VARCHAR(max)", + "VARCHAR(max)", + "VARCHAR(max)", + ] def test_fetch_native_df_respects_case_sensitivity(ctx: TestContext): adapter = ctx.engine_adapter adapter.execute("SET enable_case_sensitive_identifier TO true") - assert adapter.fetchdf('WITH t AS (SELECT 1 AS "C", 2 AS "c") SELECT * FROM t').to_dict() == { + assert adapter.fetchdf( + 'WITH t AS (SELECT 1 AS "C", 2 AS "c") SELECT * FROM t' + ).to_dict() == { "C": {0: 1}, "c": {0: 2}, } diff --git a/tests/core/engine_adapter/integration/test_integration_risingwave.py b/tests/core/engine_adapter/integration/test_integration_risingwave.py index 76b3d20a7c..544ac2a18d 100644 --- a/tests/core/engine_adapter/integration/test_integration_risingwave.py +++ b/tests/core/engine_adapter/integration/test_integration_risingwave.py @@ -1,14 +1,14 @@ import typing as t + import pytest -from sqlglot import exp from pytest import FixtureRequest +from sqlglot import exp + from sqlmesh.core.engine_adapter import RisingwaveEngineAdapter -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) @pytest.fixture(params=list(generate_pytest_params(ENGINES_BY_NAME["risingwave"]))) diff --git a/tests/core/engine_adapter/integration/test_integration_snowflake.py b/tests/core/engine_adapter/integration/test_integration_snowflake.py index 7f3c38be46..93733f986a 100644 --- a/tests/core/engine_adapter/integration/test_integration_snowflake.py +++ b/tests/core/engine_adapter/integration/test_integration_snowflake.py @@ -1,39 +1,42 @@ -import pytest import typing as t from datetime import datetime from pathlib import Path + +import pytest from pytest import FixtureRequest from pytest_mock import MockerFixture - -import sqlmesh.core.dialect as d from sqlglot import exp -from sqlmesh import Config, ExecutionContext, model from sqlglot.helper import seq_get from sqlglot.optimizer.qualify_columns import quote_identifiers + +import sqlmesh.core.dialect as d +from sqlmesh import Config, ExecutionContext, model from sqlmesh.core.config import ModelDefaultsConfig from sqlmesh.core.engine_adapter import SnowflakeEngineAdapter from sqlmesh.core.engine_adapter.shared import DataObject from sqlmesh.core.model import ModelKindName, SqlModel, load_sql_based_model from sqlmesh.core.plan import Plan from sqlmesh.core.snapshot import SnapshotId, SnapshotIdBatch -from sqlmesh.core.snapshot.execution_tracker import ( - QueryExecutionContext, - QueryExecutionTracker, -) -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) +from sqlmesh.core.snapshot.execution_tracker import (QueryExecutionContext, + QueryExecutionTracker) +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) @pytest.fixture( - params=list(generate_pytest_params(ENGINES_BY_NAME["snowflake"], show_variant_in_test_id=False)) + params=list( + generate_pytest_params( + ENGINES_BY_NAME["snowflake"], show_variant_in_test_id=False + ) + ) ) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -51,17 +54,23 @@ def test_get_alter_expressions_includes_clustering( clustered_differently_table = ctx.table("clustered_differently_table") normal_table = ctx.table("normal_table") - engine_adapter.execute(f"CREATE TABLE {clustered_table} (c1 int, c2 timestamp) CLUSTER BY (c1)") + engine_adapter.execute( + f"CREATE TABLE {clustered_table} (c1 int, c2 timestamp) CLUSTER BY (c1)" + ) engine_adapter.execute( f"CREATE TABLE {clustered_differently_table} (c1 int, c2 timestamp) CLUSTER BY (c1, to_date(c2))" ) engine_adapter.execute(f"CREATE TABLE {normal_table} (c1 int, c2 timestamp)") assert len(engine_adapter.get_alter_operations(normal_table, normal_table)) == 0 - assert len(engine_adapter.get_alter_operations(clustered_table, clustered_table)) == 0 + assert ( + len(engine_adapter.get_alter_operations(clustered_table, clustered_table)) == 0 + ) # alter table drop clustered - clustered_to_normal = engine_adapter.get_alter_operations(clustered_table, normal_table) + clustered_to_normal = engine_adapter.get_alter_operations( + clustered_table, normal_table + ) assert len(clustered_to_normal) == 1 assert ( clustered_to_normal[0].expression.sql(dialect=ctx.dialect) @@ -69,7 +78,9 @@ def test_get_alter_expressions_includes_clustering( ) # alter table add clustered - normal_to_clustered = engine_adapter.get_alter_operations(normal_table, clustered_table) + normal_to_clustered = engine_adapter.get_alter_operations( + normal_table, clustered_table + ) assert len(normal_to_clustered) == 1 assert ( normal_to_clustered[0].expression.sql(dialect=ctx.dialect) @@ -108,9 +119,7 @@ def _create_model(**kwargs: t.Any) -> SqlModel: extra_props = "\n".join([f"{k} {v}," for k, v in kwargs.items()]) return t.cast( SqlModel, - load_sql_based_model( - d.parse( - f""" + load_sql_based_model(d.parse(f""" MODEL ( name {model_name}, kind INCREMENTAL_BY_TIME_RANGE ( @@ -123,13 +132,13 @@ def _create_model(**kwargs: t.Any) -> SqlModel: ); select 1 as ID, current_timestamp() as PARTITIONDATE - """ - ) - ), + """)), ) def _get_data_object(table: exp.Table) -> DataObject: - data_object = seq_get(engine_adapter.get_data_objects(table.db, {table.name}), 0) + data_object = seq_get( + engine_adapter.get_data_objects(table.db, {table.name}), 0 + ) if not data_object: raise ValueError(f"Expected metadata for {table}") return data_object @@ -186,7 +195,9 @@ def _get_data_object(table: exp.Table) -> DataObject: assert not metadata.is_clustered -@pytest.mark.skip(reason="External volume LIST privileges not configured for CI test databases") +@pytest.mark.skip( + reason="External volume LIST privileges not configured for CI test databases" +) def test_create_iceberg_table(ctx: TestContext) -> None: # Note: this test relies on a default Catalog and External Volume being configured in Snowflake # ref: https://docs.snowflake.com/en/user-guide/tables-iceberg-configure-catalog-integration#set-a-default-catalog-at-the-account-database-or-schema-level @@ -197,8 +208,7 @@ def test_create_iceberg_table(ctx: TestContext) -> None: managed_model_name = ctx.table("TEST_DYNAMIC") sqlmesh = ctx.create_context() - model = load_sql_based_model( - d.parse(f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name {model_name}, kind FULL, @@ -207,11 +217,9 @@ def test_create_iceberg_table(ctx: TestContext) -> None: ); select 1 as "ID", 'foo' as "NAME"; - """) - ) + """)) - managed_model = load_sql_based_model( - d.parse(f""" + managed_model = load_sql_based_model(d.parse(f""" MODEL ( name {managed_model_name}, kind MANAGED, @@ -223,8 +231,7 @@ def test_create_iceberg_table(ctx: TestContext) -> None: ); select "ID", "NAME" from {model_name}; - """) - ) + """)) sqlmesh.upsert_model(model) sqlmesh.upsert_model(managed_model) @@ -254,7 +261,9 @@ def test_snowpark_concurrency(ctx: TestContext) -> None: ) def execute(context: ExecutionContext, start: datetime, **kwargs) -> DataFrame: if snowpark := context.snowpark: - return snowpark.create_dataframe([(start.day, start.date())], schema=["id", "ds"]) + return snowpark.create_dataframe( + [(start.day, start.date())], schema=["id", "ds"] + ) raise ValueError("Snowpark not present!") @@ -289,7 +298,9 @@ def test_create_drop_catalog(ctx: TestContext, engine_adapter: SnowflakeEngineAd ctx.create_catalog( non_sqlmesh_managed_catalog ) # create via TestContext so the sqlmesh_managed comment doesnt get added - ctx._catalogs.append(sqlmesh_managed_catalog) # so it still gets cleaned up if the test fails + ctx._catalogs.append( + sqlmesh_managed_catalog + ) # so it still gets cleaned up if the test fails engine_adapter.create_catalog( sqlmesh_managed_catalog @@ -304,14 +315,22 @@ def fetch_database_names() -> t.Set[str]: ) } - assert fetch_database_names() == {non_sqlmesh_managed_catalog, sqlmesh_managed_catalog} + assert fetch_database_names() == { + non_sqlmesh_managed_catalog, + sqlmesh_managed_catalog, + } engine_adapter.drop_catalog( non_sqlmesh_managed_catalog ) # no-op: catalog is not SQLMesh-managed - assert fetch_database_names() == {non_sqlmesh_managed_catalog, sqlmesh_managed_catalog} + assert fetch_database_names() == { + non_sqlmesh_managed_catalog, + sqlmesh_managed_catalog, + } - engine_adapter.drop_catalog(sqlmesh_managed_catalog) # works, catalog is SQLMesh-managed + engine_adapter.drop_catalog( + sqlmesh_managed_catalog + ) # works, catalog is SQLMesh-managed assert fetch_database_names() == {non_sqlmesh_managed_catalog} @@ -358,7 +377,9 @@ def test_unit_test(tmp_path: Path, ctx: TestContext): - c: 1 """ - (models_path / "dummy_model.sql").write_text(f"MODEL (name s.dummy); SELECT c FROM s.src_table") + (models_path / "dummy_model.sql").write_text( + f"MODEL (name s.dummy); SELECT c FROM s.src_table" + ) (tests_path / "test_dummy_model.yaml").write_text(test_payload) def _config_mutator(gateway_name: str, config: Config): diff --git a/tests/core/engine_adapter/integration/test_integration_starrocks.py b/tests/core/engine_adapter/integration/test_integration_starrocks.py index 64f0776c9d..e2df44c8e6 100644 --- a/tests/core/engine_adapter/integration/test_integration_starrocks.py +++ b/tests/core/engine_adapter/integration/test_integration_starrocks.py @@ -28,11 +28,10 @@ import pytest from sqlglot import exp +import sqlmesh.core.dialect as d from sqlmesh.core.engine_adapter.starrocks import StarRocksEngineAdapter -from sqlmesh.core.model.definition import load_sql_based_model, SqlModel +from sqlmesh.core.model.definition import SqlModel, load_sql_based_model from sqlmesh.utils.errors import SQLMeshError -import sqlmesh.core.dialect as d - from tests.core.engine_adapter.integration import TestContext # Mark as docker test (can also run against local StarRocks) @@ -48,7 +47,9 @@ def _load_sql_model(model_sql: str) -> SqlModel: return t.cast(SqlModel, load_sql_based_model(expressions)) -def _materialized_properties_from_model(model: SqlModel) -> t.Optional[t.Dict[str, t.Any]]: +def _materialized_properties_from_model( + model: SqlModel, +) -> t.Optional[t.Dict[str, t.Any]]: props: t.Dict[str, t.Any] = {} if model.partitioned_by: props["partitioned_by"] = model.partitioned_by @@ -183,7 +184,9 @@ def init_test_integration_env(starrocks_adapter: StarRocksEngineAdapter) -> None def _get_config_value(name: str) -> t.Optional[str]: try: - row = starrocks_adapter.fetchone(f"ADMIN SHOW FRONTEND CONFIG LIKE '{name}'") + row = starrocks_adapter.fetchone( + f"ADMIN SHOW FRONTEND CONFIG LIKE '{name}'" + ) except Exception as e: # pragma: no cover - defensive for older SR versions logger.warning("Skipping config lookup %s: %s", name, e) return None @@ -218,14 +221,18 @@ def _get_config_value(name: str) -> t.Optional[str]: return try: - starrocks_adapter.execute('ADMIN SET FRONTEND CONFIG ("default_replication_num" = "1")') + starrocks_adapter.execute( + 'ADMIN SET FRONTEND CONFIG ("default_replication_num" = "1")' + ) logger.info( "Set default_replication_num=1 for shared_nothing cluster with %s backends (was %s)", be_count, current_replication, ) except Exception as e: # pragma: no cover - do not break tests if lacking privilege - logger.warning("Failed to set default_replication_num for shared_nothing cluster: %s", e) + logger.warning( + "Failed to set default_replication_num for shared_nothing cluster: %s", e + ) class TestBasicOperations: @@ -236,7 +243,9 @@ class TestBasicOperations: This allows running individual tests and clear failure reporting. """ - def test_create_drop_schema(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): + def test_create_drop_schema( + self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter + ): """Test CREATE DATABASE and DROP DATABASE (TestContext version).""" db_name = ctx.schema("sr_test_create_drop_db") @@ -255,7 +264,9 @@ def test_create_drop_schema(self, ctx: TestContext, engine_adapter: StarRocksEng ) assert dropped_result is None, "DROP DATABASE failed" - def test_create_drop_table(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): + def test_create_drop_table( + self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter + ): """Test CREATE TABLE and DROP TABLE (TestContext version).""" table = ctx.table("sr_test_table") @@ -317,17 +328,20 @@ def test_create_table_like_preserves_metadata_and_copies_no_data( # Like should not copy data. src_count = fetchone_or_fail( - engine_adapter, f"SELECT COUNT(*) FROM {source.sql(dialect=ctx.dialect, identify=True)}" + engine_adapter, + f"SELECT COUNT(*) FROM {source.sql(dialect=ctx.dialect, identify=True)}", )[0] tgt_count = fetchone_or_fail( - engine_adapter, f"SELECT COUNT(*) FROM {target.sql(dialect=ctx.dialect, identify=True)}" + engine_adapter, + f"SELECT COUNT(*) FROM {target.sql(dialect=ctx.dialect, identify=True)}", )[0] assert src_count == 2 assert tgt_count == 0 # Like should preserve key metadata (engine-defined behavior). ddl = fetchone_or_fail( - engine_adapter, f"SHOW CREATE TABLE {target.sql(dialect=ctx.dialect, identify=True)}" + engine_adapter, + f"SHOW CREATE TABLE {target.sql(dialect=ctx.dialect, identify=True)}", )[1] ddl_upper = ddl.upper() assert "PRIMARY KEY" in ddl_upper @@ -373,7 +387,9 @@ def test_delete(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): count = fetchone_or_fail(engine_adapter, f"SELECT COUNT(*) FROM {table_sql}") assert count[0] == 1, "DELETE failed" - def test_rename_table(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): + def test_rename_table( + self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter + ): """Test RENAME TABLE operation (TestContext version).""" old_table = ctx.table("old_table") new_table = ctx.table("new_table") @@ -389,7 +405,9 @@ def test_rename_table(self, ctx: TestContext, engine_adapter: StarRocksEngineAda }, ) - engine_adapter.execute(f"INSERT INTO {old_table_sql} (id, name) VALUES (1, 'Test')") + engine_adapter.execute( + f"INSERT INTO {old_table_sql} (id, name) VALUES (1, 'Test')" + ) engine_adapter.rename_table(old_table, new_table) db_name = old_table.db @@ -408,10 +426,14 @@ def test_rename_table(self, ctx: TestContext, engine_adapter: StarRocksEngineAda ) assert new_exists is not None, "New table should exist after rename" - count = fetchone_or_fail(engine_adapter, f"SELECT COUNT(*) FROM {new_table_sql}") + count = fetchone_or_fail( + engine_adapter, f"SELECT COUNT(*) FROM {new_table_sql}" + ) assert count[0] == 1, "Data should be preserved after rename" - def test_create_index(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): + def test_create_index( + self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter + ): """Test CREATE INDEX operation (skipped for StarRocks) (TestContext version).""" table = ctx.table("sr_test_table") table_sql = table.sql(dialect=ctx.dialect, identify=True) @@ -428,9 +450,13 @@ def test_create_index(self, ctx: TestContext, engine_adapter: StarRocksEngineAda engine_adapter.create_index(table, "idx_name", ("name",)) count = fetchone_or_fail(engine_adapter, f"SELECT COUNT(*) FROM {table_sql}") - assert count[0] >= 0, "Table should still be functional after skipped index creation" + assert ( + count[0] >= 0 + ), "Table should still be functional after skipped index creation" - def test_create_drop_view(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): + def test_create_drop_view( + self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter + ): """Test CREATE VIEW and DROP VIEW (TestContext version).""" table = ctx.table("sr_test_table") view = ctx.table("sr_test_view") @@ -530,7 +556,9 @@ def test_create_view_replace_flag( "name": exp.DataType.build("VARCHAR(100)"), }, ) - engine_adapter.execute(f"INSERT INTO {source_sql_ident} (id, name) VALUES (1, 'A')") + engine_adapter.execute( + f"INSERT INTO {source_sql_ident} (id, name) VALUES (1, 'A')" + ) model_sql = f""" MODEL ( @@ -583,14 +611,12 @@ def _create_sales_source_table( primary_key=("order_id", "event_date"), partitioned_by="event_date", ) - engine_adapter.execute( - f""" + engine_adapter.execute(f""" INSERT INTO {table_sql} (order_id, customer_id, event_date, amount, region) VALUES (1, 1001, '2024-01-01', 10.50, 'us'), (2, 1002, '2024-01-02', 20.75, 'eu') - """ - ) + """) return table_sql def test_materialized_view_combo_with_materialized_properties( @@ -652,13 +678,18 @@ def test_materialized_view_combo_with_materialized_properties( column_descriptions=model.column_descriptions, ) - ddl = fetchone_or_fail(engine_adapter, f"SHOW CREATE MATERIALIZED VIEW {mv_sql}")[1] + ddl = fetchone_or_fail( + engine_adapter, f"SHOW CREATE MATERIALIZED VIEW {mv_sql}" + )[1] logger.debug(f"mv ddl: {ddl}") ddl_upper = normalize_sql(ddl).upper() # StarRocks renders a scheduled async refresh (ASYNC START ... EVERY ...) as # "REFRESH DEFERRED SCHEDULE START ... EVERY ..." in SHOW CREATE (newer versions); # older versions render it as "REFRESH DEFERRED ASYNC START ...". Accept either. - assert "REFRESH DEFERRED SCHEDULE" in ddl_upper or "REFRESH DEFERRED ASYNC" in ddl_upper + assert ( + "REFRESH DEFERRED SCHEDULE" in ddl_upper + or "REFRESH DEFERRED ASYNC" in ddl_upper + ) assert ( "START('2025-01-01 00:00:00')EVERY(INTERVAL 5 MINUTE)" in ddl_upper or 'START("2025-01-01 00:00:00")EVERY(INTERVAL 5 MINUTE)' in ddl_upper @@ -732,7 +763,9 @@ def test_materialized_view_combo_all_properties_block( column_descriptions=model.column_descriptions, ) - ddl = fetchone_or_fail(engine_adapter, f"SHOW CREATE MATERIALIZED VIEW {mv_sql}")[1] + ddl = fetchone_or_fail( + engine_adapter, f"SHOW CREATE MATERIALIZED VIEW {mv_sql}" + )[1] ddl_upper = normalize_sql(ddl).upper() assert "REFRESH MANUAL" in ddl_upper assert "PARTITION P202401" not in ddl_upper # ignored when MV @@ -847,7 +880,9 @@ def test_table_and_column_comments( assert column_comments["id"] == "User ID" assert column_comments["name"] == "User name" - def test_multiple_data_types(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): + def test_multiple_data_types( + self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter + ): """ Test basic data types support. @@ -893,8 +928,7 @@ def test_multiple_data_types(self, ctx: TestContext, engine_adapter: StarRocksEn assert len(columns) == 14, f"Expected 14 columns, got {len(columns)}" # Test data insertion with various types - engine_adapter.execute( - f""" + engine_adapter.execute(f""" INSERT INTO {table_sql} (col_tinyint, col_smallint, col_int, col_bigint, col_float, col_double, col_decimal, col_char, col_varchar, col_string, col_date, col_datetime, col_boolean, col_json) @@ -902,8 +936,7 @@ def test_multiple_data_types(self, ctx: TestContext, engine_adapter: StarRocksEn (127, 32767, 2147483647, 9223372036854775807, 3.14, 3.141592653589793, 12345.67, 'test', 'test varchar', 'test string', '2024-01-01', '2024-01-01 12:00:00', true, '{{"key": "value"}}') - """ - ) + """) # Verify insertion count = fetchone_or_fail(engine_adapter, f"SELECT COUNT(*) FROM {table_sql}") @@ -917,7 +950,9 @@ def test_multiple_data_types(self, ctx: TestContext, engine_adapter: StarRocksEn assert result[1] == "test varchar" # @pytest.mark.skip(reason="Complex types (ARRAY/MAP/STRUCT) may not be fully supported yet") - def test_complex_data_types(self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter): + def test_complex_data_types( + self, ctx: TestContext, engine_adapter: StarRocksEngineAdapter + ): """ Test complex and nested data types support (ARRAY, MAP, STRUCT). @@ -956,7 +991,9 @@ def test_complex_data_types(self, ctx: TestContext, engine_adapter: StarRocksEng "STRUCT, metadata MAP>" ), # ARRAY of STRUCT - "col_array_of_struct": exp.DataType.build("ARRAY>"), + "col_array_of_struct": exp.DataType.build( + "ARRAY>" + ), # Deep nesting: MAP with ARRAY of STRUCT "col_deep_nested": exp.DataType.build( "MAP>>" @@ -973,8 +1010,7 @@ def test_complex_data_types(self, ctx: TestContext, engine_adapter: StarRocksEng assert len(columns) == 9, f"Expected 9 columns, got {len(columns)}" # Test data insertion with nested types - engine_adapter.execute( - f""" + engine_adapter.execute(f""" INSERT INTO {table_sql} (id, col_array_simple, col_map_simple, col_struct_simple, col_array_nested, col_map_nested, col_struct_nested, @@ -990,8 +1026,7 @@ def test_complex_data_types(self, ctx: TestContext, engine_adapter: StarRocksEng [row(1,'Alice'), row(2,'Bob')], map{{'group1':[row(10,'field_a'), row(20,'field_b')]}} ) - """ - ) + """) # Verify insertion count = fetchone_or_fail(engine_adapter, f"SELECT COUNT(*) FROM {table_sql}") @@ -999,7 +1034,8 @@ def test_complex_data_types(self, ctx: TestContext, engine_adapter: StarRocksEng # Verify data retrieval for simple types result = fetchone_or_fail( - engine_adapter, f"SELECT col_array_simple, col_struct_simple FROM {table_sql}" + engine_adapter, + f"SELECT col_array_simple, col_struct_simple FROM {table_sql}", ) assert result is not None, "Failed to retrieve complex type data" @@ -1118,7 +1154,9 @@ def test_e2e_model_parameters(self, starrocks_adapter: StarRocksEngineAdapter): params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 1 DDL:\n{ddl}") @@ -1133,20 +1171,20 @@ def test_e2e_model_parameters(self, starrocks_adapter: StarRocksEngineAdapter): part_match = re.search(r"PARTITION BY\s*(\((?:[^()]|\([^()]*\))*\))", ddl) assert part_match, "PARTITION BY clause not found" part_cols = part_match.group(1) - assert "from_unixtime" in part_cols and "ts" in part_cols, ( - f"Expected partition expression with from_unixtime(ts), got {part_cols}" - ) - assert "region" in part_cols, ( - f"Expected 'region' partition column in PARTITION BY, got {part_cols}" - ) + assert ( + "from_unixtime" in part_cols and "ts" in part_cols + ), f"Expected partition expression with from_unixtime(ts), got {part_cols}" + assert ( + "region" in part_cols + ), f"Expected 'region' partition column in PARTITION BY, got {part_cols}" # Verify ORDER BY from clustered_by order_match = re.search(r"ORDER BY\s*\(([^)]+)\)", ddl) assert order_match, "ORDER BY clause not found" order_cols = order_match.group(1) - assert "order_id" in order_cols and "customer_id" in order_cols, ( - f"Expected ORDER BY (order_id, customer_id), got {order_cols}" - ) + assert ( + "order_id" in order_cols and "customer_id" in order_cols + ), f"Expected ORDER BY (order_id, customer_id), got {order_cols}" finally: starrocks_adapter.drop_schema(db_name, ignore_if_not_exists=True) @@ -1156,7 +1194,9 @@ def test_e2e_model_parameters(self, starrocks_adapter: StarRocksEngineAdapter): # Covers: primary_key (tuple), distributed_by (string multi-col), order_by (tuple), generic props # ======================================== - def test_e2e_physical_properties_core(self, starrocks_adapter: StarRocksEngineAdapter): + def test_e2e_physical_properties_core( + self, starrocks_adapter: StarRocksEngineAdapter + ): """ Test Case 2: Core physical_properties. @@ -1197,7 +1237,9 @@ def test_e2e_physical_properties_core(self, starrocks_adapter: StarRocksEngineAd params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 2 DDL:\n{ddl}") @@ -1213,15 +1255,17 @@ def test_e2e_physical_properties_core(self, starrocks_adapter: StarRocksEngineAd dist_match = re.search(r"DISTRIBUTED BY HASH\s*\(([^)]+)\)", ddl) assert dist_match, "DISTRIBUTED BY HASH clause not found" dist_cols = dist_match.group(1) - assert "customer_id" in dist_cols and "region" in dist_cols, ( - f"Expected HASH(customer_id, region), got HASH({dist_cols})" - ) + assert ( + "customer_id" in dist_cols and "region" in dist_cols + ), f"Expected HASH(customer_id, region), got HASH({dist_cols})" assert "BUCKETS 16" in ddl # Verify ORDER BY order_match = re.search(r"ORDER BY\s*\(([^)]+)\)", ddl) assert order_match, "ORDER BY clause not found" - assert "order_id" in order_match.group(1) and "region" in order_match.group(1) + assert "order_id" in order_match.group(1) and "region" in order_match.group( + 1 + ) # assert "replication_num" not in ddl @@ -1233,7 +1277,9 @@ def test_e2e_physical_properties_core(self, starrocks_adapter: StarRocksEngineAd # Covers: primary_key = "id, dt" auto-conversion # ======================================== - def test_e2e_string_no_paren_auto_wrap(self, starrocks_adapter: StarRocksEngineAdapter): + def test_e2e_string_no_paren_auto_wrap( + self, starrocks_adapter: StarRocksEngineAdapter + ): """ Test Case 3: String form without parentheses auto-wrap. @@ -1266,7 +1312,9 @@ def test_e2e_string_no_paren_auto_wrap(self, starrocks_adapter: StarRocksEngineA params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 3 DDL:\n{ddl}") @@ -1276,16 +1324,16 @@ def test_e2e_string_no_paren_auto_wrap(self, starrocks_adapter: StarRocksEngineA pk_match = re.search(r"PRIMARY KEY\s*\(([^)]+)\)", ddl) assert pk_match, "PRIMARY KEY clause not found" pk_clause = pk_match.group(1) - assert "order_id" in pk_clause and "event_date" in pk_clause, ( - f"Expected both order_id and event_date in PRIMARY KEY, got {pk_clause}" - ) + assert ( + "order_id" in pk_clause and "event_date" in pk_clause + ), f"Expected both order_id and event_date in PRIMARY KEY, got {pk_clause}" # Verify distributed_by with exact columns dist_match = re.search(r"DISTRIBUTED BY HASH\s*\(([^)]+)\)", ddl) assert dist_match, "DISTRIBUTED BY HASH clause not found" - assert "order_id" in dist_match.group(1), ( - f"Expected HASH(order_id), got HASH({dist_match.group(1)})" - ) + assert "order_id" in dist_match.group( + 1 + ), f"Expected HASH(order_id), got HASH({dist_match.group(1)})" assert "BUCKETS 10" in ddl finally: @@ -1296,7 +1344,9 @@ def test_e2e_string_no_paren_auto_wrap(self, starrocks_adapter: StarRocksEngineA # Covers: kind=HASH (unquoted), kind=RANDOM # ======================================== - def test_e2e_distribution_structured_hash(self, starrocks_adapter: StarRocksEngineAdapter): + def test_e2e_distribution_structured_hash( + self, starrocks_adapter: StarRocksEngineAdapter + ): """Test Case 4A: Structured HASH distribution with unquoted kind.""" db_name = "sr_e2e_dist_hash_db" table_name = f"{db_name}.sr_dist_hash_table" @@ -1324,7 +1374,9 @@ def test_e2e_distribution_structured_hash(self, starrocks_adapter: StarRocksEngi params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 4A DDL:\n{ddl}") @@ -1334,13 +1386,17 @@ def test_e2e_distribution_structured_hash(self, starrocks_adapter: StarRocksEngi assert "DISTRIBUTED BY HASH" in ddl dist_match = re.search(r"DISTRIBUTED BY HASH\s*\(([^)]+)\)", ddl) assert dist_match, "DISTRIBUTED BY HASH clause not found" - assert "customer_id" in dist_match.group(1) and "region" in dist_match.group(1) + assert "customer_id" in dist_match.group( + 1 + ) and "region" in dist_match.group(1) assert "BUCKETS 16" in ddl finally: starrocks_adapter.drop_schema(db_name, ignore_if_not_exists=True) - def test_e2e_distribution_structured_random(self, starrocks_adapter: StarRocksEngineAdapter): + def test_e2e_distribution_structured_random( + self, starrocks_adapter: StarRocksEngineAdapter + ): """Test Case 4B: Structured RANDOM distribution.""" db_name = "sr_e2e_dist_random_db" table_name = f"{db_name}.sr_dist_random_table" @@ -1369,7 +1425,9 @@ def test_e2e_distribution_structured_random(self, starrocks_adapter: StarRocksEn params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 4B DDL:\n{ddl}") @@ -1420,7 +1478,9 @@ def test_e2e_partition_range(self, starrocks_adapter: StarRocksEngineAdapter): params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 5 DDL:\n{ddl}") @@ -1477,7 +1537,9 @@ def test_e2e_partition_list(self, starrocks_adapter: StarRocksEngineAdapter): params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 6 DDL:\n{ddl}") @@ -1538,7 +1600,9 @@ def test_e2e_partition_expression_for_table( params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - ddl = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}")[1] + ddl = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + )[1] logger.info(f"Case 6B DDL:\n{ddl}") ddl_upper = ddl.upper() assert "PARTITION BY" in ddl_upper @@ -1553,7 +1617,8 @@ def test_e2e_partition_expression_for_table( assert "REGION" in after if "FROM_UNIXTIME" in partition_expr.upper(): assert "FROM_UNIXTIME" in after or ( - "FROM_UNIXTIME" in before and "__GENERATED_PARTITION_COLUMN" in after + "FROM_UNIXTIME" in before + and "__GENERATED_PARTITION_COLUMN" in after ) if "DATE_TRUNC" in partition_expr.upper(): assert "DATE_TRUNC" in after @@ -1596,12 +1661,10 @@ def test_e2e_partition_expression_for_mv( "partitioned_by": partition_clause, }, ) - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" INSERT INTO {src_table} (id, ts, event_date, region) VALUES (1, 1700000000, '2024-01-01', 'us') - """ - ) + """) model_sql = f""" MODEL ( @@ -1641,14 +1704,16 @@ def test_e2e_partition_expression_for_mv( view_properties=model.physical_properties, ) - ddl = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE MATERIALIZED VIEW {mv_table}")[ - 1 - ] + ddl = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE MATERIALIZED VIEW {mv_table}" + )[1] logger.info(f"Case 6B DDL:\n{ddl}") ddl_upper = ddl.upper() assert "PARTITION BY" in ddl_upper after = ddl_upper.split("PARTITION BY", 1)[1].lstrip() - assert after.startswith("("), f"MV partition should keep parentheses, got: {after[:50]}" + assert after.startswith( + "(" + ), f"MV partition should keep parentheses, got: {after[:50]}" assert "REGION" in after assert ("DATE_TRUNC" in after) == has_func finally: @@ -1688,7 +1753,9 @@ def test_e2e_key_type_duplicate(self, starrocks_adapter: StarRocksEngineAdapter) params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 7A DDL:\n{ddl}") @@ -1697,9 +1764,9 @@ def test_e2e_key_type_duplicate(self, starrocks_adapter: StarRocksEngineAdapter) dup_match = re.search(r"DUPLICATE KEY\s*\(([^)]+)\)", ddl) assert dup_match, "DUPLICATE KEY clause not found" - assert "id" in dup_match.group(1) and "dt" in dup_match.group(1), ( - f"Expected DUPLICATE KEY(id, dt), got DUPLICATE KEY({dup_match.group(1)})" - ) + assert "id" in dup_match.group(1) and "dt" in dup_match.group( + 1 + ), f"Expected DUPLICATE KEY(id, dt), got DUPLICATE KEY({dup_match.group(1)})" finally: starrocks_adapter.drop_schema(db_name, ignore_if_not_exists=True) @@ -1733,7 +1800,9 @@ def test_e2e_key_type_unique(self, starrocks_adapter: StarRocksEngineAdapter): params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Case 7B DDL:\n{ddl}") @@ -1765,9 +1834,10 @@ def test_e2e_key_type_aggregate(self, starrocks_adapter: StarRocksEngineAdapter) SELECT * """ - from sqlmesh.utils.errors import SQLMeshError import pytest + from sqlmesh.utils.errors import SQLMeshError + try: starrocks_adapter.create_schema(db_name, ignore_if_exists=True) @@ -1820,7 +1890,9 @@ def test_e2e_comprehensive(self, starrocks_adapter: StarRocksEngineAdapter): params = self._parse_model_and_get_all_params(model_sql) starrocks_adapter.create_table(table_name, **params) - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Comprehensive DDL:\n{ddl}") @@ -1838,9 +1910,9 @@ def test_e2e_comprehensive(self, starrocks_adapter: StarRocksEngineAdapter): part_match = re.search(r"PARTITION BY[^(]*\(([^)]+)\)", ddl) assert part_match, "PARTITION BY clause not found" part_cols = part_match.group(1) - assert "event_date" in part_cols, ( - f"Expected event_date in PARTITION BY, got {part_cols}" - ) + assert ( + "event_date" in part_cols + ), f"Expected event_date in PARTITION BY, got {part_cols}" # Verify DISTRIBUTED BY assert "DISTRIBUTED BY HASH" in ddl @@ -1849,7 +1921,9 @@ def test_e2e_comprehensive(self, starrocks_adapter: StarRocksEngineAdapter): # Verify ORDER BY order_match = re.search(r"ORDER BY\s*\(([^)]+)\)", ddl) assert order_match, "ORDER BY clause not found" - assert "order_id" in order_match.group(1) and "event_date" in order_match.group(1) + assert "order_id" in order_match.group( + 1 + ) and "event_date" in order_match.group(1) # Verify PROPERTIES assert "replication_num" in ddl @@ -1876,7 +1950,9 @@ def test_e2e_comprehensive(self, starrocks_adapter: StarRocksEngineAdapter): # Tests single quotes vs double quotes in MODEL parsing # ======================================== - def test_e2e_quote_character_handling(self, starrocks_adapter: StarRocksEngineAdapter): + def test_e2e_quote_character_handling( + self, starrocks_adapter: StarRocksEngineAdapter + ): """ Test Case: Quote Character Handling (Single vs Double Quotes). @@ -1952,7 +2028,9 @@ def test_e2e_quote_character_handling(self, starrocks_adapter: StarRocksEngineAd starrocks_adapter.create_table(table_name, **params) # Verify via SHOW CREATE TABLE - show_create = fetchone_or_fail(starrocks_adapter, f"SHOW CREATE TABLE {table_name}") + show_create = fetchone_or_fail( + starrocks_adapter, f"SHOW CREATE TABLE {table_name}" + ) ddl = show_create[1] logger.info(f"Quote Handling Test DDL:\n{ddl}") @@ -2040,8 +2118,7 @@ def test_tables( # 1. PRIMARY KEY table # Note: StarRocks PRIMARY KEY tables support complex DELETE operations (BETWEEN, subqueries, etc.) pk_table = f"{db_name}.pk_table" - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE IF NOT EXISTS {pk_table} ( id INT, dt DATE, @@ -2049,8 +2126,7 @@ def test_tables( status STRING ) PRIMARY KEY (id, dt) DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Verify table creation result = fetchone_or_fail( starrocks_adapter, @@ -2062,8 +2138,7 @@ def test_tables( # 2. DUPLICATE KEY table dup_table = f"{db_name}.dup_table" - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE IF NOT EXISTS {dup_table} ( id INT, dt DATE, @@ -2071,8 +2146,7 @@ def test_tables( status STRING ) DUPLICATE KEY (id, dt) DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Verify table creation result = fetchone_or_fail( starrocks_adapter, @@ -2084,8 +2158,7 @@ def test_tables( # 3. UNIQUE KEY table unique_table = f"{db_name}.unique_table" - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE IF NOT EXISTS {unique_table} ( id INT, dt DATE, @@ -2093,8 +2166,7 @@ def test_tables( status STRING ) UNIQUE KEY (id, dt) DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Verify table creation result = fetchone_or_fail( starrocks_adapter, @@ -2150,19 +2222,19 @@ def test_insert_select_supported(self, starrocks_adapter: StarRocksEngineAdapter try: starrocks_adapter.create_schema(db_name, ignore_if_exists=True) - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE IF NOT EXISTS {table_name} ( id INT, name VARCHAR(100) ) PRIMARY KEY (id) DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) starrocks_adapter.execute( f"INSERT INTO {table_name} (id, name) VALUES (1, 'Alice'), (2, 'Bob')" ) - rows = starrocks_adapter.fetchall(f"SELECT id, name FROM {table_name} ORDER BY id") + rows = starrocks_adapter.fetchall( + f"SELECT id, name FROM {table_name} ORDER BY id" + ) assert list(rows) == [(1, "Alice"), (2, "Bob")], f"Data mismatch: {rows}" finally: starrocks_adapter.drop_schema(db_name, ignore_if_not_exists=True) @@ -2174,16 +2246,16 @@ def test_update_supported(self, starrocks_adapter: StarRocksEngineAdapter): try: starrocks_adapter.create_schema(db_name, ignore_if_exists=True) - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE IF NOT EXISTS {table_name} ( id INT, name VARCHAR(100) ) PRIMARY KEY (id) DISTRIBUTED BY HASH(id) BUCKETS 10 - """ + """) + starrocks_adapter.execute( + f"INSERT INTO {table_name} (id, name) VALUES (1, 'Alice')" ) - starrocks_adapter.execute(f"INSERT INTO {table_name} (id, name) VALUES (1, 'Alice')") starrocks_adapter.execute( f"UPDATE {table_name} SET name = 'Alice Updated' WHERE id = 1" ) @@ -2254,7 +2326,9 @@ def test_delete_supported_syntax( starrocks_adapter.execute(f"INSERT INTO {table_name} VALUES {test_data}") # Format delete clause (for subquery/using with table reference) - delete_sql = f"DELETE FROM {table_name} {delete_clause.format(table=table_name)}" + delete_sql = ( + f"DELETE FROM {table_name} {delete_clause.format(table=table_name)}" + ) # Debug: Log the SQL before execution logger.info(f"Executing DELETE SQL: {delete_sql}") @@ -2263,11 +2337,15 @@ def test_delete_supported_syntax( starrocks_adapter.execute(delete_sql) # Verify result - count = fetchone_or_fail(starrocks_adapter, f"SELECT COUNT(*) FROM {table_name}")[0] - logger.info(f"After DELETE: {count} rows remaining (expected {expected_remaining})") - assert count == expected_remaining, ( - f"Expected {expected_remaining} rows, got {count} for {table_type} with {delete_clause}" + count = fetchone_or_fail( + starrocks_adapter, f"SELECT COUNT(*) FROM {table_name}" + )[0] + logger.info( + f"After DELETE: {count} rows remaining (expected {expected_remaining})" ) + assert ( + count == expected_remaining + ), f"Expected {expected_remaining} rows, got {count} for {table_type} with {delete_clause}" # ==================== DELETE Operations - Failure Cases ==================== @@ -2318,7 +2396,9 @@ def test_delete_unsupported_syntax( Expected: DELETE fails with specific error message. """ table_name = test_tables[table_type] - delete_sql = f"DELETE FROM {table_name} {delete_clause.format(table=table_name)}" + delete_sql = ( + f"DELETE FROM {table_name} {delete_clause.format(table=table_name)}" + ) # This should raise an exception with pytest.raises(Exception) as exc_info: @@ -2328,9 +2408,9 @@ def test_delete_unsupported_syntax( import re error_msg = str(exc_info.value).lower() - assert re.search(error_pattern, error_msg), ( - f"Expected error pattern '{error_pattern}', got: {exc_info.value}" - ) + assert re.search( + error_pattern, error_msg + ), f"Expected error pattern '{error_pattern}', got: {exc_info.value}" # ==================== COMMENT Syntax Tests ==================== @@ -2375,20 +2455,20 @@ def test_comment_syntax_variants( try: starrocks_adapter.create_schema(db_name, ignore_if_exists=True) - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE {table_name} ( id INT, col1 INT ) DUPLICATE KEY (id) -- key columns can't be changed. DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Generate SQL based on template if "table" in comment_type: - sql = sql_template.format(table=table_name, comment=f"test {comment_type}") + sql = sql_template.format( + table=table_name, comment=f"test {comment_type}" + ) else: # column sql = sql_template.format( table=table_name, column="col1", comment=f"test {comment_type}" @@ -2405,9 +2485,9 @@ def test_comment_syntax_variants( f"SELECT TABLE_COMMENT FROM information_schema.TABLES " f"WHERE TABLE_SCHEMA = '{db_name}' AND TABLE_NAME = 'test_comment'", )[0] - assert f"test {comment_type}" in result, ( - f"Comment not set correctly for {comment_type}" - ) + assert ( + f"test {comment_type}" in result + ), f"Comment not set correctly for {comment_type}" else: # column result_row = fetchone_or_fail( starrocks_adapter, @@ -2417,9 +2497,9 @@ def test_comment_syntax_variants( ) logger.info(f"Column comment: {result_row}") result = result_row[1] - assert f"test {comment_type}" in result, ( - f"Comment not set correctly for {comment_type}" - ) + assert ( + f"test {comment_type}" in result + ), f"Comment not set correctly for {comment_type}" logger.info(f"✅ {comment_type}: SUPPORTED") @@ -2459,12 +2539,10 @@ def test_comment_quote_types( try: starrocks_adapter.create_schema(db_name, ignore_if_exists=True) - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE {table_name} (id INT) DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Build SQL with appropriate quotes if "single" in quote_type: @@ -2495,8 +2573,7 @@ def test_comment_in_create_table(self, starrocks_adapter: StarRocksEngineAdapter starrocks_adapter.create_schema(db_name, ignore_if_exists=True) # Create table with comments - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE {table_name} ( id INT COMMENT 'id column', name VARCHAR(100) COMMENT 'name column' @@ -2504,8 +2581,7 @@ def test_comment_in_create_table(self, starrocks_adapter: StarRocksEngineAdapter PRIMARY KEY (id) COMMENT 'test table' DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Verify table comment table_comment = fetchone_or_fail( @@ -2513,7 +2589,9 @@ def test_comment_in_create_table(self, starrocks_adapter: StarRocksEngineAdapter f"SELECT TABLE_COMMENT FROM information_schema.TABLES " f"WHERE TABLE_SCHEMA = '{db_name}' AND TABLE_NAME = 'test_create_comment'", )[0] - assert table_comment == "test table", f"Table comment mismatch: {table_comment}" + assert ( + table_comment == "test table" + ), f"Table comment mismatch: {table_comment}" # Verify column comments column_comments = {} @@ -2525,12 +2603,12 @@ def test_comment_in_create_table(self, starrocks_adapter: StarRocksEngineAdapter if col_comment: # Skip empty comments column_comments[col_name] = col_comment - assert column_comments.get("id") == "id column", ( - f"Column comment mismatch: {column_comments}" - ) - assert column_comments.get("name") == "name column", ( - f"Column comment mismatch: {column_comments}" - ) + assert ( + column_comments.get("id") == "id column" + ), f"Column comment mismatch: {column_comments}" + assert ( + column_comments.get("name") == "name column" + ), f"Column comment mismatch: {column_comments}" finally: starrocks_adapter.drop_schema(db_name, ignore_if_not_exists=True) @@ -2548,7 +2626,9 @@ class TestCommentMethods: - View comments (depending on COMMENT_CREATION_VIEW) """ - def test_build_create_comment_table_exp(self, starrocks_adapter: StarRocksEngineAdapter): + def test_build_create_comment_table_exp( + self, starrocks_adapter: StarRocksEngineAdapter + ): """ Test _build_create_comment_table_exp generates correct ALTER TABLE COMMENT SQL. @@ -2563,8 +2643,7 @@ def test_build_create_comment_table_exp(self, starrocks_adapter: StarRocksEngine try: # Setup: Create schema and table starrocks_adapter.create_schema(db_name, ignore_if_exists=True) - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE {table_name} ( id INT, name VARCHAR(100) @@ -2572,8 +2651,7 @@ def test_build_create_comment_table_exp(self, starrocks_adapter: StarRocksEngine PRIMARY KEY (id) COMMENT 'initial comment' DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Test: Use _build_create_comment_table_exp to generate SQL table_expr = exp.to_table(table_name) @@ -2584,7 +2662,9 @@ def test_build_create_comment_table_exp(self, starrocks_adapter: StarRocksEngine # Verify: SQL format is correct assert "ALTER TABLE" in comment_sql, f"Invalid SQL format: {comment_sql}" - assert "COMMENT =" in comment_sql, f"Missing COMMENT = in SQL: {comment_sql}" + assert ( + "COMMENT =" in comment_sql + ), f"Missing COMMENT = in SQL: {comment_sql}" assert new_comment in comment_sql, f"Comment not in SQL: {comment_sql}" # Execute the generated SQL @@ -2597,16 +2677,18 @@ def test_build_create_comment_table_exp(self, starrocks_adapter: StarRocksEngine f"WHERE TABLE_SCHEMA = '{db_name}' AND TABLE_NAME = 'test_table'", ) assert result, "Table not found after comment update" - assert result[0] == new_comment, ( - f"Comment not updated. Expected: {new_comment}, Got: {result[0]}" - ) + assert ( + result[0] == new_comment + ), f"Comment not updated. Expected: {new_comment}, Got: {result[0]}" logger.info("✅ _build_create_comment_table_exp generates valid SQL") finally: starrocks_adapter.drop_schema(db_name, ignore_if_not_exists=True) - def test_build_create_comment_column_exp(self, starrocks_adapter: StarRocksEngineAdapter): + def test_build_create_comment_column_exp( + self, starrocks_adapter: StarRocksEngineAdapter + ): """ Test _build_create_comment_column_exp generates correct ALTER TABLE MODIFY COLUMN SQL. @@ -2622,8 +2704,7 @@ def test_build_create_comment_column_exp(self, starrocks_adapter: StarRocksEngin try: # Setup: Create schema and table starrocks_adapter.create_schema(db_name, ignore_if_exists=True) - starrocks_adapter.execute( - f""" + starrocks_adapter.execute(f""" CREATE TABLE {table_name} ( id INT COMMENT 'initial id comment', name VARCHAR(100) COMMENT 'initial name comment', @@ -2631,8 +2712,7 @@ def test_build_create_comment_column_exp(self, starrocks_adapter: StarRocksEngin ) PRIMARY KEY (id) DISTRIBUTED BY HASH(id) BUCKETS 10 - """ - ) + """) # Test: Use _build_create_comment_column_exp to generate SQL table_expr = exp.to_table(table_name) @@ -2646,7 +2726,9 @@ def test_build_create_comment_column_exp(self, starrocks_adapter: StarRocksEngin # Verify: SQL format is correct assert "ALTER TABLE" in comment_sql, f"Invalid SQL format: {comment_sql}" - assert "MODIFY COLUMN" in comment_sql, f"Missing MODIFY COLUMN in SQL: {comment_sql}" + assert ( + "MODIFY COLUMN" in comment_sql + ), f"Missing MODIFY COLUMN in SQL: {comment_sql}" assert "COMMENT" in comment_sql, f"Missing COMMENT in SQL: {comment_sql}" assert new_comment in comment_sql, f"Comment not in SQL: {comment_sql}" @@ -2661,14 +2743,16 @@ def test_build_create_comment_column_exp(self, starrocks_adapter: StarRocksEngin ) assert result is not None, "Column not found after comment update" column_type, column_comment = result - assert column_comment == new_comment, ( - f"Comment not updated. Expected: {new_comment}, Got: {column_comment}" + assert ( + column_comment == new_comment + ), f"Comment not updated. Expected: {new_comment}, Got: {column_comment}" + assert ( + "varchar(100)" in column_type.lower() + ), f"Column type changed unexpectedly: {column_type}" + + logger.info( + "✅ _build_create_comment_column_exp generates valid SQL with correct type" ) - assert "varchar(100)" in column_type.lower(), ( - f"Column type changed unexpectedly: {column_type}" - ) - - logger.info("✅ _build_create_comment_column_exp generates valid SQL with correct type") finally: starrocks_adapter.drop_schema(db_name, ignore_if_not_exists=True) diff --git a/tests/core/engine_adapter/integration/test_integration_trino.py b/tests/core/engine_adapter/integration/test_integration_trino.py index 81313b2a8d..c6e53c6278 100644 --- a/tests/core/engine_adapter/integration/test_integration_trino.py +++ b/tests/core/engine_adapter/integration/test_integration_trino.py @@ -1,22 +1,23 @@ import typing as t +from pathlib import Path + import pytest from pytest import FixtureRequest -from pathlib import Path +from sqlglot import exp, parse_one + from sqlmesh.core.engine_adapter import TrinoEngineAdapter -from tests.core.engine_adapter.integration import TestContext -from sqlglot import parse_one, exp -from tests.core.engine_adapter.integration import ( - TestContext, - generate_pytest_params, - ENGINES_BY_NAME, - IntegrationTestEngine, -) +from tests.core.engine_adapter.integration import (ENGINES_BY_NAME, + IntegrationTestEngine, + TestContext, + generate_pytest_params) @pytest.fixture(params=list(generate_pytest_params(ENGINES_BY_NAME["trino"]))) def ctx( request: FixtureRequest, - create_test_context: t.Callable[[IntegrationTestEngine, str, str], t.Iterable[TestContext]], + create_test_context: t.Callable[ + [IntegrationTestEngine, str, str], t.Iterable[TestContext] + ], ) -> t.Iterable[TestContext]: yield from create_test_context(*request.param) @@ -39,8 +40,7 @@ def test_macros_in_physical_properties( schema = ctx.schema() with open(models_dir / "test_model.sql", "w") as f: - f.write( - """ + f.write(""" MODEL ( name SCHEMA.test, kind FULL, @@ -51,8 +51,7 @@ def test_macros_in_physical_properties( ); select 1 as col_a, 2 as col_b; - """.replace("SCHEMA", schema) - ) + """.replace("SCHEMA", schema)) context = ctx.create_context(path=tmp_path) assert len(context.models) == 1 @@ -65,7 +64,9 @@ def test_macros_in_physical_properties( physical_table_str = snapshot.table_name() physical_table = exp.to_table(physical_table_str) - create_sql = list(engine_adapter.fetchone(f"show create table {physical_table}") or [])[0] + create_sql = list( + engine_adapter.fetchone(f"show create table {physical_table}") or [] + )[0] parsed_create_sql = parse_one(create_sql, dialect="trino") @@ -79,6 +80,11 @@ def test_macros_in_physical_properties( ) sorted_by_property = next( - p for p in parsed_create_sql.find_all(exp.Property) if "sorted_by" in p.sql(dialect="trino") + p + for p in parsed_create_sql.find_all(exp.Property) + if "sorted_by" in p.sql(dialect="trino") + ) + assert ( + sorted_by_property.sql(dialect="trino") + == "sorted_by=ARRAY['col_a ASC NULLS FIRST']" ) - assert sorted_by_property.sql(dialect="trino") == "sorted_by=ARRAY['col_a ASC NULLS FIRST']" diff --git a/tests/core/engine_adapter/test_athena.py b/tests/core/engine_adapter/test_athena.py index 19c92f66ac..697eee9d8e 100644 --- a/tests/core/engine_adapter/test_athena.py +++ b/tests/core/engine_adapter/test_athena.py @@ -1,18 +1,18 @@ import typing as t -import pytest from unittest.mock import Mock -from pytest_mock import MockerFixture -import pandas as pd # noqa: TID253 +import pandas as pd # noqa: TID253 +import pytest +from pytest_mock import MockerFixture from sqlglot import exp, parse_one + import sqlmesh.core.dialect as d from sqlmesh.core.engine_adapter import AthenaEngineAdapter from sqlmesh.core.engine_adapter.shared import DataObject from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.model.definition import SqlModel -from sqlmesh.utils.errors import SQLMeshError from sqlmesh.core.table_diff import TableDiff - +from sqlmesh.utils.errors import SQLMeshError from tests.core.engine_adapter import to_sql_calls pytestmark = [pytest.mark.athena, pytest.mark.engine] @@ -46,16 +46,31 @@ def table_diff(adapter: AthenaEngineAdapter) -> TableDiff: "s3://some/location/table/", ), # Location set to bucket - ("s3://bucket", None, exp.to_table("schema.table"), "s3://bucket/schema/table/"), + ( + "s3://bucket", + None, + exp.to_table("schema.table"), + "s3://bucket/schema/table/", + ), ("s3://bucket", {}, exp.to_table("schema.table"), "s3://bucket/schema/table/"), - ("s3://bucket", None, exp.to_table("schema.table"), "s3://bucket/schema/table/"), + ( + "s3://bucket", + None, + exp.to_table("schema.table"), + "s3://bucket/schema/table/", + ), ( "s3://bucket", {"s3_base_location": exp.Literal.string("s3://some/location/")}, exp.to_table("schema.table"), "s3://some/location/table/", ), - ("s3://bucket", {}, exp.Table(db=exp.Identifier(this="test")), "s3://bucket/test/"), + ( + "s3://bucket", + {}, + exp.Table(db=exp.Identifier(this="test")), + "s3://bucket/test/", + ), # Location set to bucket with prefix ( "s3://bucket/subpath/", @@ -63,7 +78,12 @@ def table_diff(adapter: AthenaEngineAdapter) -> TableDiff: exp.to_table("schema.table"), "s3://bucket/subpath/schema/table/", ), - ("s3://bucket/subpath/", None, exp.to_table("table"), "s3://bucket/subpath/table/"), + ( + "s3://bucket/subpath/", + None, + exp.to_table("table"), + "s3://bucket/subpath/table/", + ), ( "s3://bucket/subpath/", None, @@ -87,7 +107,9 @@ def test_table_location( ) -> None: adapter.s3_warehouse_location = config_s3_warehouse_location if expected_location is None: - with pytest.raises(SQLMeshError, match=r"Cannot figure out location for table.*"): + with pytest.raises( + SQLMeshError, match=r"Cannot figure out location for table.*" + ): adapter._table_location_or_raise(table_properties, table) else: location = adapter._table_location_or_raise( @@ -113,8 +135,7 @@ def test_create_schema(adapter: AthenaEngineAdapter) -> None: def test_create_table_hive(adapter: AthenaEngineAdapter) -> None: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_table, kind FULL, @@ -127,8 +148,7 @@ def test_create_table_hive(adapter: AthenaEngineAdapter) -> None: ); SELECT 1::timestamp AS cola, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) adapter.create_table( @@ -145,8 +165,7 @@ def test_create_table_hive(adapter: AthenaEngineAdapter) -> None: def test_create_table_iceberg(adapter: AthenaEngineAdapter) -> None: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_table, kind FULL, @@ -159,8 +178,7 @@ def test_create_table_iceberg(adapter: AthenaEngineAdapter) -> None: ); SELECT 1::timestamp AS cola, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) adapter.create_table( @@ -178,16 +196,14 @@ def test_create_table_iceberg(adapter: AthenaEngineAdapter) -> None: def test_create_table_no_location(adapter: AthenaEngineAdapter) -> None: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_table, kind FULL ); SELECT a::int FROM foo; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) with pytest.raises(SQLMeshError, match=r"Cannot figure out location.*"): @@ -251,8 +267,7 @@ def test_ctas_iceberg_no_specific_location(adapter: AthenaEngineAdapter): def test_ctas_iceberg_partitioned(adapter: AthenaEngineAdapter): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -263,8 +278,7 @@ def test_ctas_iceberg_partitioned(adapter: AthenaEngineAdapter): ); SELECT 1::timestamp AS business_date, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) adapter.s3_warehouse_location = "s3://bucket/prefix/" @@ -283,7 +297,8 @@ def test_ctas_iceberg_partitioned(adapter: AthenaEngineAdapter): def test_replace_query(adapter: AthenaEngineAdapter, mocker: MockerFixture): mocker.patch( - "sqlmesh.core.engine_adapter.athena.AthenaEngineAdapter.table_exists", return_value=True + "sqlmesh.core.engine_adapter.athena.AthenaEngineAdapter.table_exists", + return_value=True, ) mocker.patch( "sqlmesh.core.engine_adapter.athena.AthenaEngineAdapter._query_table_type", @@ -308,7 +323,8 @@ def test_replace_query(adapter: AthenaEngineAdapter, mocker: MockerFixture): ] mocker.patch( - "sqlmesh.core.engine_adapter.athena.AthenaEngineAdapter.table_exists", return_value=False + "sqlmesh.core.engine_adapter.athena.AthenaEngineAdapter.table_exists", + return_value=False, ) mocker.patch.object(adapter, "_get_data_objects", return_value=[]) adapter.cursor.execute.reset_mock() @@ -332,7 +348,8 @@ def test_columns(adapter: AthenaEngineAdapter, mocker: MockerFixture): mock = mocker.patch( "pandas.io.sql.read_sql_query", return_value=pd.DataFrame( - data=[["col1", "int"], ["col2", "varchar"]], columns=["column_name", "data_type"] + data=[["col1", "int"], ["col2", "varchar"]], + columns=["column_name", "data_type"], ), ) @@ -374,7 +391,9 @@ def test_truncate_table_hive(adapter: AthenaEngineAdapter, mocker: MockerFixture "_is_hive_partitioned_table", return_value=False, ) - mocker.patch.object(adapter, "_query_table_s3_location", return_value="s3://foo/bar") + mocker.patch.object( + adapter, "_query_table_s3_location", return_value="s3://foo/bar" + ) mocker.patch.multiple( adapter, _clear_partition_data=mocker.DEFAULT, _clear_s3_location=mocker.DEFAULT ) @@ -386,7 +405,9 @@ def test_truncate_table_hive(adapter: AthenaEngineAdapter, mocker: MockerFixture t.cast(Mock, adapter._clear_s3_location).assert_called_with("s3://foo/bar") -def test_truncate_table_hive_partitioned(adapter: AthenaEngineAdapter, mocker: MockerFixture): +def test_truncate_table_hive_partitioned( + adapter: AthenaEngineAdapter, mocker: MockerFixture +): mocker.patch.object( adapter, "_query_table_type", @@ -420,7 +441,9 @@ def test_create_state_table(adapter: AthenaEngineAdapter): def test_drop_partitions_from_metastore_uses_batches( adapter: AthenaEngineAdapter, mocker: MockerFixture ): - glue_client_mock = mocker.patch.object(AthenaEngineAdapter, "_glue_client", autospec=True) + glue_client_mock = mocker.patch.object( + AthenaEngineAdapter, "_glue_client", autospec=True + ) glue_client_mock.batch_delete_partition.assert_not_called() @@ -457,8 +480,7 @@ def test_drop_partitions_from_metastore_uses_batches( def test_iceberg_partition_transforms(adapter: AthenaEngineAdapter): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_table, kind FULL, @@ -467,8 +489,7 @@ def test_iceberg_partition_transforms(adapter: AthenaEngineAdapter): ); SELECT 1::timestamp AS business_date, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) assert model.partitioned_by == [ @@ -511,15 +532,30 @@ def test_iceberg_partition_transforms(adapter: AthenaEngineAdapter): ("iceberg", "hive", None, True), # Expect error for mismatched formats ("hive", "iceberg", None, True), # Expect error for mismatched formats ("iceberg", "iceberg", "iceberg", False), - (None, "iceberg", None, True), # Source doesn't exist or type unknown, target is iceberg + ( + None, + "iceberg", + None, + True, + ), # Source doesn't exist or type unknown, target is iceberg ( "iceberg", None, "iceberg", True, ), # Target doesn't exist or type unknown, source is iceberg - (None, "hive", None, False), # Source doesn't exist or type unknown, target is hive - ("hive", None, None, False), # Target doesn't exist or type unknown, source is hive + ( + None, + "hive", + None, + False, + ), # Source doesn't exist or type unknown, target is hive + ( + "hive", + None, + None, + False, + ), # Target doesn't exist or type unknown, source is hive (None, None, None, False), # Both don't exist or types unknown ], ) @@ -550,7 +586,9 @@ def mock_query_table_type(table_name: exp.Table) -> t.Optional[str]: # Mock fetchdf and other calls made within row_diff to avoid actual DB interaction mocker.patch.object(adapter, "fetchdf", return_value=pd.DataFrame()) mocker.patch.object(adapter, "get_data_objects", return_value=[]) - mocker.patch.object(adapter, "columns", return_value={"id": exp.DataType.build("int")}) + mocker.patch.object( + adapter, "columns", return_value={"id": exp.DataType.build("int")} + ) if expect_error: with pytest.raises( diff --git a/tests/core/engine_adapter/test_base.py b/tests/core/engine_adapter/test_base.py index 1971ba3bbc..181da025bf 100644 --- a/tests/core/engine_adapter/test_base.py +++ b/tests/core/engine_adapter/test_base.py @@ -12,15 +12,17 @@ from sqlmesh.core import dialect as d from sqlmesh.core.dialect import normalize_model_name -from sqlmesh.core.engine_adapter import EngineAdapter, EngineAdapterWithIndexSupport -from sqlmesh.core.engine_adapter.shared import InsertOverwriteStrategy, DataObject -from sqlmesh.core.schema_diff import SchemaDiffer, TableAlterOperation, NestedSupport +from sqlmesh.core.engine_adapter import (EngineAdapter, + EngineAdapterWithIndexSupport) +from sqlmesh.core.engine_adapter.shared import (DataObject, + InsertOverwriteStrategy) +from sqlmesh.core.schema_diff import (NestedSupport, SchemaDiffer, + TableAlterOperation) from sqlmesh.utils import columns_to_types_to_struct from sqlmesh.utils.date import to_ds from sqlmesh.utils.errors import SQLMeshError, UnsupportedCatalogOperationError from tests.core.engine_adapter import to_sql_calls - pytestmark = pytest.mark.engine @@ -128,7 +130,9 @@ def test_create_materialized_view(make_mocked_engine_adapter: t.Callable): adapter.cursor.execute.assert_has_calls( [ - call('CREATE OR REPLACE MATERIALIZED VIEW "test_view" AS SELECT "a" FROM "tbl"'), + call( + 'CREATE OR REPLACE MATERIALIZED VIEW "test_view" AS SELECT "a" FROM "tbl"' + ), call('CREATE MATERIALIZED VIEW "test_view" AS SELECT "a" FROM "tbl"'), ] ) @@ -139,7 +143,10 @@ def test_create_materialized_view(make_mocked_engine_adapter: t.Callable): parse_one("SELECT a, b FROM tbl"), replace=False, materialized=True, - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("INT"), + }, ) adapter.create_view( "test_view", parse_one("SELECT a, b FROM tbl"), replace=False, materialized=True @@ -147,7 +154,9 @@ def test_create_materialized_view(make_mocked_engine_adapter: t.Callable): adapter.cursor.execute.assert_has_calls( [ - call('CREATE MATERIALIZED VIEW "test_view" ("a", "b") AS SELECT "a", "b" FROM "tbl"'), + call( + 'CREATE MATERIALIZED VIEW "test_view" ("a", "b") AS SELECT "a", "b" FROM "tbl"' + ), call('CREATE MATERIALIZED VIEW "test_view" AS SELECT "a", "b" FROM "tbl"'), ] ) @@ -223,7 +232,10 @@ def test_insert_overwrite_by_time_partition(make_mocked_engine_adapter: t.Callab end="2022-01-02", time_column="b", time_formatter=lambda x, _: exp.Literal.string(to_ds(x)), - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("STRING"), + }, ) adapter.cursor.begin.assert_called_once() @@ -241,7 +253,10 @@ def test_insert_overwrite_by_time_partition_missing_time_column_type( adapter = make_mocked_engine_adapter(EngineAdapter) columns_mock = mocker.patch.object(adapter, "columns") - columns_mock.return_value = {"a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING")} + columns_mock.return_value = { + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("STRING"), + } adapter.insert_overwrite_by_time_partition( "test_table", @@ -280,7 +295,10 @@ def test_insert_overwrite_by_time_partition_supports_insert_overwrite( end="2022-01-02", time_column="b", time_formatter=lambda x, _: exp.Literal.string(to_ds(x)), - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("STRING"), + }, ) adapter.cursor.execute.assert_called_once_with( @@ -360,7 +378,9 @@ def test_insert_overwrite_by_time_partition_supports_insert_overwrite_query_sour ] -def test_insert_overwrite_by_time_partition_replace_where(make_mocked_engine_adapter: t.Callable): +def test_insert_overwrite_by_time_partition_replace_where( + make_mocked_engine_adapter: t.Callable, +): adapter = make_mocked_engine_adapter(EngineAdapter) adapter.INSERT_OVERWRITE_STRATEGY = InsertOverwriteStrategy.REPLACE_WHERE @@ -371,7 +391,10 @@ def test_insert_overwrite_by_time_partition_replace_where(make_mocked_engine_ada end="2022-01-02", time_column="b", time_formatter=lambda x, _: exp.Literal.string(to_ds(x)), - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("STRING"), + }, ) assert to_sql_calls(adapter) == [ @@ -521,7 +544,10 @@ def test_insert_append_query_select_star(make_mocked_engine_adapter: t.Callable) adapter.insert_append( "test_table", parse_one("SELECT 1 AS a, * FROM tbl"), - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("INT"), + }, ) assert to_sql_calls(adapter) == [ @@ -692,7 +718,9 @@ def test_comments(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture) ] # verify comments aren't registered if the config flag is False - adapter_no_comments = make_mocked_engine_adapter(EngineAdapter, register_comments=False) + adapter_no_comments = make_mocked_engine_adapter( + EngineAdapter, register_comments=False + ) adapter_no_comments.create_table( "test_table", @@ -1085,7 +1113,9 @@ def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: adapter.columns = table_columns - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) adapter.cursor.begin.assert_called_once() adapter.cursor.commit.assert_called_once() @@ -1275,7 +1305,9 @@ def test_merge_when_matched(make_mocked_engine_adapter: t.Callable, assert_exp_e ) -def test_merge_when_matched_multiple(make_mocked_engine_adapter: t.Callable, assert_exp_eq): +def test_merge_when_matched_multiple( + make_mocked_engine_adapter: t.Callable, assert_exp_eq +): adapter = make_mocked_engine_adapter(EngineAdapter) adapter.merge( @@ -1291,7 +1323,9 @@ def test_merge_when_matched_multiple(make_mocked_engine_adapter: t.Callable, ass expressions=[ exp.When( matched=True, - condition=exp.column("ID", "__MERGE_SOURCE__").eq(exp.Literal.number(1)), + condition=exp.column("ID", "__MERGE_SOURCE__").eq( + exp.Literal.number(1) + ), then=exp.Update( expressions=[ exp.column("val", "__MERGE_TARGET__").eq( @@ -1445,10 +1479,7 @@ def test_scd_type_2_by_time(make_mocked_engine_adapter: t.Callable): execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON (COALESCE("id", '') || '|' || COALESCE("name", ''), COALESCE("name", '')) @@ -1615,9 +1646,7 @@ def test_scd_type_2_by_time(make_mocked_engine_adapter: t.Callable): UNION ALL SELECT "id", "name", "price", "test_UPDATED_at", "test_valid_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id", "name", "price", "test_UPDATED_at", "test_valid_from", "test_valid_to" FROM "inserted_rows" ) AS "_subquery" - """ - ).sql() - ) + """).sql() def test_scd_type_2_by_time_source_columns(make_mocked_engine_adapter: t.Callable): @@ -1655,9 +1684,7 @@ def test_scd_type_2_by_time_source_columns(make_mocked_engine_adapter: t.Callabl is_restatement=True, ) sql_calls = to_sql_calls(adapter) - assert ( - parse_one(sql_calls[1]).sql() - == parse_one(""" + assert parse_one(sql_calls[1]).sql() == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON ("id") @@ -1831,10 +1858,11 @@ def test_scd_type_2_by_time_source_columns(make_mocked_engine_adapter: t.Callabl FROM "inserted_rows" ) AS "_subquery" """).sql() - ) -def test_scd_type_2_by_time_no_invalidate_hard_deletes(make_mocked_engine_adapter: t.Callable): +def test_scd_type_2_by_time_no_invalidate_hard_deletes( + make_mocked_engine_adapter: t.Callable, +): adapter = make_mocked_engine_adapter(EngineAdapter) adapter.scd_type_2_by_time( @@ -1858,10 +1886,7 @@ def test_scd_type_2_by_time_no_invalidate_hard_deletes(make_mocked_engine_adapte execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON (COALESCE("id", '')) @@ -2006,9 +2031,7 @@ def test_scd_type_2_by_time_no_invalidate_hard_deletes(make_mocked_engine_adapte UNION ALL SELECT "id", "name", "price", "test_updated_at", "test_valid_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id", "name", "price", "test_updated_at", "test_valid_from", "test_valid_to" FROM "inserted_rows" ) AS "_subquery" - """ - ).sql() - ) + """).sql() def test_merge_scd_type_2_pandas(make_mocked_engine_adapter: t.Callable): @@ -2046,10 +2069,7 @@ def test_merge_scd_type_2_pandas(make_mocked_engine_adapter: t.Callable): execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON ("id1", "id2") @@ -2202,9 +2222,7 @@ def test_merge_scd_type_2_pandas(make_mocked_engine_adapter: t.Callable): "joined"."test_updated_at" > "joined"."t_test_updated_at" ) SELECT CAST("id1" AS INT) AS "id1", CAST("id2" AS INT) AS "id2", CAST("name" AS VARCHAR) AS "name", CAST("price" AS DOUBLE) AS "price", CAST("test_updated_at" AS TIMESTAMPTZ) AS "test_updated_at", CAST("test_valid_from" AS TIMESTAMPTZ) AS "test_valid_from", CAST("test_valid_to" AS TIMESTAMPTZ) AS "test_valid_to" FROM (SELECT "id1", "id2", "name", "price", "test_updated_at", "test_valid_from", "test_valid_to" FROM "static" UNION ALL SELECT "id1", "id2", "name", "price", "test_updated_at", "test_valid_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id1", "id2", "name", "price", "test_updated_at", "test_valid_from", "test_valid_to" FROM "inserted_rows") AS "_subquery" -""" - ).sql() - ) +""").sql() def test_scd_type_2_by_column(make_mocked_engine_adapter: t.Callable): @@ -2212,7 +2230,9 @@ def test_scd_type_2_by_column(make_mocked_engine_adapter: t.Callable): adapter.scd_type_2_by_column( target_table="target", - source_table=t.cast(exp.Select, parse_one("SELECT id, name, price FROM source")), + source_table=t.cast( + exp.Select, parse_one("SELECT id, name, price FROM source") + ), unique_key=[exp.column("id")], valid_from_col=exp.column("test_VALID_from", quoted=True), valid_to_col=exp.column("test_valid_to", quoted=True), @@ -2228,10 +2248,7 @@ def test_scd_type_2_by_column(make_mocked_engine_adapter: t.Callable): extra_col_ignore="testing", ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON ("id") @@ -2383,9 +2400,7 @@ def test_scd_type_2_by_column(make_mocked_engine_adapter: t.Callable): ) ) SELECT CAST("id" AS INT) AS "id", CAST("name" AS VARCHAR) AS "name", CAST("price" AS DOUBLE) AS "price", CAST("test_VALID_from" AS TIMESTAMP) AS "test_VALID_from", CAST("test_valid_to" AS TIMESTAMP) AS "test_valid_to" FROM (SELECT "id", "name", "price", "test_VALID_from", "test_valid_to" FROM "static" UNION ALL SELECT "id", "name", "price", "test_VALID_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id", "name", "price", "test_VALID_from", "test_valid_to" FROM "inserted_rows") AS "_subquery" - """ - ).sql() - ) + """).sql() def test_scd_type_2_by_column_composite_key(make_mocked_engine_adapter: t.Callable): @@ -2393,7 +2408,9 @@ def test_scd_type_2_by_column_composite_key(make_mocked_engine_adapter: t.Callab adapter.scd_type_2_by_column( target_table="target", - source_table=t.cast(exp.Select, parse_one("SELECT id_a, id_b, name, price FROM source")), + source_table=t.cast( + exp.Select, parse_one("SELECT id_a, id_b, name, price FROM source") + ), unique_key=[exp.func("CONCAT", exp.column("id_a"), exp.column("id_b"))], valid_from_col=exp.column("test_VALID_from", quoted=True), valid_to_col=exp.column("test_valid_to", quoted=True), @@ -2409,10 +2426,7 @@ def test_scd_type_2_by_column_composite_key(make_mocked_engine_adapter: t.Callab execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON (CONCAT("id_a", "id_b")) @@ -2575,9 +2589,7 @@ def test_scd_type_2_by_column_composite_key(make_mocked_engine_adapter: t.Callab ) ) SELECT CAST("id_a" AS VARCHAR) AS "id_a", CAST("id_b" AS VARCHAR) AS "id_b", CAST("name" AS VARCHAR) AS "name", CAST("price" AS DOUBLE) AS "price", CAST("test_VALID_from" AS TIMESTAMP) AS "test_VALID_from", CAST("test_valid_to" AS TIMESTAMP) AS "test_valid_to" FROM (SELECT "id_a", "id_b", "name", "price", "test_VALID_from", "test_valid_to" FROM "static" UNION ALL SELECT "id_a", "id_b", "name", "price", "test_VALID_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id_a", "id_b", "name", "price", "test_VALID_from", "test_valid_to" FROM "inserted_rows") AS "_subquery" - """ - ).sql() - ) + """).sql() def test_scd_type_2_truncate(make_mocked_engine_adapter: t.Callable): @@ -2585,7 +2597,9 @@ def test_scd_type_2_truncate(make_mocked_engine_adapter: t.Callable): adapter.scd_type_2_by_column( target_table="target", - source_table=t.cast(exp.Select, parse_one("SELECT id, name, price FROM source")), + source_table=t.cast( + exp.Select, parse_one("SELECT id, name, price FROM source") + ), unique_key=[exp.column("id")], valid_from_col=exp.column("test_valid_from", quoted=True), valid_to_col=exp.column("test_valid_to", quoted=True), @@ -2601,10 +2615,7 @@ def test_scd_type_2_truncate(make_mocked_engine_adapter: t.Callable): truncate=True, ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON ("id") @@ -2758,9 +2769,7 @@ def test_scd_type_2_truncate(make_mocked_engine_adapter: t.Callable): ) ) SELECT CAST("id" AS INT) AS "id", CAST("name" AS VARCHAR) AS "name", CAST("price" AS DOUBLE) AS "price", CAST("test_valid_from" AS TIMESTAMP) AS "test_valid_from", CAST("test_valid_to" AS TIMESTAMP) AS "test_valid_to" FROM (SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "static" UNION ALL SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "inserted_rows") AS "_subquery" - """ - ).sql() - ) + """).sql() def test_scd_type_2_by_column_star_check(make_mocked_engine_adapter: t.Callable): @@ -2768,7 +2777,9 @@ def test_scd_type_2_by_column_star_check(make_mocked_engine_adapter: t.Callable) adapter.scd_type_2_by_column( target_table="target", - source_table=t.cast(exp.Select, parse_one("SELECT id, name, price FROM source")), + source_table=t.cast( + exp.Select, parse_one("SELECT id, name, price FROM source") + ), unique_key=[exp.column("id")], valid_from_col=exp.column("test_valid_from", quoted=True), valid_to_col=exp.column("test_valid_to", quoted=True), @@ -2783,10 +2794,7 @@ def test_scd_type_2_by_column_star_check(make_mocked_engine_adapter: t.Callable) execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON ("id") @@ -2952,17 +2960,19 @@ def test_scd_type_2_by_column_star_check(make_mocked_engine_adapter: t.Callable) ) ) SELECT CAST("id" AS INT) AS "id", CAST("name" AS VARCHAR) AS "name", CAST("price" AS DOUBLE) AS "price", CAST("test_valid_from" AS TIMESTAMP) AS "test_valid_from", CAST("test_valid_to" AS TIMESTAMP) AS "test_valid_to" FROM (SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "static" UNION ALL SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "inserted_rows") AS "_subquery" - """ - ).sql() - ) + """).sql() -def test_scd_type_2_by_column_no_invalidate_hard_deletes(make_mocked_engine_adapter: t.Callable): +def test_scd_type_2_by_column_no_invalidate_hard_deletes( + make_mocked_engine_adapter: t.Callable, +): adapter = make_mocked_engine_adapter(EngineAdapter) adapter.scd_type_2_by_column( target_table="target", - source_table=t.cast(exp.Select, parse_one("SELECT id, name, price FROM source")), + source_table=t.cast( + exp.Select, parse_one("SELECT id, name, price FROM source") + ), unique_key=[exp.column("id")], valid_from_col=exp.column("test_valid_from", quoted=True), valid_to_col=exp.column("test_valid_to", quoted=True), @@ -2978,10 +2988,7 @@ def test_scd_type_2_by_column_no_invalidate_hard_deletes(make_mocked_engine_adap execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert ( - adapter.cursor.execute.call_args[0][0] - == parse_one( - """ + assert adapter.cursor.execute.call_args[0][0] == parse_one(""" CREATE OR REPLACE TABLE "target" AS WITH "source" AS ( SELECT DISTINCT ON ("id") @@ -3130,9 +3137,7 @@ def test_scd_type_2_by_column_no_invalidate_hard_deletes(make_mocked_engine_adap ) ) SELECT CAST("id" AS INT) AS "id", CAST("name" AS VARCHAR) AS "name", CAST("price" AS DOUBLE) AS "price", CAST("test_valid_from" AS TIMESTAMP) AS "test_valid_from", CAST("test_valid_to" AS TIMESTAMP) AS "test_valid_to" FROM (SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "static" UNION ALL SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id", "name", "price", "test_valid_from", "test_valid_to" FROM "inserted_rows") AS "_subquery" - """ - ).sql() - ) + """).sql() def test_replace_query(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): @@ -3145,7 +3150,9 @@ def test_replace_query(make_mocked_engine_adapter: t.Callable, mocker: MockerFix ) adapter.replace_query("test_table", parse_one("SELECT a FROM tbl")) adapter.replace_query( - "test_table", parse_one("SELECT a FROM tbl"), {"a": exp.DataType.build("UNKNOWN")} + "test_table", + parse_one("SELECT a FROM tbl"), + {"a": exp.DataType.build("UNKNOWN")}, ) # TODO: Shouldn't we enforce that `a` is casted to an int? @@ -3182,7 +3189,9 @@ def test_replace_query_pandas(make_mocked_engine_adapter: t.Callable): df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) adapter.replace_query( - "test_table", df, {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")} + "test_table", + df, + {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")}, ) assert to_sql_calls(adapter) == [ @@ -3289,7 +3298,9 @@ def test_replace_query_self_referencing_not_exists_known( ] -def test_create_table_like(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_create_table_like( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(EngineAdapter) columns_to_types = { @@ -3333,7 +3344,9 @@ def test_rename_table(make_mocked_engine_adapter: t.Callable): adapter = make_mocked_engine_adapter(EngineAdapter) adapter.rename_table("old_table", "new_table") - adapter.cursor.execute.assert_called_once_with('ALTER TABLE "old_table" RENAME TO "new_table"') + adapter.cursor.execute.assert_called_once_with( + 'ALTER TABLE "old_table" RENAME TO "new_table"' + ) def test_clone_table(make_mocked_engine_adapter: t.Callable): @@ -3483,13 +3496,17 @@ def test_get_temp_table(mocker: MockerFixture, make_mocked_engine_adapter: t.Cal mocker.patch("sqlmesh.core.engine_adapter.base.random_id", return_value="abcdefgh") value = adapter._get_temp_table( - normalize_model_name("catalog.db.test_table", default_catalog=None, dialect=None) + normalize_model_name( + "catalog.db.test_table", default_catalog=None, dialect=None + ) ) assert value.sql() == '"catalog"."db"."__temp_test_table_abcdefgh"' -def test_get_data_objects_batching(mocker: MockerFixture, make_mocked_engine_adapter: t.Callable): +def test_get_data_objects_batching( + mocker: MockerFixture, make_mocked_engine_adapter: t.Callable +): adapter = make_mocked_engine_adapter(EngineAdapter) adapter._get_data_objects = mocker.Mock(return_value=[]) @@ -3546,7 +3563,9 @@ def test_insert_overwrite_by_partition_query( ): adapter = make_mocked_engine_adapter(EngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test_schema.test_table" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -3578,12 +3597,16 @@ def test_insert_overwrite_by_partition_query( def test_insert_overwrite_by_partition_query_insert_overwrite_strategy( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(EngineAdapter) adapter.INSERT_OVERWRITE_STRATEGY = InsertOverwriteStrategy.INSERT_OVERWRITE - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test_schema.test_table" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -3623,7 +3646,10 @@ def test_log_sql(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): assert mock_logger.log.call_count == 5 assert mock_logger.log.call_args_list[0][0][2] == "SELECT 1" - assert mock_logger.log.call_args_list[1][0][2] == 'INSERT INTO "test" SELECT * FROM "source"' + assert ( + mock_logger.log.call_args_list[1][0][2] + == 'INSERT INTO "test" SELECT * FROM "source"' + ) assert ( mock_logger.log.call_args_list[2][0][2] == 'INSERT INTO "test" ("id", "value") VALUES ""' @@ -3699,7 +3725,9 @@ def test_select_columns( ], ) def test_casted_columns( - columns_to_types: t.Dict[str, exp.DataType], source_columns: t.List[str], expected: t.List[str] + columns_to_types: t.Dict[str, exp.DataType], + source_columns: t.List[str], + expected: t.List[str], ) -> None: assert [ x.sql() for x in EngineAdapter._casted_columns(columns_to_types, source_columns) @@ -3718,11 +3746,15 @@ def test_data_object_cache_get_data_objects( adapter, "_get_data_objects", return_value=[table1, table2] ) - result1 = adapter.get_data_objects("test_schema", {"table1", "table2"}, safe_to_cache=True) + result1 = adapter.get_data_objects( + "test_schema", {"table1", "table2"}, safe_to_cache=True + ) assert len(result1) == 2 assert mock_get_data_objects.call_count == 1 - result2 = adapter.get_data_objects("test_schema", {"table1", "table2"}, safe_to_cache=True) + result2 = adapter.get_data_objects( + "test_schema", {"table1", "table2"}, safe_to_cache=True + ) assert len(result2) == 2 assert mock_get_data_objects.call_count == 1 # Should not increase @@ -3777,7 +3809,9 @@ def test_data_object_cache_get_data_objects_no_object_names( assert len(result1) == 2 assert mock_get_data_objects.call_count == 1 - result2 = adapter.get_data_objects("test_schema", {"table1", "table2"}, safe_to_cache=True) + result2 = adapter.get_data_objects( + "test_schema", {"table1", "table2"}, safe_to_cache=True + ) assert len(result2) == 2 assert mock_get_data_objects.call_count == 1 # Should not increase @@ -3787,9 +3821,13 @@ def test_data_object_cache_get_data_object( ): adapter = make_mocked_engine_adapter(EngineAdapter, patch_get_data_objects=False) - table = DataObject(catalog=None, schema="test_schema", name="test_table", type="table") + table = DataObject( + catalog=None, schema="test_schema", name="test_table", type="table" + ) - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[table]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[table] + ) result1 = adapter.get_data_object("test_schema.test_table", safe_to_cache=True) assert result1 is not None @@ -3807,9 +3845,13 @@ def test_data_object_cache_cleared_on_drop_table( ): adapter = make_mocked_engine_adapter(EngineAdapter, patch_get_data_objects=False) - table = DataObject(catalog=None, schema="test_schema", name="test_table", type="table") + table = DataObject( + catalog=None, schema="test_schema", name="test_table", type="table" + ) - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[table]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[table] + ) adapter.get_data_object("test_schema.test_table", safe_to_cache=True) assert mock_get_data_objects.call_count == 1 @@ -3829,7 +3871,9 @@ def test_data_object_cache_cleared_on_drop_view( view = DataObject(catalog=None, schema="test_schema", name="test_view", type="view") - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[view]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[view] + ) adapter.get_data_object("test_schema.test_view", safe_to_cache=True) assert mock_get_data_objects.call_count == 1 @@ -3847,9 +3891,13 @@ def test_data_object_cache_cleared_on_drop_data_object( ): adapter = make_mocked_engine_adapter(EngineAdapter, patch_get_data_objects=False) - table = DataObject(catalog=None, schema="test_schema", name="test_table", type="table") + table = DataObject( + catalog=None, schema="test_schema", name="test_table", type="table" + ) - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[table]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[table] + ) adapter.get_data_object("test_schema.test_table", safe_to_cache=True) assert mock_get_data_objects.call_count == 1 @@ -3870,13 +3918,17 @@ def test_data_object_cache_cleared_on_create_table( adapter = make_mocked_engine_adapter(EngineAdapter, patch_get_data_objects=False) # Initially cache that table doesn't exist - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[] + ) result = adapter.get_data_object("test_schema.test_table", safe_to_cache=True) assert result is None assert mock_get_data_objects.call_count == 1 # Create the table - table = DataObject(catalog=None, schema="test_schema", name="test_table", type="table") + table = DataObject( + catalog=None, schema="test_schema", name="test_table", type="table" + ) mock_get_data_objects.return_value = [table] adapter.create_table( "test_schema.test_table", @@ -3897,7 +3949,9 @@ def test_data_object_cache_cleared_on_create_view( adapter = make_mocked_engine_adapter(EngineAdapter, patch_get_data_objects=False) # Initially cache that view doesn't exist - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[] + ) result = adapter.get_data_object("test_schema.test_view", safe_to_cache=True) assert result is None assert mock_get_data_objects.call_count == 1 @@ -3922,11 +3976,15 @@ def test_data_object_cache_cleared_on_clone_table( from sqlmesh.core.engine_adapter.snowflake import SnowflakeEngineAdapter adapter = make_mocked_engine_adapter( - SnowflakeEngineAdapter, patch_get_data_objects=False, default_catalog="test_catalog" + SnowflakeEngineAdapter, + patch_get_data_objects=False, + default_catalog="test_catalog", ) # Initially cache that target table doesn't exist - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[] + ) result = adapter.get_data_object("test_schema.test_target", safe_to_cache=True) assert result is None assert mock_get_data_objects.call_count == 1 @@ -3950,21 +4008,29 @@ def test_data_object_cache_with_catalog( from sqlmesh.core.engine_adapter.snowflake import SnowflakeEngineAdapter adapter = make_mocked_engine_adapter( - SnowflakeEngineAdapter, patch_get_data_objects=False, default_catalog="test_catalog" + SnowflakeEngineAdapter, + patch_get_data_objects=False, + default_catalog="test_catalog", ) table = DataObject( catalog="test_catalog", schema="test_schema", name="test_table", type="table" ) - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[table]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[table] + ) - result1 = adapter.get_data_object("test_catalog.test_schema.test_table", safe_to_cache=True) + result1 = adapter.get_data_object( + "test_catalog.test_schema.test_table", safe_to_cache=True + ) assert result1 is not None assert result1.catalog == "test_catalog" assert mock_get_data_objects.call_count == 1 - result2 = adapter.get_data_object("test_catalog.test_schema.test_table", safe_to_cache=True) + result2 = adapter.get_data_object( + "test_catalog.test_schema.test_table", safe_to_cache=True + ) assert result2 is not None assert result2.catalog == "test_catalog" assert mock_get_data_objects.call_count == 1 # Should not increase @@ -3987,7 +4053,9 @@ def test_data_object_cache_partial_cache_hit( assert mock_get_data_objects.call_count == 1 mock_get_data_objects.return_value = [table3] - result = adapter.get_data_objects("test_schema", {"table1", "table3"}, safe_to_cache=True) + result = adapter.get_data_objects( + "test_schema", {"table1", "table3"}, safe_to_cache=True + ) assert len(result) == 2 assert {obj.name for obj in result} == {"table1", "table3"} @@ -4002,13 +4070,19 @@ def test_data_object_cache_get_data_objects_missing_objects( table1 = DataObject(catalog=None, schema="test_schema", name="table1", type="table") table2 = DataObject(catalog=None, schema="test_schema", name="table2", type="table") - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[] + ) - result1 = adapter.get_data_objects("test_schema", {"table1", "table2"}, safe_to_cache=True) + result1 = adapter.get_data_objects( + "test_schema", {"table1", "table2"}, safe_to_cache=True + ) assert not result1 assert mock_get_data_objects.call_count == 1 - result2 = adapter.get_data_objects("test_schema", {"table1", "table2"}, safe_to_cache=True) + result2 = adapter.get_data_objects( + "test_schema", {"table1", "table2"}, safe_to_cache=True + ) assert not result2 assert mock_get_data_objects.call_count == 1 # Should not increase @@ -4022,7 +4096,9 @@ def test_data_object_cache_cleared_on_rename_table( ): adapter = make_mocked_engine_adapter(EngineAdapter, patch_get_data_objects=False) - old_table = DataObject(catalog=None, schema="test_schema", name="old_table", type="table") + old_table = DataObject( + catalog=None, schema="test_schema", name="old_table", type="table" + ) mock_get_data_objects = mocker.patch.object( adapter, "_get_data_objects", return_value=[old_table] ) @@ -4032,7 +4108,9 @@ def test_data_object_cache_cleared_on_rename_table( assert result.name == "old_table" assert mock_get_data_objects.call_count == 1 - new_table = DataObject(catalog=None, schema="test_schema", name="new_table", type="table") + new_table = DataObject( + catalog=None, schema="test_schema", name="new_table", type="table" + ) mock_get_data_objects.return_value = [new_table] adapter.rename_table("test_schema.old_table", "test_schema.new_table") @@ -4061,12 +4139,16 @@ def test_data_object_cache_cleared_on_create_table_like( } mocker.patch.object(adapter, "columns", return_value=columns_to_types) - mock_get_data_objects = mocker.patch.object(adapter, "_get_data_objects", return_value=[]) + mock_get_data_objects = mocker.patch.object( + adapter, "_get_data_objects", return_value=[] + ) result = adapter.get_data_object("test_schema.target_table", safe_to_cache=True) assert result is None assert mock_get_data_objects.call_count == 1 - target_table = DataObject(catalog=None, schema="test_schema", name="target_table", type="table") + target_table = DataObject( + catalog=None, schema="test_schema", name="target_table", type="table" + ) mock_get_data_objects.return_value = [target_table] adapter.create_table_like("test_schema.target_table", "test_schema.source_table") @@ -4173,7 +4255,9 @@ def test_sync_grants_config_unsupported_engine(make_mocked_engine_adapter: t.Cal adapter.sync_grants_config(relation, grants_config) -def test_get_current_grants_config_not_implemented(make_mocked_engine_adapter: t.Callable): +def test_get_current_grants_config_not_implemented( + make_mocked_engine_adapter: t.Callable, +): adapter = make_mocked_engine_adapter(EngineAdapter) relation = exp.to_table("test_table") diff --git a/tests/core/engine_adapter/test_base_postgres.py b/tests/core/engine_adapter/test_base_postgres.py index f286c47c56..e7721d5ebc 100644 --- a/tests/core/engine_adapter/test_base_postgres.py +++ b/tests/core/engine_adapter/test_base_postgres.py @@ -78,10 +78,14 @@ def test_drop_view(make_mocked_engine_adapter: t.Callable): ) -def test_get_current_schema(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_get_current_schema( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(BasePostgresEngineAdapter) - fetchone_mock = mocker.patch.object(adapter, "fetchone", return_value=("test_schema",)) + fetchone_mock = mocker.patch.object( + adapter, "fetchone", return_value=("test_schema",) + ) result = adapter._get_current_schema() assert result == "test_schema" diff --git a/tests/core/engine_adapter/test_bigquery.py b/tests/core/engine_adapter/test_bigquery.py index 134f144df1..176ac5c3df 100644 --- a/tests/core/engine_adapter/test_bigquery.py +++ b/tests/core/engine_adapter/test_bigquery.py @@ -22,7 +22,9 @@ @pytest.fixture -def adapter(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture) -> BigQueryEngineAdapter: +def adapter( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +) -> BigQueryEngineAdapter: mocked_adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) mocker.patch("sqlmesh.core.engine_adapter.bigquery.BigQueryEngineAdapter.execute") return mocked_adapter @@ -54,14 +56,18 @@ def test_insert_overwrite_by_time_partition_query( def test_insert_overwrite_by_partition_query( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) execute_mock = mocker.patch( "sqlmesh.core.engine_adapter.bigquery.BigQueryEngineAdapter.execute" ) adapter._default_catalog = "test_project" - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test_schema.test_table" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -89,7 +95,9 @@ def test_insert_overwrite_by_partition_query( def test_insert_overwrite_by_partition_query_unknown_column_types( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) execute_mock = mocker.patch( @@ -104,7 +112,9 @@ def test_insert_overwrite_by_partition_query_unknown_column_types( "ds": exp.DataType.build("DATETIME"), } adapter._default_catalog = "test_project" - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test_schema.test_table" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -143,7 +153,10 @@ def test_insert_overwrite_by_time_partition_pandas( def temp_table_exists(table: exp.Table) -> bool: nonlocal temp_table_exists_counter temp_table_exists_counter += 1 - if table.sql() == "project.dataset.temp_table" and temp_table_exists_counter == 1: + if ( + table.sql() == "project.dataset.temp_table" + and temp_table_exists_counter == 1 + ): return False return True @@ -172,7 +185,9 @@ def temp_table_exists(table: exp.Table) -> bool: retry_resp_call.errors = None retry_mock.return_value = retry_resp db_call_mock.return_value = AttributeDict({"errors": None}) - df = pd.DataFrame({"a": [1, 2, 3], "ds": ["2020-01-01", "2020-01-02", "2020-01-03"]}) + df = pd.DataFrame( + {"a": [1, 2, 3], "ds": ["2020-01-01", "2020-01-02", "2020-01-03"]} + ) adapter.insert_overwrite_by_time_partition( "test_table", df, @@ -233,7 +248,9 @@ def test_replace_query(make_mocked_engine_adapter: t.Callable, mocker: MockerFix ] -def test_replace_query_pandas(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_replace_query_pandas( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) get_bq_table_value = AttributeDict( @@ -247,7 +264,10 @@ def test_replace_query_pandas(make_mocked_engine_adapter: t.Callable, mocker: Mo def temp_table_exists(table: exp.Table) -> bool: nonlocal temp_table_exists_counter temp_table_exists_counter += 1 - if table.sql() == "project.dataset.temp_table" and temp_table_exists_counter == 1: + if ( + table.sql() == "project.dataset.temp_table" + and temp_table_exists_counter == 1 + ): return False return True @@ -274,7 +294,9 @@ def temp_table_exists(table: exp.Table) -> bool: df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) adapter.replace_query( - "test_table", df, {"a": exp.DataType.build("int"), "b": exp.DataType.build("int")} + "test_table", + df, + {"a": exp.DataType.build("int"), "b": exp.DataType.build("int")}, ) assert db_call_mock.call_count == 1 @@ -308,7 +330,10 @@ def temp_table_exists(table: exp.Table) -> bool: "partition_by_cols, partition_by_statement", [ ([exp.to_column("ds")], "`ds`"), - ([d.parse_one("DATE_TRUNC(ds, MONTH)", dialect="bigquery")], "DATE_TRUNC(`ds`, MONTH)"), + ( + [d.parse_one("DATE_TRUNC(ds, MONTH)", dialect="bigquery")], + "DATE_TRUNC(`ds`, MONTH)", + ), ], ) def test_create_table_date_partition( @@ -344,14 +369,54 @@ def test_create_table_date_partition( ([exp.to_column("ds")], "date", IntervalUnit.DAY, "`ds`"), ([exp.to_column("ds")], "date", IntervalUnit.MONTH, "DATE_TRUNC(`ds`, MONTH)"), ([exp.to_column("ds")], "date", IntervalUnit.YEAR, "DATE_TRUNC(`ds`, YEAR)"), - ([exp.to_column("ds")], "datetime", IntervalUnit.HOUR, "DATETIME_TRUNC(`ds`, HOUR)"), - ([exp.to_column("ds")], "datetime", IntervalUnit.DAY, "DATETIME_TRUNC(`ds`, DAY)"), - ([exp.to_column("ds")], "datetime", IntervalUnit.MONTH, "DATETIME_TRUNC(`ds`, MONTH)"), - ([exp.to_column("ds")], "datetime", IntervalUnit.YEAR, "DATETIME_TRUNC(`ds`, YEAR)"), - ([exp.to_column("ds")], "timestamp", IntervalUnit.HOUR, "TIMESTAMP_TRUNC(`ds`, HOUR)"), - ([exp.to_column("ds")], "timestamp", IntervalUnit.DAY, "TIMESTAMP_TRUNC(`ds`, DAY)"), - ([exp.to_column("ds")], "timestamp", IntervalUnit.MONTH, "TIMESTAMP_TRUNC(`ds`, MONTH)"), - ([exp.to_column("ds")], "timestamp", IntervalUnit.YEAR, "TIMESTAMP_TRUNC(`ds`, YEAR)"), + ( + [exp.to_column("ds")], + "datetime", + IntervalUnit.HOUR, + "DATETIME_TRUNC(`ds`, HOUR)", + ), + ( + [exp.to_column("ds")], + "datetime", + IntervalUnit.DAY, + "DATETIME_TRUNC(`ds`, DAY)", + ), + ( + [exp.to_column("ds")], + "datetime", + IntervalUnit.MONTH, + "DATETIME_TRUNC(`ds`, MONTH)", + ), + ( + [exp.to_column("ds")], + "datetime", + IntervalUnit.YEAR, + "DATETIME_TRUNC(`ds`, YEAR)", + ), + ( + [exp.to_column("ds")], + "timestamp", + IntervalUnit.HOUR, + "TIMESTAMP_TRUNC(`ds`, HOUR)", + ), + ( + [exp.to_column("ds")], + "timestamp", + IntervalUnit.DAY, + "TIMESTAMP_TRUNC(`ds`, DAY)", + ), + ( + [exp.to_column("ds")], + "timestamp", + IntervalUnit.MONTH, + "TIMESTAMP_TRUNC(`ds`, MONTH)", + ), + ( + [exp.to_column("ds")], + "timestamp", + IntervalUnit.YEAR, + "TIMESTAMP_TRUNC(`ds`, YEAR)", + ), ( [d.parse_one("TIMESTAMP_TRUNC(ds, HOUR)", dialect="bigquery")], "timestamp", @@ -380,7 +445,9 @@ def test_create_table_time_partition( partition_by_statement, mocker: MockerFixture, ): - partition_column_sql_type = exp.DataType.build(partition_column_type, dialect="bigquery") + partition_column_sql_type = exp.DataType.build( + partition_column_type, dialect="bigquery" + ) adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) @@ -458,7 +525,10 @@ def test_merge_pandas(make_mocked_engine_adapter: t.Callable, mocker: MockerFixt def temp_table_exists(table: exp.Table) -> bool: nonlocal temp_table_exists_counter temp_table_exists_counter += 1 - if table.sql() == "project.dataset.temp_table" and temp_table_exists_counter == 1: + if ( + table.sql() == "project.dataset.temp_table" + and temp_table_exists_counter == 1 + ): return False return True @@ -564,10 +634,15 @@ def test_begin_end_session(mocker: MockerFixture): assert not execute_b_call[1]["job_config"].connection_properties # starting a new session with session property query_label and array value - with adapter.session({"query_label": parse_one("[('key1', 'value1'), ('key2', 'value2')]")}): + with adapter.session( + {"query_label": parse_one("[('key1', 'value1'), ('key2', 'value2')]")} + ): adapter.execute("SELECT 4;") begin_new_session_call = connection_mock._client.query.call_args_list[3] - assert begin_new_session_call[0][0] == 'SET @@query_label = "key1:value1,key2:value2";SELECT 1;' + assert ( + begin_new_session_call[0][0] + == 'SET @@query_label = "key1:value1,key2:value2";SELECT 1;' + ) # starting a new session with session property query_label and Paren value with adapter.session({"query_label": parse_one("(('key1', 'value1'))")}): @@ -600,7 +675,9 @@ def _to_sql_calls(execute_mock: t.Any, identify: bool = True) -> t.List[str]: return output -def test_create_table_table_options(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_create_table_table_options( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) execute_mock = mocker.patch( @@ -741,7 +818,9 @@ def test_nested_comments(make_mocked_engine_adapter: t.Callable, mocker: MockerF "repeated_record.nested_repeated_record.struct_field with space.nested_field": "Nested Repeated Record Nested Field", "same_name_": "Level 1", "same_name_.same_name_.same_name_": "Level 3", - "same_name_.same_name_.same_name_.same_name_": "4" * allowed_column_comment_length + "X", + "same_name_.same_name_.same_name_.same_name_": "4" + * allowed_column_comment_length + + "X", } adapter.create_table( @@ -958,7 +1037,9 @@ def test_materialized_view_properties( ] -def test_nested_fields_update(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_nested_fields_update( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) current_schema = [ @@ -977,7 +1058,10 @@ def test_nested_fields_update(make_mocked_engine_adapter: t.Callable, mocker: Mo ), ) ] - new_nested_fields = [("year", "INT64", ["user", "orders"]), ("active", "BOOL", ["user"])] + new_nested_fields = [ + ("year", "INT64", ["user", "orders"]), + ("active", "BOOL", ["user"]), + ] expected = [ bigquery.SchemaField( "user", @@ -1087,7 +1171,9 @@ def test_get_alter_expressions_includes_catalog( assert tables == {"bing"} -def test_job_cancellation_on_keyboard_interrupt_job_still_running(mocker: MockerFixture): +def test_job_cancellation_on_keyboard_interrupt_job_still_running( + mocker: MockerFixture, +): # Create a mock connection connection_mock = mocker.NonCallableMock() cursor_mock = mocker.Mock() @@ -1171,7 +1257,10 @@ def test_drop_cascade(adapter: BigQueryEngineAdapter): # BigQuery doesnt support DROP CASCADE for tables # ref: https://cloud.google.com/bigquery/docs/reference/standard-sql/data-definition-language#drop_table_statement - assert _to_sql_calls(adapter) == ["DROP TABLE IF EXISTS `foo`", "DROP TABLE IF EXISTS `foo`"] + assert _to_sql_calls(adapter) == [ + "DROP TABLE IF EXISTS `foo`", + "DROP TABLE IF EXISTS `foo`", + ] adapter.execute.reset_mock() # type: ignore # But, it does for schemas @@ -1217,11 +1306,16 @@ def test_scd_type_2_by_partitioning(adapter: BigQueryEngineAdapter): assert "PARTITION BY TIMESTAMP_TRUNC(`valid_from`, DAY)" in calls[1] -def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_sync_grants_config( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(BigQueryEngineAdapter) relation = exp.to_table("project.dataset.test_table", dialect="bigquery") new_grants_config = { - "roles/bigquery.dataViewer": ["user:analyst@example.com", "group:data-team@example.com"], + "roles/bigquery.dataViewer": [ + "user:analyst@example.com", + "group:data-team@example.com", + ], "roles/bigquery.dataEditor": ["user:admin@example.com"], } current_grants = [ @@ -1229,7 +1323,9 @@ def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: Mock ("roles/bigquery.admin", "user:old_admin@example.com"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) execute_mock = mocker.patch.object(adapter, "execute") mocker.patch.object(adapter, "get_current_catalog", return_value="project") mocker.patch.object(adapter.client, "location", "us-central1") @@ -1281,7 +1377,10 @@ def test_sync_grants_config_with_overlaps( "user:analyst2@example.com", "user:analyst3@example.com", ], - "roles/bigquery.dataEditor": ["user:analyst2@example.com", "user:editor@example.com"], + "roles/bigquery.dataEditor": [ + "user:analyst2@example.com", + "user:editor@example.com", + ], } current_grants = [ ("roles/bigquery.dataViewer", "user:analyst1@example.com"), # Keep @@ -1290,7 +1389,9 @@ def test_sync_grants_config_with_overlaps( ("roles/bigquery.admin", "user:admin@example.com"), # Remove ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) execute_mock = mocker.patch.object(adapter, "execute") mocker.patch.object(adapter, "get_current_catalog", return_value="project") mocker.patch.object(adapter.client, "location", "us-central1") @@ -1378,5 +1479,7 @@ def test_sync_grants_config_no_schema( "roles/bigquery.dataEditor": ["user:editor@example.com"], } - with pytest.raises(ValueError, match="Table test_table does not have a schema \\(dataset\\)"): + with pytest.raises( + ValueError, match="Table test_table does not have a schema \\(dataset\\)" + ): adapter.sync_grants_config(relation, new_grants_config) diff --git a/tests/core/engine_adapter/test_clickhouse.py b/tests/core/engine_adapter/test_clickhouse.py index a3dfe0fdda..8ae0552390 100644 --- a/tests/core/engine_adapter/test_clickhouse.py +++ b/tests/core/engine_adapter/test_clickhouse.py @@ -1,16 +1,18 @@ +import typing as t +from datetime import datetime + import pytest +from pytest_mock.plugin import MockerFixture +from sqlglot import exp, parse_one +from sqlglot.optimizer.qualify_columns import quote_identifiers + +from sqlmesh.core import dialect as d +from sqlmesh.core.dialect import parse from sqlmesh.core.engine_adapter import ClickhouseEngineAdapter +from sqlmesh.core.engine_adapter.shared import DataObject, EngineRunMode from sqlmesh.core.model.definition import load_sql_based_model from sqlmesh.core.model.kind import ModelKindName -from sqlmesh.core.engine_adapter.shared import EngineRunMode, DataObject from tests.core.engine_adapter import to_sql_calls -from sqlmesh.core.dialect import parse -from sqlglot import exp, parse_one -import typing as t -from datetime import datetime -from pytest_mock.plugin import MockerFixture -from sqlmesh.core import dialect as d -from sqlglot.optimizer.qualify_columns import quote_identifiers pytestmark = [pytest.mark.clickhouse, pytest.mark.engine] @@ -82,7 +84,9 @@ def test_create_table(adapter: ClickhouseEngineAdapter, mocker): ) # ON CLUSTER not added because engine_run_mode.is_cluster=False - adapter.create_table("foo", {"a": exp.DataType.build("Int8", dialect=adapter.dialect)}) + adapter.create_table( + "foo", {"a": exp.DataType.build("Int8", dialect=adapter.dialect)} + ) # adapter.create_table_like("target", "source") mocker.patch.object( @@ -90,7 +94,9 @@ def test_create_table(adapter: ClickhouseEngineAdapter, mocker): "engine_run_mode", new_callable=mocker.PropertyMock(return_value=EngineRunMode.CLUSTER), ) - adapter.create_table("foo", {"a": exp.DataType.build("Int8", dialect=adapter.dialect)}) + adapter.create_table( + "foo", {"a": exp.DataType.build("Int8", dialect=adapter.dialect)} + ) # adapter.create_table_like("target", "source") assert to_sql_calls(adapter) == [ @@ -164,14 +170,20 @@ def test_alter_table( def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: if table_name == current_table_name: return { - k: exp.DataType.build(v, dialect=adapter.dialect) for k, v in current_table.items() + k: exp.DataType.build(v, dialect=adapter.dialect) + for k, v in current_table.items() } - return {k: exp.DataType.build(v, dialect=adapter.dialect) for k, v in target_table.items()} + return { + k: exp.DataType.build(v, dialect=adapter.dialect) + for k, v in target_table.items() + } adapter.columns = table_columns # type: ignore # ON CLUSTER not added because engine_run_mode.is_cluster=False - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) mocker.patch.object( ClickhouseEngineAdapter, @@ -184,7 +196,9 @@ def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: new_callable=mocker.PropertyMock(return_value=EngineRunMode.CLUSTER), ) - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) assert to_sql_calls(adapter) == [ 'ALTER TABLE "test_table" DROP COLUMN "c"', @@ -234,7 +248,8 @@ def test_nullable_datatypes_in_model_columns(adapter: ClickhouseEngineAdapter): ) rendered_columns_to_types = { - k: v.sql(dialect="clickhouse") for k, v in model.columns_to_types_or_raise.items() + k: v.sql(dialect="clickhouse") + for k, v in model.columns_to_types_or_raise.items() } assert rendered_columns_to_types["id"] == "Int64" @@ -244,7 +259,9 @@ def test_nullable_datatypes_in_model_columns(adapter: ClickhouseEngineAdapter): def test_model_properties(adapter: ClickhouseEngineAdapter): - def build_properties_sql(storage_format="", order_by="", primary_key="", properties=""): + def build_properties_sql( + storage_format="", order_by="", primary_key="", properties="" + ): model = load_sql_based_model( parse( f""" @@ -268,7 +285,8 @@ def build_properties_sql(storage_format="", order_by="", primary_key="", propert ) return adapter._build_table_properties_exp( - storage_format=model.storage_format, table_properties=model.physical_properties + storage_format=model.storage_format, + table_properties=model.physical_properties, ).sql("clickhouse") # no order by or primary key because table engine is not part of "MergeTree" engine family @@ -296,12 +314,16 @@ def build_properties_sql(storage_format="", order_by="", primary_key="", propert ) assert ( - build_properties_sql(order_by='ORDER_BY = "a",', primary_key='PRIMARY_KEY = "a",') + build_properties_sql( + order_by='ORDER_BY = "a",', primary_key='PRIMARY_KEY = "a",' + ) == 'ENGINE=MergeTree ORDER BY ("a") PRIMARY KEY ("a")' ) assert ( - build_properties_sql(order_by="ORDER_BY = (a),", primary_key="PRIMARY_KEY = (a)") + build_properties_sql( + order_by="ORDER_BY = (a),", primary_key="PRIMARY_KEY = (a)" + ) == "ENGINE=MergeTree ORDER BY (a) PRIMARY KEY (a)" ) @@ -315,7 +337,9 @@ def build_properties_sql(storage_format="", order_by="", primary_key="", propert ) assert ( - build_properties_sql(order_by="ORDER_BY = (a, b),", primary_key="PRIMARY_KEY = (a)") + build_properties_sql( + order_by="ORDER_BY = (a, b),", primary_key="PRIMARY_KEY = (a)" + ) == "ENGINE=MergeTree ORDER BY (a, b) PRIMARY KEY (a)" ) @@ -324,14 +348,20 @@ def build_properties_sql(storage_format="", order_by="", primary_key="", propert == "ENGINE=MergeTree ORDER BY () PRIMARY KEY (a)" ) - assert build_properties_sql(order_by="ORDER_BY = a + 1,") == "ENGINE=MergeTree ORDER BY (a + 1)" + assert ( + build_properties_sql(order_by="ORDER_BY = a + 1,") + == "ENGINE=MergeTree ORDER BY (a + 1)" + ) assert ( - build_properties_sql(order_by="ORDER_BY = (a + 1),") == "ENGINE=MergeTree ORDER BY (a + 1)" + build_properties_sql(order_by="ORDER_BY = (a + 1),") + == "ENGINE=MergeTree ORDER BY (a + 1)" ) assert ( - build_properties_sql(order_by="ORDER_BY = (a, b + 1),", primary_key="PRIMARY_KEY = (a, b)") + build_properties_sql( + order_by="ORDER_BY = (a, b + 1),", primary_key="PRIMARY_KEY = (a, b)" + ) == "ENGINE=MergeTree ORDER BY (a, b + 1) PRIMARY KEY (a, b)" ) @@ -574,7 +604,9 @@ def test_create_table_properties(make_mocked_engine_adapter: t.Callable, mocker) ] -def test_nulls_after_join(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_nulls_after_join( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) query = exp.select("col1").from_("table") @@ -597,7 +629,8 @@ def test_nulls_after_join(make_mocked_engine_adapter: t.Callable, mocker: Mocker ) assert ( - adapter.use_server_nulls_for_unmatched_after_join(query_with_setting) == query_with_setting + adapter.use_server_nulls_for_unmatched_after_join(query_with_setting) + == query_with_setting ) # Server default of 0 != method default of 1, so we inject 0 @@ -613,11 +646,15 @@ def test_nulls_after_join(make_mocked_engine_adapter: t.Callable, mocker: Mocker def test_scd_type_2_by_time( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.side_effect = [ make_temp_table_name(table_name, "efgh"), @@ -630,7 +667,9 @@ def test_scd_type_2_by_time( return_value=[DataObject(schema="", name=table_name, type="table")], ) - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = None # The SCD query we build must specify the setting join_use_nulls = 1. We need to ensure that our @@ -639,7 +678,9 @@ def test_scd_type_2_by_time( # This test's user query does not contain a setting "join_use_nulls", so we determine whether or not # to inject it based on the current server value. The mocked server value is 1, so we should not # inject. - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = "1" adapter.scd_type_2_by_time( @@ -665,8 +706,10 @@ def test_scd_type_2_by_time( execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert to_sql_calls(adapter)[3] == parse_one( - """ + assert ( + to_sql_calls(adapter)[3] + == parse_one( + """ INSERT INTO "__temp_target_abcd" ("id", "name", "price", "test_UPDATED_at", "test_valid_from", "test_valid_to") WITH "source" AS ( SELECT DISTINCT ON (COALESCE("id", '') || '|' || COALESCE("name", ''), COALESCE("name", '')) @@ -826,16 +869,21 @@ def test_scd_type_2_by_time( UNION ALL SELECT "id", "name", "price", "test_UPDATED_at", "test_valid_from", "test_valid_to" FROM "inserted_rows" SETTINGS join_use_nulls = 1 """, - dialect=adapter.dialect, - ).sql(adapter.dialect) + dialect=adapter.dialect, + ).sql(adapter.dialect) + ) def test_scd_type_2_by_column( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.side_effect = [ make_temp_table_name(table_name, "efgh"), @@ -848,7 +896,9 @@ def test_scd_type_2_by_column( return_value=[DataObject(schema="", name=table_name, type="table")], ) - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = None # The SCD query we build must specify the setting join_use_nulls = 1. We need to ensure that our @@ -856,12 +906,16 @@ def test_scd_type_2_by_column( # # This test's user query does not contain a setting "join_use_nulls", so we determine whether or not # to inject it based on the current server value. The mocked server value is 0, so we should inject. - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = "0" adapter.scd_type_2_by_column( target_table="target", - source_table=t.cast(exp.Select, parse_one("SELECT id, name, price FROM source")), + source_table=t.cast( + exp.Select, parse_one("SELECT id, name, price FROM source") + ), unique_key=[exp.column("id")], valid_from_col=exp.column("test_VALID_from", quoted=True), valid_to_col=exp.column("test_valid_to", quoted=True), @@ -876,8 +930,10 @@ def test_scd_type_2_by_column( execution_time=datetime(2020, 1, 1, 0, 0, 0), ) - assert to_sql_calls(adapter)[3] == parse_one( - """ + assert ( + to_sql_calls(adapter)[3] + == parse_one( + """ INSERT INTO "__temp_target_abcd" ("id", "name", "price", "test_VALID_from", "test_valid_to") WITH "source" AS ( SELECT DISTINCT ON ("id") @@ -1031,20 +1087,27 @@ def test_scd_type_2_by_column( ) SELECT "id", "name", "price", "test_VALID_from", "test_valid_to" FROM "static" UNION ALL SELECT "id", "name", "price", "test_VALID_from", "test_valid_to" FROM "updated_rows" UNION ALL SELECT "id", "name", "price", "test_VALID_from", "test_valid_to" FROM "inserted_rows" SETTINGS join_use_nulls = 1 """, - dialect=adapter.dialect, - ).sql(adapter.dialect) + dialect=adapter.dialect, + ).sql(adapter.dialect) + ) def test_insert_overwrite_by_condition_replace_partitioned( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.return_value = make_temp_table_name(table_name, "abcd") - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = "dateTrunc('WEEK', ds)" insert_table_name = make_temp_table_name("new_records", "abcd") @@ -1074,15 +1137,21 @@ def test_insert_overwrite_by_condition_replace_partitioned( def test_insert_overwrite_by_condition_replace( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.return_value = make_temp_table_name(table_name, "abcd") - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = None insert_table_name = make_temp_table_name("new_records", "abcd") @@ -1112,18 +1181,26 @@ def test_insert_overwrite_by_condition_replace( def test_insert_overwrite_by_condition_where_partitioned( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.return_value = make_temp_table_name(table_name, "abcd") - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = "dateTrunc('WEEK', ds)" - fetchall_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchall") + fetchall_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchall" + ) fetchall_mock.side_effect = [ [("1",), ("2",), ("3",), ("4",)], ["1", "2", "4"], @@ -1163,15 +1240,21 @@ def test_insert_overwrite_by_condition_where_partitioned( def test_insert_overwrite_by_condition_by_key( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.return_value = make_temp_table_name(table_name, "abcd") - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = None insert_table_name = make_temp_table_name("new_records", "abcd") @@ -1218,18 +1301,26 @@ def test_insert_overwrite_by_condition_by_key( def test_insert_overwrite_by_condition_by_key_partitioned( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.return_value = make_temp_table_name(table_name, "abcd") - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.side_effect = ["dateTrunc('WEEK', ds)", "dateTrunc('WEEK', ds)"] - fetchall_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchall") + fetchall_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchall" + ) fetchall_mock.side_effect = [ [("1",), ("2",), ("3",), ("4",)], ["1", "2", "4"], @@ -1283,18 +1374,26 @@ def test_insert_overwrite_by_condition_by_key_partitioned( def test_insert_overwrite_by_condition_inc_by_partition( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_mock.return_value = make_temp_table_name(table_name, "abcd") - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchone" + ) fetchone_mock.return_value = "dateTrunc('WEEK', ds)" - fetchall_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchall") + fetchall_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.fetchall" + ) fetchall_mock.return_value = [("1",), ("2",), ("4",)] insert_table_name = make_temp_table_name("new_records", "abcd") @@ -1325,8 +1424,7 @@ def test_insert_overwrite_by_condition_inc_by_partition( def test_to_time_column(): # we should get DateTime64(6) back for any temporal type other than explicit DateTime64 - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -1336,8 +1434,7 @@ def test_to_time_column(): ); SELECT ds::datetime - """ - ) + """) model = load_sql_based_model(expressions) assert ( model.convert_to_time_column("2022-01-01 00:00:00.000001").sql("clickhouse") @@ -1345,8 +1442,7 @@ def test_to_time_column(): ) # We should respect the user's DateTime64 precision if specified - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -1356,8 +1452,7 @@ def test_to_time_column(): ); SELECT ds::DateTime64(4) - """ - ) + """) model = load_sql_based_model(expressions) assert ( model.convert_to_time_column("2022-01-01 00:00:00.000001").sql("clickhouse") @@ -1367,8 +1462,7 @@ def test_to_time_column(): # We should respect the user's DateTime64 precision if specified, even if we're making it nullable from sqlmesh.utils.date import to_time_column - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -1378,8 +1472,7 @@ def test_to_time_column(): ); SELECT ds::DateTime64(4) - """ - ) + """) model = load_sql_based_model(expressions) assert ( to_time_column( @@ -1393,16 +1486,23 @@ def test_to_time_column(): def test_exchange_tables( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): - from clickhouse_connect.driver.exceptions import DatabaseError # type: ignore + from clickhouse_connect.driver.exceptions import \ + DatabaseError # type: ignore adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) temp_table_mock.return_value = make_temp_table_name("table1", "abcd") - execute_mock = mocker.patch("sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.execute") + execute_mock = mocker.patch( + "sqlmesh.core.engine_adapter.ClickhouseEngineAdapter.execute" + ) execute_mock.side_effect = [ DatabaseError( "DB::Exception: Moving tables between databases of different engines is not supported. (NOT_IMPLEMENTED)" @@ -1416,9 +1516,11 @@ def test_exchange_tables( # The EXCHANGE TABLES call errored, so we RENAME TABLE instead assert [ - quote_identifiers(call.args[0]).sql("clickhouse") - if isinstance(call.args[0], exp.Expr) - else call.args[0] + ( + quote_identifiers(call.args[0]).sql("clickhouse") + if isinstance(call.args[0], exp.Expr) + else call.args[0] + ) for call in execute_mock.call_args_list ] == [ 'EXCHANGE TABLES "table1" AND "table2"', @@ -1430,7 +1532,8 @@ def test_exchange_tables( def test_virtual_catalog_ddl_stripping(make_mocked_engine_adapter: t.Callable): """After inject_virtual_catalog(), create_schema() with the virtual catalog prefix must strip - the catalog and execute without raising, and with a wrong catalog must raise SQLMeshError.""" + the catalog and execute without raising, and with a wrong catalog must raise SQLMeshError. + """ from sqlmesh.utils.errors import SQLMeshError adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) @@ -1466,7 +1569,9 @@ def test_supports_virtual_catalog_returns_true(): assert adapter._default_catalog is None -def test_inject_virtual_catalog_uses_custom_config(make_mocked_engine_adapter: t.Callable): +def test_inject_virtual_catalog_uses_custom_config( + make_mocked_engine_adapter: t.Callable, +): """When virtual_catalog is set in _extra_config, inject_virtual_catalog uses that value instead of the synthetic __gateway_name__ default.""" adapter = make_mocked_engine_adapter( @@ -1505,13 +1610,19 @@ def test_clickhouse_connection_config_virtual_catalog_empty_string_rejected(): from sqlmesh.utils.errors import ConfigError with pytest.raises(ConfigError, match="virtual_catalog cannot be an empty string"): - ClickhouseConnectionConfig(host="localhost", username="user", virtual_catalog="") + ClickhouseConnectionConfig( + host="localhost", username="user", virtual_catalog="" + ) with pytest.raises(ConfigError, match="virtual_catalog cannot be an empty string"): - ClickhouseConnectionConfig(host="localhost", username="user", virtual_catalog=" ") + ClickhouseConnectionConfig( + host="localhost", username="user", virtual_catalog=" " + ) -def test_virtual_catalog_stripped_in_delete_from(make_mocked_engine_adapter: t.Callable): +def test_virtual_catalog_stripped_in_delete_from( + make_mocked_engine_adapter: t.Callable, +): """delete_from() must strip the virtual catalog prefix before building the DELETE expression.""" adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) adapter.inject_virtual_catalog("ch_gw") @@ -1564,7 +1675,11 @@ def test_virtual_catalog_stripped_in_insert_overwrite_by_partition( assert source_query_mock.called target_table_kwarg = source_query_mock.call_args[1].get( "target_table", - source_query_mock.call_args[0][2] if len(source_query_mock.call_args[0]) > 2 else None, + ( + source_query_mock.call_args[0][2] + if len(source_query_mock.call_args[0]) > 2 + else None + ), ) if target_table_kwarg is not None: target_sql = ( @@ -1575,7 +1690,9 @@ def test_virtual_catalog_stripped_in_insert_overwrite_by_partition( assert "__ch_gw__" not in target_sql -def test_virtual_catalog_stripped_in_alter_table(make_mocked_engine_adapter: t.Callable): +def test_virtual_catalog_stripped_in_alter_table( + make_mocked_engine_adapter: t.Callable, +): """alter_table() must strip the virtual catalog prefix from each ALTER TABLE statement.""" adapter = make_mocked_engine_adapter(ClickhouseEngineAdapter) adapter.inject_virtual_catalog("ch_gw") @@ -1604,7 +1721,9 @@ def test_virtual_catalog_stripped_from_create_view_source( cluster="my_cluster", ) adapter.inject_virtual_catalog("clickhouse_gw") - query = parse_one("SELECT * FROM __clickhouse_gw__.my_db.my_db__connection_test__1234567890") + query = parse_one( + "SELECT * FROM __clickhouse_gw__.my_db.my_db__connection_test__1234567890" + ) adapter.create_view( "__clickhouse_gw__.my_db.connection_test__dev", diff --git a/tests/core/engine_adapter/test_databricks.py b/tests/core/engine_adapter/test_databricks.py index 65d8e0388f..7f5b3dd02c 100644 --- a/tests/core/engine_adapter/test_databricks.py +++ b/tests/core/engine_adapter/test_databricks.py @@ -21,13 +21,16 @@ def _query_tags_map(*items: t.Optional[str]) -> exp.Map: keys=exp.Array(expressions=[exp.Literal.string(item) for item in items[::2]]), values=exp.Array( expressions=[ - exp.Null() if item is None else exp.Literal.string(item) for item in items[1::2] + exp.Null() if item is None else exp.Literal.string(item) + for item in items[1::2] ] ), ) -def test_replace_query_not_exists(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_replace_query_not_exists( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.table_exists", return_value=False, @@ -35,7 +38,9 @@ def test_replace_query_not_exists(mocker: MockFixture, make_mocked_engine_adapte mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) adapter.replace_query( "test_table", parse_one("SELECT a FROM tbl"), {"a": exp.DataType.build("INT")} ) @@ -45,7 +50,9 @@ def test_replace_query_not_exists(mocker: MockFixture, make_mocked_engine_adapte ] -def test_replace_query_exists(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_replace_query_exists( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.table_exists", return_value=True, @@ -53,7 +60,9 @@ def test_replace_query_exists(mocker: MockFixture, make_mocked_engine_adapter: t mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) mocker.patch.object( adapter, "_get_data_objects", @@ -76,10 +85,14 @@ def test_replace_query_pandas_not_exists( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) adapter.replace_query( - "test_table", df, {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")} + "test_table", + df, + {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")}, ) assert to_sql_calls(adapter) == [ @@ -87,7 +100,9 @@ def test_replace_query_pandas_not_exists( ] -def test_replace_query_pandas_exists(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_replace_query_pandas_exists( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.table_exists", return_value=True, @@ -95,7 +110,9 @@ def test_replace_query_pandas_exists(mocker: MockFixture, make_mocked_engine_ada mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) mocker.patch.object( adapter, "_get_data_objects", @@ -103,7 +120,9 @@ def test_replace_query_pandas_exists(mocker: MockFixture, make_mocked_engine_ada ) df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) adapter.replace_query( - "test_table", df, {"a": exp.DataType.build("int"), "b": exp.DataType.build("int")} + "test_table", + df, + {"a": exp.DataType.build("int"), "b": exp.DataType.build("int")}, ) assert to_sql_calls(adapter) == [ @@ -115,25 +134,35 @@ def test_clone_table(mocker: MockFixture, make_mocked_engine_adapter: t.Callable mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) adapter.clone_table("target_table", "source_table") adapter.cursor.execute.assert_called_once_with( "CREATE TABLE IF NOT EXISTS `target_table` SHALLOW CLONE `source_table`" ) -def test_set_current_catalog(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") +def test_set_current_catalog( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) adapter.set_current_catalog("test_catalog2") assert to_sql_calls(adapter) == ["USE CATALOG `test_catalog2`"] -def test_session_query_tags(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_session_query_tags( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) with adapter.session( { @@ -159,9 +188,13 @@ def test_session_query_tags_allow_none_values( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) - with adapter.session({"query_tags": _query_tags_map("team", "data-eng", "feature", None)}): + with adapter.session( + {"query_tags": _query_tags_map("team", "data-eng", "feature", None)} + ): adapter.execute("SELECT 1") adapter.cursor.execute.assert_called_with( @@ -175,12 +208,16 @@ def test_session_query_tags_do_not_override_explicit_query_tags( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) with adapter.session({"query_tags": _query_tags_map("team", "data-eng")}): adapter.execute("SELECT 1", query_tags={"team": "analytics"}) - adapter.cursor.execute.assert_called_with("SELECT 1", query_tags={"team": "analytics"}) + adapter.cursor.execute.assert_called_with( + "SELECT 1", query_tags={"team": "analytics"} + ) def test_session_query_tags_not_applied_to_spark_session_connection( @@ -189,7 +226,9 @@ def test_session_query_tags_not_applied_to_spark_session_connection( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) mocker.patch.object( DatabricksEngineAdapter, "is_spark_session_connection", @@ -209,7 +248,9 @@ def test_session_query_tags_not_applied_to_spark_engine_adapter( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) spark_cursor = mocker.Mock() adapter._spark_engine_adapter = mocker.Mock(cursor=spark_cursor) adapter._connection_pool.set_attribute("use_spark_engine_adapter", True) @@ -236,37 +277,51 @@ def test_session_query_tags_not_applied_to_spark_engine_adapter( ], ) def test_session_query_tags_invalid(query_tags, make_mocked_engine_adapter: t.Callable): - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) with pytest.raises(SQLMeshError, match="session_properties.query_tags"): with adapter.session({"query_tags": query_tags}): pass -def test_get_current_catalog(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_get_current_catalog( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) adapter.cursor.fetchone.return_value = ("test_catalog",) assert adapter.get_current_catalog() == "test_catalog" assert to_sql_calls(adapter) == ["SELECT CURRENT_CATALOG()"] -def test_get_current_schema(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_get_current_schema( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) adapter.cursor.fetchone.return_value = ("test_database",) assert adapter._get_current_schema() == "test_database" assert to_sql_calls(adapter) == ["SELECT CURRENT_DATABASE()"] -def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: MockFixture): - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="main") +def test_sync_grants_config( + make_mocked_engine_adapter: t.Callable, mocker: MockFixture +): + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="main" + ) relation = exp.to_table("main.test_schema.test_table", dialect="databricks") new_grants_config = { "SELECT": ["group1", "group2"], @@ -277,7 +332,9 @@ def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: Mock ("SELECT", "legacy"), ("REFRESH", "stale"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -294,17 +351,34 @@ def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: Mock sql_calls = to_sql_calls(adapter) assert len(sql_calls) == 5 - assert "GRANT SELECT ON TABLE `main`.`test_schema`.`test_table` TO `group1`" in sql_calls - assert "GRANT SELECT ON TABLE `main`.`test_schema`.`test_table` TO `group2`" in sql_calls - assert "GRANT MODIFY ON TABLE `main`.`test_schema`.`test_table` TO `writers`" in sql_calls - assert "REVOKE SELECT ON TABLE `main`.`test_schema`.`test_table` FROM `legacy`" in sql_calls - assert "REVOKE REFRESH ON TABLE `main`.`test_schema`.`test_table` FROM `stale`" in sql_calls + assert ( + "GRANT SELECT ON TABLE `main`.`test_schema`.`test_table` TO `group1`" + in sql_calls + ) + assert ( + "GRANT SELECT ON TABLE `main`.`test_schema`.`test_table` TO `group2`" + in sql_calls + ) + assert ( + "GRANT MODIFY ON TABLE `main`.`test_schema`.`test_table` TO `writers`" + in sql_calls + ) + assert ( + "REVOKE SELECT ON TABLE `main`.`test_schema`.`test_table` FROM `legacy`" + in sql_calls + ) + assert ( + "REVOKE REFRESH ON TABLE `main`.`test_schema`.`test_table` FROM `stale`" + in sql_calls + ) def test_sync_grants_config_with_overlaps( make_mocked_engine_adapter: t.Callable, mocker: MockFixture ): - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="main") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="main" + ) relation = exp.to_table("main.test_schema.test_table", dialect="databricks") new_grants_config = { "SELECT": ["shared", "new_role"], @@ -316,7 +390,9 @@ def test_sync_grants_config_with_overlaps( ("SELECT", "legacy"), ("MODIFY", "shared"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -333,9 +409,18 @@ def test_sync_grants_config_with_overlaps( sql_calls = to_sql_calls(adapter) assert len(sql_calls) == 3 - assert "GRANT SELECT ON TABLE `main`.`test_schema`.`test_table` TO `new_role`" in sql_calls - assert "GRANT MODIFY ON TABLE `main`.`test_schema`.`test_table` TO `writer`" in sql_calls - assert "REVOKE SELECT ON TABLE `main`.`test_schema`.`test_table` FROM `legacy`" in sql_calls + assert ( + "GRANT SELECT ON TABLE `main`.`test_schema`.`test_table` TO `new_role`" + in sql_calls + ) + assert ( + "GRANT MODIFY ON TABLE `main`.`test_schema`.`test_table` TO `writer`" + in sql_calls + ) + assert ( + "REVOKE SELECT ON TABLE `main`.`test_schema`.`test_table` FROM `legacy`" + in sql_calls + ) @pytest.mark.parametrize( @@ -353,7 +438,9 @@ def test_sync_grants_config_object_kind( table_type: DataObjectType, expected_keyword: str, ) -> None: - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="main") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="main" + ) relation = exp.to_table("main.test_schema.test_object", dialect="databricks") mocker.patch.object(adapter, "fetchall", return_value=[]) @@ -366,9 +453,15 @@ def test_sync_grants_config_object_kind( ] -def test_sync_grants_config_quotes(make_mocked_engine_adapter: t.Callable, mocker: MockFixture): - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="`test_db`") - relation = exp.to_table("`test_db`.`test_schema`.`test_table`", dialect="databricks") +def test_sync_grants_config_quotes( + make_mocked_engine_adapter: t.Callable, mocker: MockFixture +): + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="`test_db`" + ) + relation = exp.to_table( + "`test_db`.`test_schema`.`test_table`", dialect="databricks" + ) new_grants_config = { "SELECT": ["group1", "group2"], "MODIFY": ["writers"], @@ -378,7 +471,9 @@ def test_sync_grants_config_quotes(make_mocked_engine_adapter: t.Callable, mocke ("SELECT", "legacy"), ("REFRESH", "stale"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -395,17 +490,34 @@ def test_sync_grants_config_quotes(make_mocked_engine_adapter: t.Callable, mocke sql_calls = to_sql_calls(adapter) assert len(sql_calls) == 5 - assert "GRANT SELECT ON TABLE `test_db`.`test_schema`.`test_table` TO `group1`" in sql_calls - assert "GRANT SELECT ON TABLE `test_db`.`test_schema`.`test_table` TO `group2`" in sql_calls - assert "GRANT MODIFY ON TABLE `test_db`.`test_schema`.`test_table` TO `writers`" in sql_calls - assert "REVOKE SELECT ON TABLE `test_db`.`test_schema`.`test_table` FROM `legacy`" in sql_calls - assert "REVOKE REFRESH ON TABLE `test_db`.`test_schema`.`test_table` FROM `stale`" in sql_calls + assert ( + "GRANT SELECT ON TABLE `test_db`.`test_schema`.`test_table` TO `group1`" + in sql_calls + ) + assert ( + "GRANT SELECT ON TABLE `test_db`.`test_schema`.`test_table` TO `group2`" + in sql_calls + ) + assert ( + "GRANT MODIFY ON TABLE `test_db`.`test_schema`.`test_table` TO `writers`" + in sql_calls + ) + assert ( + "REVOKE SELECT ON TABLE `test_db`.`test_schema`.`test_table` FROM `legacy`" + in sql_calls + ) + assert ( + "REVOKE REFRESH ON TABLE `test_db`.`test_schema`.`test_table` FROM `stale`" + in sql_calls + ) def test_sync_grants_config_no_catalog_or_schema( make_mocked_engine_adapter: t.Callable, mocker: MockFixture ): - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="main_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="main_catalog" + ) relation = exp.to_table("test_table", dialect="databricks") new_grants_config = { "SELECT": ["group1", "group2"], @@ -416,7 +528,9 @@ def test_sync_grants_config_no_catalog_or_schema( ("SELECT", "legacy"), ("REFRESH", "stale"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) mocker.patch.object(adapter, "_get_current_schema", return_value="schema") mocker.patch.object(adapter, "get_current_catalog", return_value="main_catalog") @@ -443,14 +557,20 @@ def test_sync_grants_config_no_catalog_or_schema( def test_insert_overwrite_by_partition_query( - make_mocked_engine_adapter: t.Callable, mocker: MockFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test_schema.test_table" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -477,11 +597,15 @@ def test_insert_overwrite_by_partition_query( ] -def test_materialized_view_properties(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_materialized_view_properties( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) adapter.create_view( "test_table", @@ -508,7 +632,9 @@ def test_materialized_view_with_column_comments( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) mocker.patch.object(adapter, "get_current_catalog", return_value="test_catalog") adapter.create_view( @@ -538,7 +664,9 @@ def test_regular_view_with_column_comments( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) mocker.patch.object(adapter, "get_current_catalog", return_value="test_catalog") adapter.create_view( @@ -562,11 +690,15 @@ def test_regular_view_with_column_comments( ] -def test_create_table_clustered_by(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): +def test_create_table_clustered_by( + mocker: MockFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) columns_to_types = { "cola": exp.DataType.build("INT"), @@ -591,7 +723,9 @@ def test_create_table_clustered_by_keyword( mocker.patch( "sqlmesh.core.engine_adapter.databricks.DatabricksEngineAdapter.set_current_catalog" ) - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) columns_to_types = { "cola": exp.DataType.build("INT"), @@ -680,7 +814,9 @@ def test_drop_data_object_materialized_view_calls_correct_drop(mocker: MockFixtu def test_columns(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): - adapter = make_mocked_engine_adapter(DatabricksEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + DatabricksEngineAdapter, default_catalog="test_catalog" + ) # Override/mock get_current_catalog to return default current_catalog_mock = mocker.patch.object( @@ -720,10 +856,14 @@ def test_columns(mocker: MockFixture, make_mocked_engine_adapter: t.Callable): "small_int": exp.DataType.build("smallint", dialect=adapter.dialect), "string_col": exp.DataType.build("string", dialect=adapter.dialect), "timestamp_col": exp.DataType.build("timestamp", dialect=adapter.dialect), - "timestamp_ntz_col": exp.DataType.build("timestamp_ntz", dialect=adapter.dialect), + "timestamp_ntz_col": exp.DataType.build( + "timestamp_ntz", dialect=adapter.dialect + ), "tinyint_col": exp.DataType.build("tinyint", dialect=adapter.dialect), "array_col": exp.DataType.build("array", dialect=adapter.dialect), - "simple_struct_col": exp.DataType.build("struct", dialect=adapter.dialect), + "simple_struct_col": exp.DataType.build( + "struct", dialect=adapter.dialect + ), "long_struct_col": exp.DataType.build( f"struct<{','.join(long_struct_cols)}>", dialect=adapter.dialect ), @@ -757,7 +897,10 @@ def _make_databricks_connect_adapter( mock_databricks_module = types.ModuleType("databricks") mocker.patch.dict( sys.modules, - {"databricks": mock_databricks_module, "databricks.connect": mock_connect_module}, + { + "databricks": mock_databricks_module, + "databricks.connect": mock_connect_module, + }, ) mocker.patch( diff --git a/tests/core/engine_adapter/test_duckdb.py b/tests/core/engine_adapter/test_duckdb.py index 9fd65a6e66..f52ad710c5 100644 --- a/tests/core/engine_adapter/test_duckdb.py +++ b/tests/core/engine_adapter/test_duckdb.py @@ -5,6 +5,7 @@ from pytest_mock.plugin import MockerFixture from sqlglot import expressions as exp from sqlglot import parse_one + from sqlmesh.core.engine_adapter import DuckDBEngineAdapter, EngineAdapter from tests.core.engine_adapter import to_sql_calls @@ -32,7 +33,9 @@ def test_create_schema(adapter: EngineAdapter, duck_conn): "SELECT 1 FROM information_schema.schemata WHERE schema_name = 'test_schema'" ).fetchall() == [(1,)] with pytest.raises(Exception): - adapter.create_schema("test_schema", ignore_if_exists=False, warn_on_error=False) + adapter.create_schema( + "test_schema", ignore_if_exists=False, warn_on_error=False + ) def test_table_exists(adapter: EngineAdapter, duck_conn): @@ -63,7 +66,9 @@ def test_create_table(adapter: EngineAdapter, duck_conn): def test_replace_query_pandas(adapter: EngineAdapter, duck_conn): df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) adapter.replace_query( - "test_table", df, {"a": exp.DataType.build("long"), "b": exp.DataType.build("long")} + "test_table", + df, + {"a": exp.DataType.build("long"), "b": exp.DataType.build("long")}, ) pd.testing.assert_frame_equal(adapter.fetchdf("SELECT * FROM test_table"), df) @@ -86,7 +91,9 @@ def test_temporary_table(make_mocked_engine_adapter: t.Callable, mocker: MockerF adapter.create_table( "test_table", {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")}, - table_properties={"creatable_type": exp.Column(this=exp.Identifier(this="Temporary"))}, + table_properties={ + "creatable_type": exp.Column(this=exp.Identifier(this="Temporary")) + }, ) assert to_sql_calls(adapter) == [ diff --git a/tests/core/engine_adapter/test_fabric.py b/tests/core/engine_adapter/test_fabric.py index d16e973e8a..061d4664cd 100644 --- a/tests/core/engine_adapter/test_fabric.py +++ b/tests/core/engine_adapter/test_fabric.py @@ -8,8 +8,8 @@ from sqlglot import exp, parse_one from sqlmesh.core.engine_adapter import FabricEngineAdapter -from tests.core.engine_adapter import to_sql_calls from sqlmesh.core.engine_adapter.shared import DataObject +from tests.core.engine_adapter import to_sql_calls pytestmark = [pytest.mark.engine, pytest.mark.fabric] @@ -209,7 +209,10 @@ def test_insert_overwrite_by_time_partition(adapter: FabricEngineAdapter): end="2022-01-02", time_column="b", time_formatter=lambda x, _: exp.Literal.string(x.strftime("%Y-%m-%d")), - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("STRING"), + }, ) # Fabric adapter should use DELETE/INSERT strategy, not MERGE. @@ -236,7 +239,9 @@ def test_replace_query(adapter: FabricEngineAdapter, mocker: MockerFixture): ] -def test_alter_table_column_type_workaround(adapter: FabricEngineAdapter, mocker: MockerFixture): +def test_alter_table_column_type_workaround( + adapter: FabricEngineAdapter, mocker: MockerFixture +): """ Tests the alter_table method's workaround for changing a column's data type. """ @@ -269,7 +274,9 @@ def test_alter_table_column_type_workaround(adapter: FabricEngineAdapter, mocker assert to_sql_calls(adapter) == expected_calls -def test_alter_table_direct_alteration(adapter: FabricEngineAdapter, mocker: MockerFixture): +def test_alter_table_direct_alteration( + adapter: FabricEngineAdapter, mocker: MockerFixture +): """ Tests the alter_table method for direct alterations like adding a column. """ @@ -277,7 +284,9 @@ def test_alter_table_direct_alteration(adapter: FabricEngineAdapter, mocker: Moc alter_expression = exp.Alter( this=exp.to_table("my_db.my_schema.my_table"), - actions=[exp.ColumnDef(this=exp.to_column("new_col"), kind=exp.DataType.build("INT"))], + actions=[ + exp.ColumnDef(this=exp.to_column("new_col"), kind=exp.DataType.build("INT")) + ], ) adapter.alter_table([alter_expression]) @@ -292,7 +301,9 @@ def test_alter_table_direct_alteration(adapter: FabricEngineAdapter, mocker: Moc def test_merge_pandas( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.fabric.FabricEngineAdapter.table_exists", @@ -301,7 +312,9 @@ def test_merge_pandas( adapter = make_mocked_engine_adapter(FabricEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -355,7 +368,9 @@ def test_merge_pandas( def test_merge_exists( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.fabric.FabricEngineAdapter.table_exists", @@ -364,7 +379,9 @@ def test_merge_exists( adapter = make_mocked_engine_adapter(FabricEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -438,7 +455,9 @@ def test_comments(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture) comment = "\\" create_table_comment_mock = mocker.patch.object(adapter, "_create_table_comment") - create_column_comments_mock = mocker.patch.object(adapter, "_create_column_comments") + create_column_comments_mock = mocker.patch.object( + adapter, "_create_column_comments" + ) mocker.patch.object(adapter, "_create_table") adapter.create_table( diff --git a/tests/core/engine_adapter/test_mixins.py b/tests/core/engine_adapter/test_mixins.py index 50bef59d6e..fa8800c1a7 100644 --- a/tests/core/engine_adapter/test_mixins.py +++ b/tests/core/engine_adapter/test_mixins.py @@ -6,10 +6,8 @@ from pytest_mock.plugin import MockerFixture from sqlglot import exp, parse_one -from sqlmesh.core.engine_adapter.mixins import ( - LogicalMergeMixin, - NonTransactionalTruncateMixin, -) +from sqlmesh.core.engine_adapter.mixins import (LogicalMergeMixin, + NonTransactionalTruncateMixin) from tests.core.engine_adapter import to_sql_calls pytestmark = pytest.mark.engine @@ -17,7 +15,9 @@ def test_logical_merge(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): adapter = make_mocked_engine_adapter(LogicalMergeMixin, "duckdb") - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table" + ) temp_table_mock.return_value = exp.to_table("temporary") adapter.merge( @@ -36,7 +36,9 @@ def test_logical_merge(make_mocked_engine_adapter: t.Callable, mocker: MockerFix call( '''CREATE TABLE "temporary" AS SELECT CAST("id" AS INT) AS "id", CAST("ts" AS TIMESTAMP) AS "ts", CAST("val" AS INT) AS "val" FROM (SELECT "id", "ts", "val" FROM "source") AS "_subquery"''' ), - call("""DELETE FROM "target" WHERE "id" IN (SELECT "id" FROM "temporary")"""), + call( + """DELETE FROM "target" WHERE "id" IN (SELECT "id" FROM "temporary")""" + ), call( 'INSERT INTO "target" ("id", "ts", "val") SELECT DISTINCT ON ("id") "id", "ts", "val" FROM "temporary"' ), @@ -73,7 +75,9 @@ def test_logical_merge(make_mocked_engine_adapter: t.Callable, mocker: MockerFix def test_non_transaction_truncate_mixin( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(NonTransactionalTruncateMixin, "redshift") adapter._truncate_table(table_name="test_table") @@ -82,7 +86,9 @@ def test_non_transaction_truncate_mixin( def test_non_transaction_truncate_mixin_within_transaction( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): adapter = make_mocked_engine_adapter(NonTransactionalTruncateMixin, "redshift") adapter._connection_pool = mocker.MagicMock() diff --git a/tests/core/engine_adapter/test_mssql.py b/tests/core/engine_adapter/test_mssql.py index 007a59b365..21b678008e 100644 --- a/tests/core/engine_adapter/test_mssql.py +++ b/tests/core/engine_adapter/test_mssql.py @@ -1,6 +1,7 @@ # type: ignore import typing as t from datetime import date +from pathlib import Path from unittest import mock import pandas as pd # noqa: TID253 @@ -9,14 +10,15 @@ from sqlglot import expressions as exp from sqlglot import parse_one -from pathlib import Path from sqlmesh import model +from sqlmesh.core import dialect as d from sqlmesh.core.engine_adapter.mssql import MSSQLEngineAdapter -from sqlmesh.core.snapshot import SnapshotEvaluator, SnapshotChangeCategory, Snapshot +from sqlmesh.core.engine_adapter.shared import (DataObject, DataObjectType, + SourceQuery) from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.model.kind import SCDType2ByTimeKind -from sqlmesh.core import dialect as d -from sqlmesh.core.engine_adapter.shared import DataObject, DataObjectType, SourceQuery +from sqlmesh.core.snapshot import (Snapshot, SnapshotChangeCategory, + SnapshotEvaluator) from sqlmesh.utils.date import to_ds from tests.core.engine_adapter import to_sql_calls @@ -82,7 +84,9 @@ def test_columns(adapter: MSSQLEngineAdapter): ) -def test_varchar_workaround_to_max(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_varchar_workaround_to_max( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(MSSQLEngineAdapter) columns = { @@ -202,8 +206,7 @@ def test_table_exists(make_mocked_engine_adapter: t.Callable): ], ) def test_to_time_column(select_expr, input_time, expected_sql): - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -213,8 +216,7 @@ def test_to_time_column(select_expr, input_time, expected_sql): ); {select_expr} - """ - ) + """) model = load_sql_based_model(expressions) assert model.convert_to_time_column(input_time).sql("tsql") == expected_sql @@ -266,7 +268,9 @@ def test_incremental_by_time_datetimeoffset_precision( def test_insert_overwrite_by_time_partition_supports_insert_overwrite_pandas_not_exists( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.mssql.MSSQLEngineAdapter.table_exists", @@ -306,7 +310,9 @@ def test_insert_overwrite_by_time_partition_supports_insert_overwrite_pandas_not def test_insert_overwrite_by_time_partition_supports_insert_overwrite_pandas_exists( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.mssql.MSSQLEngineAdapter.table_exists", @@ -342,7 +348,9 @@ def test_insert_overwrite_by_time_partition_supports_insert_overwrite_pandas_exi def test_insert_append_pandas( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.mssql.MSSQLEngineAdapter.table_exists", @@ -351,7 +359,9 @@ def test_insert_append_pandas( adapter = make_mocked_engine_adapter(MSSQLEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test_table" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -410,7 +420,9 @@ def test_create_physical_properties(make_mocked_engine_adapter: t.Callable): def test_merge_pandas( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.mssql.MSSQLEngineAdapter.table_exists", @@ -419,7 +431,9 @@ def test_merge_pandas( adapter = make_mocked_engine_adapter(MSSQLEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -473,7 +487,9 @@ def test_merge_pandas( def test_merge_exists( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.mssql.MSSQLEngineAdapter.table_exists", @@ -482,7 +498,9 @@ def test_merge_exists( adapter = make_mocked_engine_adapter(MSSQLEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "target" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -567,7 +585,9 @@ def test_replace_query(make_mocked_engine_adapter: t.Callable, mocker: MockerFix def test_replace_query_pandas( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): temp_table_exists_counter = 0 @@ -584,7 +604,9 @@ def test_replace_query_pandas( ) adapter.cursor.fetchone.return_value = (1,) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test_table" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -644,7 +666,9 @@ def test_create_index(make_mocked_engine_adapter: t.Callable): ) -def test_drop_schema_with_catalog(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_drop_schema_with_catalog( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(MSSQLEngineAdapter) adapter.get_current_catalog = mocker.MagicMock(return_value="other_catalog") @@ -658,8 +682,12 @@ def test_drop_schema_with_catalog(make_mocked_engine_adapter: t.Callable, mocker ] -def test_get_data_objects_catalog(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): - adapter = make_mocked_engine_adapter(MSSQLEngineAdapter, patch_get_data_objects=False) +def test_get_data_objects_catalog( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): + adapter = make_mocked_engine_adapter( + MSSQLEngineAdapter, patch_get_data_objects=False + ) original_set_current_catalog = adapter.set_current_catalog local_state = {} @@ -673,7 +701,9 @@ def set_local_catalog(catalog, local_state): adapter.set_current_catalog = mocker.MagicMock( side_effect=lambda x: set_local_catalog(x, local_state) ) - adapter.cursor.fetchall.return_value = [("test_catalog", "test_table", "test_schema", "TABLE")] + adapter.cursor.fetchall.return_value = [ + ("test_catalog", "test_table", "test_schema", "TABLE") + ] adapter.cursor.description = [["catalog_name"], ["name"], ["schema_name"], ["type"]] result = adapter.get_data_objects("test_catalog.test_schema") @@ -794,7 +824,9 @@ def test_rename_table(make_mocked_engine_adapter: t.Callable, mocker: MockerFixt ] -def test_create_table_from_query(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_create_table_from_query( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(MSSQLEngineAdapter) mocker.patch( "sqlmesh.core.engine_adapter.base.random_id", @@ -830,7 +862,9 @@ def test_create_table_from_query(make_mocked_engine_adapter: t.Callable, mocker: "CREATE TABLE [test_schema].[test_table] ([a] VARCHAR(MAX), [b] VARCHAR(60), [c] VARCHAR(MAX), [d] VARCHAR(MAX), [e] DATETIME2);", ] - columns_mock.assert_called_once_with(exp.table_("__temp_ctas_test_random_id", quoted=True)) + columns_mock.assert_called_once_with( + exp.table_("__temp_ctas_test_random_id", quoted=True) + ) # We don't want to drop anything other than LIMIT 0 # See https://github.com/SQLMesh/sqlmesh/issues/4048 @@ -850,8 +884,7 @@ def test_create_table_from_query(make_mocked_engine_adapter: t.Callable, mocker: def test_replace_query_strategy(adapter: MSSQLEngineAdapter, mocker: MockerFixture): # ref issue 4472: https://github.com/SQLMesh/sqlmesh/issues/4472 # The FULL strategy calls EngineAdapter.replace_query() which calls _insert_overwrite_by_condition() should use DELETE+INSERT and not MERGE - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, kind FULL, @@ -859,8 +892,7 @@ def test_replace_query_strategy(adapter: MSSQLEngineAdapter, mocker: MockerFixtu ); select a, b from db.upstream_table; - """ - ) + """) model = load_sql_based_model(expressions) exists_mock = mocker.patch( @@ -988,7 +1020,11 @@ def python_scd2_model(context, **kwargs): import pandas as pd return pd.DataFrame( - {"id": [1, 2], "value": ["a", "b"], "updated_at": ["2024-01-01", "2024-01-02"]} + { + "id": [1, 2], + "value": ["a", "b"], + "updated_at": ["2024-01-01", "2024-01-02"], + } ) m = model.get_registry()["test_schema.python_scd2_with_mssql_merge"].model( diff --git a/tests/core/engine_adapter/test_mysql.py b/tests/core/engine_adapter/test_mysql.py index d09b1c1780..6888a05dd8 100644 --- a/tests/core/engine_adapter/test_mysql.py +++ b/tests/core/engine_adapter/test_mysql.py @@ -20,7 +20,9 @@ def test_comments(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture) truncated_column_comment = "c" * allowed_column_comment_length long_column_comment = truncated_column_comment + "d" - fetchone_mock = mocker.patch("sqlmesh.core.engine_adapter.mysql.MySQLEngineAdapter.fetchone") + fetchone_mock = mocker.patch( + "sqlmesh.core.engine_adapter.mysql.MySQLEngineAdapter.fetchone" + ) fetchone_mock.return_value = ["test_table", "CREATE TABLE test_table (a INT)"] adapter.create_table( @@ -92,7 +94,9 @@ def test_replace_by_key_composite_uses_join_delete( ): """Composite key DELETE uses JOIN instead of CONCAT_WS to allow index usage.""" adapter = make_mocked_engine_adapter(MySQLEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table" + ) temp_table_mock.return_value = exp.to_table("temporary") adapter.merge( @@ -109,12 +113,12 @@ def test_replace_by_key_composite_uses_join_delete( sql_calls = to_sql_calls(adapter) # The DELETE should use a JOIN instead of CONCAT_WS - assert any("CONCAT_WS" in s for s in sql_calls) is False, ( - "DELETE should not use CONCAT_WS for composite keys" - ) - assert any("INNER JOIN" in s for s in sql_calls) is True, ( - "DELETE should use INNER JOIN for composite keys" - ) + assert ( + any("CONCAT_WS" in s for s in sql_calls) is False + ), "DELETE should not use CONCAT_WS for composite keys" + assert ( + any("INNER JOIN" in s for s in sql_calls) is True + ), "DELETE should use INNER JOIN for composite keys" # Verify the full sequence of SQL calls adapter.cursor.execute.assert_has_calls( @@ -138,12 +142,16 @@ def test_replace_by_key_three_column_composite_key( ): """3-column composite key matching the original issue scenario (#5711).""" adapter = make_mocked_engine_adapter(MySQLEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table" + ) temp_table_mock.return_value = exp.to_table("temporary") adapter.merge( target_table="target", - source_table=t.cast(exp.Select, parse_one("SELECT id, region, ts, val FROM source")), + source_table=t.cast( + exp.Select, parse_one("SELECT id, region, ts, val FROM source") + ), target_columns_to_types={ "id": exp.DataType(this=exp.DataType.Type.INT), "region": exp.DataType(this=exp.DataType.Type.VARCHAR), @@ -172,7 +180,9 @@ def test_replace_by_key_expression_based_composite_key( ): """Expression-based composite keys (e.g. DATE_TRUNC + column) via insert_overwrite_by_partition.""" adapter = make_mocked_engine_adapter(MySQLEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table" + ) temp_table_mock.return_value = exp.to_table("temporary") adapter.insert_overwrite_by_partition( @@ -207,7 +217,9 @@ def test_replace_by_key_single_key_uses_in( ): """Single key DELETE still uses the IN-based approach (indexes work fine for single column).""" adapter = make_mocked_engine_adapter(MySQLEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.base.EngineAdapter._get_temp_table" + ) temp_table_mock.return_value = exp.to_table("temporary") adapter.merge( diff --git a/tests/core/engine_adapter/test_postgres.py b/tests/core/engine_adapter/test_postgres.py index ebcdd03f55..2b62505da1 100644 --- a/tests/core/engine_adapter/test_postgres.py +++ b/tests/core/engine_adapter/test_postgres.py @@ -54,13 +54,16 @@ def test_drop_schema(kwargs, expected, make_mocked_engine_adapter: t.Callable): assert to_sql_calls(adapter) == ensure_list(expected) -def test_drop_schema_with_catalog(make_mocked_engine_adapter: t.Callable, mocker: MockFixture): +def test_drop_schema_with_catalog( + make_mocked_engine_adapter: t.Callable, mocker: MockFixture +): adapter = make_mocked_engine_adapter(PostgresEngineAdapter) adapter.get_current_catalog = mocker.MagicMock(return_value="other_catalog") with pytest.raises( - SQLMeshError, match="requires that all catalog operations be against a single catalog" + SQLMeshError, + match="requires that all catalog operations be against a single catalog", ): adapter.drop_schema("test_catalog.test_schema") @@ -114,12 +117,16 @@ def test_merge_version_gte_15(make_mocked_engine_adapter: t.Callable): def test_merge_version_lt_15( - make_mocked_engine_adapter: t.Callable, make_temp_table_name: t.Callable, mocker: MockerFixture + make_mocked_engine_adapter: t.Callable, + make_temp_table_name: t.Callable, + mocker: MockerFixture, ): adapter = make_mocked_engine_adapter(PostgresEngineAdapter) adapter.server_version = (14, 0) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) @@ -152,12 +159,17 @@ def test_alter_table_drop_column_cascade(make_mocked_engine_adapter: t.Callable) def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: if table_name == current_table_name: - return {"id": exp.DataType.build("int"), "test_column": exp.DataType.build("int")} + return { + "id": exp.DataType.build("int"), + "test_column": exp.DataType.build("int"), + } return {"id": exp.DataType.build("int")} adapter.columns = table_columns - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) assert to_sql_calls(adapter) == [ 'ALTER TABLE "test_table" DROP COLUMN "test_column" CASCADE', ] @@ -179,13 +191,17 @@ def test_server_version(make_mocked_engine_adapter: t.Callable, mocker: MockerFi assert adapter.server_version == (15, 13) -def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_sync_grants_config( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(PostgresEngineAdapter) relation = exp.to_table("test_schema.test_table", dialect="postgres") new_grants_config = {"SELECT": ["user1", "user2"], "INSERT": ["user3"]} current_grants = [("SELECT", "old_user"), ("UPDATE", "admin_user")] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -213,7 +229,10 @@ def test_sync_grants_config_with_overlaps( ): adapter = make_mocked_engine_adapter(PostgresEngineAdapter) relation = exp.to_table("test_schema.test_table", dialect="postgres") - new_grants_config = {"SELECT": ["user1", "user2", "user3"], "INSERT": ["user2", "user4"]} + new_grants_config = { + "SELECT": ["user1", "user2", "user3"], + "INSERT": ["user2", "user4"], + } current_grants = [ ("SELECT", "user1"), @@ -221,7 +240,9 @@ def test_sync_grants_config_with_overlaps( ("INSERT", "user2"), ("UPDATE", "user3"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -266,8 +287,12 @@ def test_sync_grants_config_with_default_schema( new_grants_config = {"SELECT": ["user1"], "INSERT": ["user2"]} currrent_grants = [("UPDATE", "old_user")] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=currrent_grants) - get_schema_mock = mocker.patch.object(adapter, "_get_current_schema", return_value="public") + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=currrent_grants + ) + get_schema_mock = mocker.patch.object( + adapter, "_get_current_schema", return_value="public" + ) adapter.sync_grants_config(relation, new_grants_config) diff --git a/tests/core/engine_adapter/test_redshift.py b/tests/core/engine_adapter/test_redshift.py index ddd2c7c2c8..dbb07109ec 100644 --- a/tests/core/engine_adapter/test_redshift.py +++ b/tests/core/engine_adapter/test_redshift.py @@ -1,10 +1,10 @@ # type: ignore import typing as t +from unittest.mock import PropertyMock import pandas as pd # noqa: TID253 import pytest from pytest_mock.plugin import MockerFixture -from unittest.mock import PropertyMock from sqlglot import expressions as exp from sqlglot import parse_one @@ -158,9 +158,7 @@ def test_create_table_physical_properties_from_model_definition( adapter = make_mocked_engine_adapter(RedshiftEngineAdapter) model: SqlModel = t.cast( SqlModel, - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_table, kind full, @@ -171,9 +169,7 @@ def test_create_table_physical_properties_from_model_definition( ) ); SELECT id_file::INT, batch_time::TIMESTAMP; - """ - ) - ), + """)), ) adapter.create_table( @@ -187,7 +183,9 @@ def test_create_table_physical_properties_from_model_definition( ] -def test_varchar_size_workaround(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_varchar_size_workaround( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(RedshiftEngineAdapter) columns = { @@ -238,13 +236,17 @@ def test_varchar_size_workaround(make_mocked_engine_adapter: t.Callable, mocker: ] -def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_sync_grants_config( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(RedshiftEngineAdapter) relation = exp.to_table("test_schema.test_table", dialect="redshift") new_grants_config = {"SELECT": ["user1", "user2"], "INSERT": ["user3"]} current_grants = [("SELECT", "old_user"), ("UPDATE", "legacy_user")] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -281,7 +283,9 @@ def test_sync_grants_config_with_overlaps( ("SELECT", "user_legacy"), ("INSERT", "user_shared"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -327,13 +331,17 @@ def test_sync_grants_config_object_kind( assert sql_calls == [f'GRANT SELECT ON "test_schema"."test_object" TO "user_test"'] -def test_sync_grants_config_quotes(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_sync_grants_config_quotes( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(RedshiftEngineAdapter) relation = exp.to_table('"TestSchema"."TestTable"', dialect="redshift") new_grants_config = {"SELECT": ["user1", "user2"], "INSERT": ["user3"]} current_grants = [("SELECT", "user_old"), ("UPDATE", "user_legacy")] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -363,8 +371,12 @@ def test_sync_grants_config_no_schema( new_grants_config = {"SELECT": ["user1"], "INSERT": ["user2"]} current_grants = [("UPDATE", "user_old")] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) - get_schema_mock = mocker.patch.object(adapter, "_get_current_schema", return_value="public") + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) + get_schema_mock = mocker.patch.object( + adapter, "_get_current_schema", return_value="public" + ) adapter.sync_grants_config(relation, new_grants_config) @@ -424,7 +436,9 @@ def test_create_table_from_query_exists_no_if_not_exists( 'CREATE TABLE "test_schema"."test_table" ("a" VARCHAR(MAX), "b" VARCHAR(60), "c" VARCHAR(MAX), "d" VARCHAR(MAX), "e" TIMESTAMP)', ] - columns_mock.assert_called_once_with(exp.table_("__temp_ctas_test_random_id", quoted=True)) + columns_mock.assert_called_once_with( + exp.table_("__temp_ctas_test_random_id", quoted=True) + ) def test_create_table_recursive_cte(adapter: t.Callable, mocker: MockerFixture): @@ -464,7 +478,9 @@ def test_create_table_recursive_cte(adapter: t.Callable, mocker: MockerFixture): 'CREATE TABLE "test_schema"."test_table" ("a" VARCHAR(MAX), "b" VARCHAR(60), "c" VARCHAR(MAX), "d" VARCHAR(MAX), "e" TIMESTAMP)', ] - columns_mock.assert_called_once_with(exp.table_("__temp_ctas_test_random_id", quoted=True)) + columns_mock.assert_called_once_with( + exp.table_("__temp_ctas_test_random_id", quoted=True) + ) def test_create_table_from_query_exists_and_if_not_exists( @@ -523,7 +539,10 @@ def test_values_to_sql(adapter: t.Callable, mocker: MockerFixture): df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) result = adapter._values_to_sql( values=list(df.itertuples(index=False, name=None)), - target_columns_to_types={"a": exp.DataType.build("int"), "b": exp.DataType.build("int")}, + target_columns_to_types={ + "a": exp.DataType.build("int"), + "b": exp.DataType.build("int"), + }, batch_start=0, batch_end=2, ) @@ -544,7 +563,9 @@ def test_replace_query_with_query(adapter: t.Callable, mocker: MockerFixture): return_value={"cola": exp.DataType(this=exp.DataType.Type.INT)}, ) - adapter.replace_query(table_name="test_table", query_or_df=parse_one("SELECT cola FROM table")) + adapter.replace_query( + table_name="test_table", query_or_df=parse_one("SELECT cola FROM table") + ) assert to_sql_calls(adapter) == [ 'CREATE TABLE "test_table" AS SELECT "cola" FROM "table"', @@ -565,7 +586,9 @@ def mock_table(*args, **kwargs): return f"temp_table_{call_counter}" mock_temp_table = mocker.MagicMock(side_effect=mock_table) - mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table", mock_temp_table) + mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table", mock_temp_table + ) mocker.patch.object( adapter, "_get_data_objects", @@ -593,7 +616,9 @@ def mock_table(*args, **kwargs): ] -def test_replace_query_with_df_table_not_exists(adapter: t.Callable, mocker: MockerFixture): +def test_replace_query_with_df_table_not_exists( + adapter: t.Callable, mocker: MockerFixture +): df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) mocker.patch( "sqlmesh.core.engine_adapter.redshift.RedshiftEngineAdapter.table_exists", @@ -663,12 +688,17 @@ def test_alter_table_drop_column_cascade(adapter: t.Callable): def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: if table_name == current_table_name: - return {"id": exp.DataType.build("int"), "test_column": exp.DataType.build("int")} + return { + "id": exp.DataType.build("int"), + "test_column": exp.DataType.build("int"), + } return {"id": exp.DataType.build("int")} adapter.columns = table_columns - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) assert to_sql_calls(adapter) == [ 'ALTER TABLE "test_table" DROP COLUMN "test_column" CASCADE', ] @@ -691,7 +721,9 @@ def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: adapter.columns = table_columns - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) assert to_sql_calls(adapter) == [ 'ALTER TABLE "test_table" ALTER COLUMN "test_column" TYPE VARCHAR(20)', ] @@ -714,7 +746,9 @@ def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: adapter.columns = table_columns - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) assert to_sql_calls(adapter) == [ 'ALTER TABLE "test_table" DROP COLUMN "test_column" CASCADE', 'ALTER TABLE "test_table" ADD COLUMN "test_column" DECIMAL(25, 10)', @@ -762,7 +796,9 @@ def test_merge(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): ] -def test_merge_when_matched_error(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_merge_when_matched_error( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(RedshiftEngineAdapter) mocker.patch( "sqlmesh.core.engine_adapter.redshift.RedshiftEngineAdapter.enable_merge", @@ -785,7 +821,9 @@ def test_merge_when_matched_error(make_mocked_engine_adapter: t.Callable, mocker expressions=[ exp.When( matched=True, - condition=exp.column("ID", "__MERGE_SOURCE__").eq(exp.Literal.number(1)), + condition=exp.column("ID", "__MERGE_SOURCE__").eq( + exp.Literal.number(1) + ), then=exp.Update( expressions=[ exp.column("val", "__MERGE_TARGET__").eq( @@ -810,7 +848,9 @@ def test_merge_when_matched_error(make_mocked_engine_adapter: t.Callable, mocker ) -def test_merge_logical_filter_error(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_merge_logical_filter_error( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(RedshiftEngineAdapter) mocker.patch( "sqlmesh.core.engine_adapter.redshift.RedshiftEngineAdapter.enable_merge", @@ -831,17 +871,22 @@ def test_merge_logical_filter_error(make_mocked_engine_adapter: t.Callable, mock unique_key=[exp.to_identifier("ID", quoted=True)], merge_filter=exp.and_( exp.and_(exp.column("ID", "__MERGE_SOURCE__") > 0), - exp.column("ts", "__MERGE_TARGET__") < exp.column("ts", "__MERGE_SOURCE__"), + exp.column("ts", "__MERGE_TARGET__") + < exp.column("ts", "__MERGE_SOURCE__"), ), ) def test_merge_logical( - make_mocked_engine_adapter: t.Callable, make_temp_table_name: t.Callable, mocker: MockerFixture + make_mocked_engine_adapter: t.Callable, + make_temp_table_name: t.Callable, + mocker: MockerFixture, ): adapter = make_mocked_engine_adapter(RedshiftEngineAdapter) - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "test" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) diff --git a/tests/core/engine_adapter/test_risingwave.py b/tests/core/engine_adapter/test_risingwave.py index ed3cd77a3f..f73e720980 100644 --- a/tests/core/engine_adapter/test_risingwave.py +++ b/tests/core/engine_adapter/test_risingwave.py @@ -3,7 +3,8 @@ from unittest.mock import call import pytest -from sqlglot import parse_one, exp +from sqlglot import exp, parse_one + from sqlmesh.core.engine_adapter.risingwave import RisingwaveEngineAdapter pytestmark = [pytest.mark.engine, pytest.mark.risingwave] diff --git a/tests/core/engine_adapter/test_snowflake.py b/tests/core/engine_adapter/test_snowflake.py index 085c51098b..39e1a2bda6 100644 --- a/tests/core/engine_adapter/test_snowflake.py +++ b/tests/core/engine_adapter/test_snowflake.py @@ -13,11 +13,11 @@ from sqlmesh.core.engine_adapter.shared import DataObjectType from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.model.definition import SqlModel +from sqlmesh.core.model.kind import ViewKind from sqlmesh.core.node import IntervalUnit -from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils import optional_import +from sqlmesh.utils.errors import SQLMeshError from tests.core.engine_adapter import to_sql_calls -from sqlmesh.core.model.kind import ViewKind pytestmark = [pytest.mark.engine, pytest.mark.snowflake] @@ -35,16 +35,23 @@ def test_get_temp_table(mocker: MockerFixture, make_mocked_engine_adapter: t.Cal mocker.patch("sqlmesh.core.engine_adapter.base.random_id", return_value="abcdefgh") value = adapter._get_temp_table( - normalize_model_name("catalog.db.test_table", default_catalog=None, dialect=adapter.dialect) + normalize_model_name( + "catalog.db.test_table", default_catalog=None, dialect=adapter.dialect + ) ) - assert value.sql(dialect=adapter.dialect) == '"CATALOG"."DB"."__temp_TEST_TABLE_abcdefgh"' + assert ( + value.sql(dialect=adapter.dialect) + == '"CATALOG"."DB"."__temp_TEST_TABLE_abcdefgh"' + ) def test_get_data_objects_lowercases_columns( make_mocked_engine_adapter: t.Callable, mocker: MockerFixture ) -> None: - adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter, patch_get_data_objects=False) + adapter = make_mocked_engine_adapter( + SnowflakeEngineAdapter, patch_get_data_objects=False + ) adapter.get_current_catalog = mocker.Mock(return_value="TEST_CATALOG") @@ -83,10 +90,34 @@ def test_get_data_objects_lowercases_columns( '"test_warehouse"', False, ), - ("test_warehouse", '"test_warehouse"', "test_warehouse", '"TEST_WAREHOUSE"', True), - ("TEST_WAREHOUSE", '"TEST_WAREHOUSE"', "test_warehouse", '"TEST_WAREHOUSE"', False), - ("test warehouse", '"test warehouse"', "test warehouse", '"test warehouse"', False), - ("test warehouse", '"test warehouse"', "another warehouse", '"another warehouse"', True), + ( + "test_warehouse", + '"test_warehouse"', + "test_warehouse", + '"TEST_WAREHOUSE"', + True, + ), + ( + "TEST_WAREHOUSE", + '"TEST_WAREHOUSE"', + "test_warehouse", + '"TEST_WAREHOUSE"', + False, + ), + ( + "test warehouse", + '"test warehouse"', + "test warehouse", + '"test warehouse"', + False, + ), + ( + "test warehouse", + '"test warehouse"', + "another warehouse", + '"another warehouse"', + True, + ), ( "test warehouse", '"test warehouse"', @@ -94,9 +125,27 @@ def test_get_data_objects_lowercases_columns( '"another warehouse"', True, ), - ("test warehouse", '"test warehouse"', "another_warehouse", '"ANOTHER_WAREHOUSE"', True), - ("TEST_WAREHOUSE", '"TEST_WAREHOUSE"', "another_warehouse", '"ANOTHER_WAREHOUSE"', True), - ("test_warehouse", '"test_warehouse"', "another_warehouse", '"ANOTHER_WAREHOUSE"', True), + ( + "test warehouse", + '"test warehouse"', + "another_warehouse", + '"ANOTHER_WAREHOUSE"', + True, + ), + ( + "TEST_WAREHOUSE", + '"TEST_WAREHOUSE"', + "another_warehouse", + '"ANOTHER_WAREHOUSE"', + True, + ), + ( + "test_warehouse", + '"test_warehouse"', + "another_warehouse", + '"ANOTHER_WAREHOUSE"', + True, + ), ( "test_warehouse", '"test_warehouse"', @@ -217,7 +266,9 @@ def test_comments(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture) ] -def test_multiple_column_comments(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_multiple_column_comments( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) adapter.create_table( @@ -246,18 +297,26 @@ def test_multiple_column_comments(make_mocked_engine_adapter: t.Callable, mocker ] -def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_sync_grants_config( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) relation = normalize_identifiers( - exp.to_table("test_db.test_schema.test_table", dialect="snowflake"), dialect="snowflake" + exp.to_table("test_db.test_schema.test_table", dialect="snowflake"), + dialect="snowflake", ) - new_grants_config = {"SELECT": ["ROLE role1", "ROLE role2"], "INSERT": ["ROLE role3"]} + new_grants_config = { + "SELECT": ["ROLE role1", "ROLE role2"], + "INSERT": ["ROLE role3"], + } current_grants = [ ("SELECT", "ROLE old_role"), ("UPDATE", "ROLE legacy_role"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -274,9 +333,18 @@ def test_sync_grants_config(make_mocked_engine_adapter: t.Callable, mocker: Mock sql_calls = to_sql_calls(adapter) assert len(sql_calls) == 5 - assert 'GRANT SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "ROLE1"' in sql_calls - assert 'GRANT SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "ROLE2"' in sql_calls - assert 'GRANT INSERT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "ROLE3"' in sql_calls + assert ( + 'GRANT SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "ROLE1"' + in sql_calls + ) + assert ( + 'GRANT SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "ROLE2"' + in sql_calls + ) + assert ( + 'GRANT INSERT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "ROLE3"' + in sql_calls + ) assert ( 'REVOKE SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" FROM ROLE "OLD_ROLE"' in sql_calls @@ -292,7 +360,8 @@ def test_sync_grants_config_with_overlaps( ): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) relation = normalize_identifiers( - exp.to_table("test_db.test_schema.test_table", dialect="snowflake"), dialect="snowflake" + exp.to_table("test_db.test_schema.test_table", dialect="snowflake"), + dialect="snowflake", ) new_grants_config = { "SELECT": ["ROLE shared", "ROLE new_role"], @@ -304,7 +373,9 @@ def test_sync_grants_config_with_overlaps( ("SELECT", "ROLE legacy"), ("INSERT", "ROLE shared"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -322,10 +393,12 @@ def test_sync_grants_config_with_overlaps( assert len(sql_calls) == 3 assert ( - 'GRANT SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "NEW_ROLE"' in sql_calls + 'GRANT SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "NEW_ROLE"' + in sql_calls ) assert ( - 'GRANT INSERT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "WRITER"' in sql_calls + 'GRANT INSERT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" TO ROLE "WRITER"' + in sql_calls ) assert ( 'REVOKE SELECT ON TABLE "TEST_DB"."TEST_SCHEMA"."TEST_TABLE" FROM ROLE "LEGACY"' @@ -350,7 +423,8 @@ def test_sync_grants_config_object_kind( ) -> None: adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) relation = normalize_identifiers( - exp.to_table("test_db.test_schema.test_object", dialect="snowflake"), dialect="snowflake" + exp.to_table("test_db.test_schema.test_object", dialect="snowflake"), + dialect="snowflake", ) mocker.patch.object(adapter, "fetchall", return_value=[]) @@ -363,19 +437,26 @@ def test_sync_grants_config_object_kind( ] -def test_sync_grants_config_quotes(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_sync_grants_config_quotes( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) relation = normalize_identifiers( exp.to_table('"test_db"."test_schema"."test_table"', dialect="snowflake"), dialect="snowflake", ) - new_grants_config = {"SELECT": ["ROLE role1", "ROLE role2"], "INSERT": ["ROLE role3"]} + new_grants_config = { + "SELECT": ["ROLE role1", "ROLE role2"], + "INSERT": ["ROLE role3"], + } current_grants = [ ("SELECT", "ROLE old_role"), ("UPDATE", "ROLE legacy_role"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) adapter.sync_grants_config(relation, new_grants_config) @@ -392,9 +473,18 @@ def test_sync_grants_config_quotes(make_mocked_engine_adapter: t.Callable, mocke sql_calls = to_sql_calls(adapter) assert len(sql_calls) == 5 - assert 'GRANT SELECT ON TABLE "test_db"."test_schema"."test_table" TO ROLE "ROLE1"' in sql_calls - assert 'GRANT SELECT ON TABLE "test_db"."test_schema"."test_table" TO ROLE "ROLE2"' in sql_calls - assert 'GRANT INSERT ON TABLE "test_db"."test_schema"."test_table" TO ROLE "ROLE3"' in sql_calls + assert ( + 'GRANT SELECT ON TABLE "test_db"."test_schema"."test_table" TO ROLE "ROLE1"' + in sql_calls + ) + assert ( + 'GRANT SELECT ON TABLE "test_db"."test_schema"."test_table" TO ROLE "ROLE2"' + in sql_calls + ) + assert ( + 'GRANT INSERT ON TABLE "test_db"."test_schema"."test_table" TO ROLE "ROLE3"' + in sql_calls + ) assert ( 'REVOKE SELECT ON TABLE "test_db"."test_schema"."test_table" FROM ROLE "OLD_ROLE"' in sql_calls @@ -412,13 +502,18 @@ def test_sync_grants_config_no_catalog_or_schema( relation = normalize_identifiers( exp.to_table('"TesT_Table"', dialect="snowflake"), dialect="snowflake" ) - new_grants_config = {"SELECT": ["ROLE role1", "ROLE role2"], "INSERT": ["ROLE role3"]} + new_grants_config = { + "SELECT": ["ROLE role1", "ROLE role2"], + "INSERT": ["ROLE role3"], + } current_grants = [ ("SELECT", "ROLE old_role"), ("UPDATE", "ROLE legacy_role"), ] - fetchall_mock = mocker.patch.object(adapter, "fetchall", return_value=current_grants) + fetchall_mock = mocker.patch.object( + adapter, "fetchall", return_value=current_grants + ) mocker.patch.object(adapter, "get_current_catalog", return_value="caTalog") mocker.patch.object(adapter, "_get_current_schema", return_value="sChema") @@ -457,7 +552,9 @@ def test_df_to_source_queries_use_schema( df = pd.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]}) adapter.replace_query( - "other_db.test_table", df, {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")} + "other_db.test_table", + df, + {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")}, ) assert 'USE SCHEMA "other_db"' in to_sql_calls(adapter) @@ -476,12 +573,16 @@ def test_df_to_source_queries_reset_non_default_index( "sqlmesh.core.engine_adapter.snowflake.SnowflakeEngineAdapter.table_exists", return_value=False, ) - write_pandas = mocker.patch("snowflake.connector.pandas_tools.write_pandas", return_value=None) + write_pandas = mocker.patch( + "snowflake.connector.pandas_tools.write_pandas", return_value=None + ) adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) df = pd.DataFrame({"a": [2, 3], "b": [5, 6]}, index=[1, 2]) adapter.replace_query( - "other_db.test_table", df, {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")} + "other_db.test_table", + df, + {"a": exp.DataType.build("INT"), "b": exp.DataType.build("INT")}, ) uploaded_df = write_pandas.call_args.args[1] @@ -489,7 +590,9 @@ def test_df_to_source_queries_reset_non_default_index( assert uploaded_df.to_dict("list") == {"a": [2, 3], "b": [5, 6]} -def test_create_managed_table(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_create_managed_table( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) mocker.patch( @@ -574,7 +677,9 @@ def test_create_managed_table(make_mocked_engine_adapter: t.Callable, mocker: Mo ] -def test_drop_managed_table(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_drop_managed_table( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) adapter.drop_managed_table(table_name="foo.bar", exists=False) @@ -621,9 +726,7 @@ def test_set_current_catalog(make_mocked_engine_adapter: t.Callable): model_a: SqlModel = t.cast( SqlModel, - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name external.test.table, kind full, @@ -631,16 +734,12 @@ def test_set_current_catalog(make_mocked_engine_adapter: t.Callable): ); SELECT 1; - """ - ) - ), + """)), ) model_b: SqlModel = t.cast( SqlModel, - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name "exTERnal".test.table, kind full, @@ -648,9 +747,7 @@ def test_set_current_catalog(make_mocked_engine_adapter: t.Callable): ); SELECT 1; - """ - ) - ), + """)), ) assert model_a.catalog == "external" @@ -695,15 +792,17 @@ def test_replace_query_snowpark_dataframe( if not optional_import("snowflake.snowpark"): pytest.skip("Snowpark not available in this environment") - from snowflake.snowpark.session import Session from snowflake.snowpark.dataframe import DataFrame as SnowparkDataFrame + from snowflake.snowpark.session import Session session = Session.builder.config("local_testing", True).create() # df.createOrReplaceTempView() throws "[Local Testing] Mocking SnowflakePlan Rename is not supported" when used against the Snowflake local_testing session # since we cant trace any queries from the Snowpark library anyway, we just suppress this and verify the cleanup queries issued by our EngineAdapter session._conn._suppress_not_implemented_error = True - df: SnowparkDataFrame = session.create_dataframe([(1, "name")], schema=["ID", "NAME"]) + df: SnowparkDataFrame = session.create_dataframe( + [(1, "name")], schema=["ID", "NAME"] + ) assert isinstance(df, SnowparkDataFrame) mocker.patch("sqlmesh.core.engine_adapter.base.random_id", return_value="e6wjkjj6") @@ -728,7 +827,9 @@ def test_replace_query_snowpark_dataframe( ] -def test_creatable_type_materialized_view_properties(make_mocked_engine_adapter: t.Callable): +def test_creatable_type_materialized_view_properties( + make_mocked_engine_adapter: t.Callable, +): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) adapter.create_view( @@ -768,7 +869,9 @@ def test_creatable_type_secure_view(make_mocked_engine_adapter: t.Callable): ] -def test_creatable_type_secure_materialized_view(make_mocked_engine_adapter: t.Callable): +def test_creatable_type_secure_materialized_view( + make_mocked_engine_adapter: t.Callable, +): adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter) adapter.create_view( @@ -860,9 +963,7 @@ def test_creatable_type_transient_type_from_model_definition( model: SqlModel = t.cast( SqlModel, - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name external.test.table, kind full, @@ -871,9 +972,7 @@ def test_creatable_type_transient_type_from_model_definition( ) ); SELECT a::INT; - """ - ) - ), + """)), ) adapter.create_table( model.name, @@ -894,9 +993,7 @@ def test_creatable_type_transient_type_from_model_definition_with_other_property model: SqlModel = t.cast( SqlModel, - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name external.test.table, kind full, @@ -906,9 +1003,7 @@ def test_creatable_type_transient_type_from_model_definition_with_other_property ) ); SELECT a::INT; - """ - ) - ), + """)), ) adapter.create_table( model.name, @@ -936,8 +1031,12 @@ def test_create_view(make_mocked_engine_adapter: t.Callable): def test_clone_table(mocker: MockerFixture, make_mocked_engine_adapter: t.Callable): - mocker.patch("sqlmesh.core.engine_adapter.snowflake.SnowflakeEngineAdapter.set_current_catalog") - adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter, default_catalog="test_catalog") + mocker.patch( + "sqlmesh.core.engine_adapter.snowflake.SnowflakeEngineAdapter.set_current_catalog" + ) + adapter = make_mocked_engine_adapter( + SnowflakeEngineAdapter, default_catalog="test_catalog" + ) adapter.clone_table("target_table", "source_table") adapter.cursor.execute.assert_called_once_with( 'CREATE TABLE IF NOT EXISTS "target_table" CLONE "source_table"' @@ -947,9 +1046,13 @@ def test_clone_table(mocker: MockerFixture, make_mocked_engine_adapter: t.Callab rendered_physical_properties = { "creatable_type": exp.column("transient"), } - adapter = make_mocked_engine_adapter(SnowflakeEngineAdapter, default_catalog="test_catalog") + adapter = make_mocked_engine_adapter( + SnowflakeEngineAdapter, default_catalog="test_catalog" + ) adapter.clone_table( - "target_table", "source_table", rendered_physical_properties=rendered_physical_properties + "target_table", + "source_table", + rendered_physical_properties=rendered_physical_properties, ) adapter.cursor.execute.assert_called_once_with( 'CREATE TRANSIENT TABLE IF NOT EXISTS "target_table" CLONE "source_table"' @@ -959,18 +1062,21 @@ def test_clone_table(mocker: MockerFixture, make_mocked_engine_adapter: t.Callab adapter = make_mocked_engine_adapter(EngineAdapter, default_catalog="test_catalog") adapter.SUPPORTS_CLONING = True adapter.clone_table( - "target_table", "source_table", rendered_physical_properties=rendered_physical_properties + "target_table", + "source_table", + rendered_physical_properties=rendered_physical_properties, ) adapter.cursor.execute.assert_called_once_with( 'CREATE TABLE IF NOT EXISTS "target_table" CLONE "source_table"' ) -def test_table_format_iceberg(snowflake_mocked_engine_adapter: SnowflakeEngineAdapter) -> None: +def test_table_format_iceberg( + snowflake_mocked_engine_adapter: SnowflakeEngineAdapter, +) -> None: adapter = snowflake_mocked_engine_adapter - model = load_sql_based_model( - expressions=d.parse(""" + model = load_sql_based_model(expressions=d.parse(""" MODEL ( name test.table, kind full, @@ -981,8 +1087,7 @@ def test_table_format_iceberg(snowflake_mocked_engine_adapter: SnowflakeEngineAd ) ); SELECT a::INT; - """) - ) + """)) assert isinstance(model, SqlModel) assert model.table_format == "iceberg" @@ -1012,8 +1117,7 @@ def test_create_view_with_schema_and_grants( ): adapter = snowflake_mocked_engine_adapter - model_v = load_sql_based_model( - d.parse(f""" + model_v = load_sql_based_model(d.parse(f""" MODEL ( name test.v, kind VIEW, @@ -1022,11 +1126,9 @@ def test_create_view_with_schema_and_grants( ); select 1 as "ID", 'foo' as "NAME"; - """) - ) + """)) - model_mv = load_sql_based_model( - d.parse(f""" + model_mv = load_sql_based_model(d.parse(f""" MODEL ( name test.mv, kind VIEW ( @@ -1037,8 +1139,7 @@ def test_create_view_with_schema_and_grants( ); select 1 as "ID", 'foo' as "NAME"; - """) - ) + """)) assert isinstance(model_v.kind, ViewKind) assert isinstance(model_mv.kind, ViewKind) @@ -1071,7 +1172,9 @@ def test_create_view_with_schema_and_grants( ] -def test_create_catalog(snowflake_mocked_engine_adapter: SnowflakeEngineAdapter) -> None: +def test_create_catalog( + snowflake_mocked_engine_adapter: SnowflakeEngineAdapter, +) -> None: adapter = snowflake_mocked_engine_adapter adapter.create_catalog(exp.to_identifier("foo")) diff --git a/tests/core/engine_adapter/test_spark.py b/tests/core/engine_adapter/test_spark.py index d7c3127f05..152ed7ebda 100644 --- a/tests/core/engine_adapter/test_spark.py +++ b/tests/core/engine_adapter/test_spark.py @@ -9,13 +9,13 @@ from sqlglot import expressions as exp from sqlglot import parse_one +import sqlmesh.core.dialect as d from sqlmesh.core.engine_adapter import SparkEngineAdapter from sqlmesh.core.engine_adapter.shared import DataObject -from sqlmesh.utils.errors import SQLMeshError -from tests.core.engine_adapter import to_sql_calls -import sqlmesh.core.dialect as d from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.model.definition import SqlModel +from sqlmesh.utils.errors import SQLMeshError +from tests.core.engine_adapter import to_sql_calls pytestmark = [pytest.mark.engine, pytest.mark.spark] @@ -136,7 +136,9 @@ def test_create_view_properties(make_mocked_engine_adapter: t.Callable): adapter = make_mocked_engine_adapter(SparkEngineAdapter) adapter.create_view( - "test_view", parse_one("SELECT a FROM tbl"), view_properties={"a": exp.convert(1)} + "test_view", + parse_one("SELECT a FROM tbl"), + view_properties={"a": exp.convert(1)}, ) # type: ignore adapter.cursor.execute.assert_called_once_with( "CREATE OR REPLACE VIEW test_view TBLPROPERTIES ('a'=1) AS SELECT a FROM tbl" @@ -154,7 +156,9 @@ def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: "id": exp.DataType.build("INT"), "a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING"), - "complex": exp.DataType.build("STRUCT"), + "complex": exp.DataType.build( + "STRUCT" + ), "ds": exp.DataType.build("STRING"), } return { @@ -166,7 +170,9 @@ def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: adapter.columns = table_columns - adapter.alter_table(adapter.get_alter_operations(current_table_name, target_table_name)) + adapter.alter_table( + adapter.get_alter_operations(current_table_name, target_table_name) + ) adapter.cursor.execute.assert_has_calls( [ @@ -176,14 +182,18 @@ def table_columns(table_name: str) -> t.Dict[str, exp.DataType]: call("""ALTER TABLE `test_table` DROP COLUMN `a`"""), call("""ALTER TABLE `test_table` ADD COLUMN `a` STRING"""), call("""ALTER TABLE `test_table` DROP COLUMN `complex`"""), - call("""ALTER TABLE `test_table` ADD COLUMN `complex` STRUCT<`complex_a`: INT>"""), + call( + """ALTER TABLE `test_table` ADD COLUMN `complex` STRUCT<`complex_a`: INT>""" + ), call("""ALTER TABLE `test_table` DROP COLUMN `ds`"""), call("""ALTER TABLE `test_table` ADD COLUMN `ds` INT"""), ] ) -def test_replace_query_not_exists(mocker: MockerFixture, make_mocked_engine_adapter: t.Callable): +def test_replace_query_not_exists( + mocker: MockerFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.spark.SparkEngineAdapter.table_exists", return_value=False, @@ -198,7 +208,9 @@ def test_replace_query_not_exists(mocker: MockerFixture, make_mocked_engine_adap ] -def test_replace_query_exists(mocker: MockerFixture, make_mocked_engine_adapter: t.Callable): +def test_replace_query_exists( + mocker: MockerFixture, make_mocked_engine_adapter: t.Callable +): mocker.patch( "sqlmesh.core.engine_adapter.spark.SparkEngineAdapter.table_exists", return_value=True, @@ -217,7 +229,9 @@ def test_replace_query_exists(mocker: MockerFixture, make_mocked_engine_adapter: def test_replace_query_self_ref_not_exists( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.spark.SparkEngineAdapter.get_current_catalog", @@ -245,7 +259,10 @@ def test_replace_query_self_ref_not_exists( def check_table_exists(table_name: exp.Table) -> bool: for sql in to_sql_calls(adapter): - if f"CREATE TABLE IF NOT EXISTS {table_name.sql(dialect=adapter.dialect)}" in sql: + if ( + f"CREATE TABLE IF NOT EXISTS {table_name.sql(dialect=adapter.dialect)}" + in sql + ): return True return False @@ -260,7 +277,9 @@ def check_table_exists(table_name: exp.Table) -> bool: return_value=[DataObject(schema="db", name="table", type="table")], ) - adapter.replace_query(table_name, parse_one(f"SELECT col + 1 AS col FROM {table_name}")) + adapter.replace_query( + table_name, parse_one(f"SELECT col + 1 AS col FROM {table_name}") + ) assert to_sql_calls(adapter) == [ "CREATE TABLE IF NOT EXISTS `db`.`table` (`col` INT)", @@ -272,7 +291,9 @@ def check_table_exists(table_name: exp.Table) -> bool: def test_replace_query_self_ref_exists( - make_mocked_engine_adapter: t.Callable, mocker: MockerFixture, make_temp_table_name: t.Callable + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, + make_temp_table_name: t.Callable, ): mocker.patch( "sqlmesh.core.engine_adapter.spark.SparkEngineAdapter.table_exists", @@ -307,7 +328,9 @@ def test_replace_query_self_ref_exists( return_value={"col": exp.DataType(this=exp.DataType.Type.INT)}, ) - adapter.replace_query(table_name, parse_one(f"SELECT col + 1 AS col FROM {table_name}")) + adapter.replace_query( + table_name, parse_one(f"SELECT col + 1 AS col FROM {table_name}") + ) assert to_sql_calls(adapter) == [ "CREATE TABLE IF NOT EXISTS `db`.`table` (`col` INT)", @@ -379,7 +402,9 @@ def test_col_to_types_to_spark_schema_primitives(type_name, spark_type): assert SparkEngineAdapter.sqlglot_to_spark_types( {f"col_{type_name}": exp.DataType.build(type_name, dialect="spark")} - ) == spark_types.StructType([spark_types.StructField(f"col_{type_name}", spark_type)]) + ) == spark_types.StructType( + [spark_types.StructField(f"col_{type_name}", spark_type)] + ) test_complex_params = [ @@ -442,10 +467,14 @@ def test_col_to_types_to_spark_schema_primitives(type_name, spark_type): spark_types.ArrayType( spark_types.StructType( [ - spark_types.StructField("cola", spark_types.IntegerType()), + spark_types.StructField( + "cola", spark_types.IntegerType() + ), spark_types.StructField( "colb", - spark_types.ArrayType(spark_types.StringType(), True), + spark_types.ArrayType( + spark_types.StringType(), True + ), ), ] ), @@ -468,7 +497,8 @@ def test_col_to_types_to_spark_schema_primitives(type_name, spark_type): "array_of_maps", "array>", spark_types.ArrayType( - spark_types.MapType(spark_types.IntegerType(), spark_types.StringType()), True + spark_types.MapType(spark_types.IntegerType(), spark_types.StringType()), + True, ), ), ( @@ -512,7 +542,9 @@ def test_col_to_types_to_spark_schema_complex(type_name, spark_type): actual = SparkEngineAdapter.sqlglot_to_spark_types( {f"col_{type_name}": exp.DataType.build(type_name, dialect="spark")} ) - expected = spark_types.StructType([spark_types.StructField(f"col_{type_name}", spark_type)]) + expected = spark_types.StructType( + [spark_types.StructField(f"col_{type_name}", spark_type)] + ) assert actual == expected @@ -523,7 +555,9 @@ def test_col_to_types_to_spark_schema_complex(type_name, spark_type): ) def test_spark_struct_primitives_to_col_to_types(type_name, spark_type): actual = SparkEngineAdapter.spark_to_sqlglot_types( - spark_types.StructType([spark_types.StructField(f"col_{type_name}", spark_type)]) + spark_types.StructType( + [spark_types.StructField(f"col_{type_name}", spark_type)] + ) ) expected_type = ( @@ -542,28 +576,37 @@ def test_spark_struct_primitives_to_col_to_types(type_name, spark_type): ) def test_spark_struct_complex_to_col_to_types(type_name, spark_type): actual = SparkEngineAdapter.spark_to_sqlglot_types( - spark_types.StructType([spark_types.StructField(f"col_{type_name}", spark_type)]) + spark_types.StructType( + [spark_types.StructField(f"col_{type_name}", spark_type)] + ) ) expected = {f"col_{type_name}": exp.DataType.build(type_name, dialect="spark")} assert actual == expected def test_scd_type_2_by_time( - make_mocked_engine_adapter: t.Callable, make_temp_table_name: t.Callable, mocker: MockerFixture + make_mocked_engine_adapter: t.Callable, + make_temp_table_name: t.Callable, + mocker: MockerFixture, ): adapter = make_mocked_engine_adapter(SparkEngineAdapter) adapter._default_catalog = "spark_catalog" adapter.spark.catalog.currentCatalog.return_value = "spark_catalog" adapter.spark.catalog.currentDatabase.return_value = "default" - temp_table_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table") + temp_table_mock = mocker.patch( + "sqlmesh.core.engine_adapter.EngineAdapter._get_temp_table" + ) table_name = "db.target" temp_table_id = "abcdefgh" temp_table_mock.return_value = make_temp_table_name(table_name, temp_table_id) def check_table_exists(table_name: exp.Table) -> bool: for sql in to_sql_calls(adapter): - if f"CREATE TABLE IF NOT EXISTS {table_name.sql(dialect=adapter.dialect)}" in sql: + if ( + f"CREATE TABLE IF NOT EXISTS {table_name.sql(dialect=adapter.dialect)}" + in sql + ): return True return False @@ -580,7 +623,8 @@ def check_table_exists(table_name: exp.Table) -> bool: adapter.scd_type_2_by_time( target_table="db.target", source_table=t.cast( - exp.Select, parse_one("SELECT id, name, price, test_updated_at FROM db.source") + exp.Select, + parse_one("SELECT id, name, price, test_updated_at FROM db.source"), ), unique_key=[exp.func("COALESCE", "id", "''")], valid_from_col=exp.column("test_valid_from", quoted=True), @@ -920,7 +964,9 @@ def test_comments_hive(mocker: MockerFixture, make_mocked_engine_adapter: t.Call ] -def test_comments_iceberg(mocker: MockerFixture, make_mocked_engine_adapter: t.Callable): +def test_comments_iceberg( + mocker: MockerFixture, make_mocked_engine_adapter: t.Callable +): adapter = make_mocked_engine_adapter(SparkEngineAdapter) current_catalog_mock = mocker.patch( @@ -978,7 +1024,9 @@ def test_comments_iceberg(mocker: MockerFixture, make_mocked_engine_adapter: t.C ] -def test_create_table_with_wap(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_create_table_with_wap( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): mocker.patch( "sqlmesh.core.engine_adapter.spark.SparkEngineAdapter.table_exists", return_value=False, @@ -1040,8 +1088,7 @@ def test_table_format(adapter: SparkEngineAdapter, mocker: MockerFixture): return_value=True, ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_table, kind FULL, @@ -1050,8 +1097,7 @@ def test_table_format(adapter: SparkEngineAdapter, mocker: MockerFixture): ); SELECT 1::timestamp AS cola, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) # both table_format and storage_format @@ -1092,8 +1138,12 @@ def test_table_format(adapter: SparkEngineAdapter, mocker: MockerFixture): ] -def test_get_data_object_wap_branch(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): - adapter = make_mocked_engine_adapter(SparkEngineAdapter, patch_get_data_objects=False) +def test_get_data_object_wap_branch( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): + adapter = make_mocked_engine_adapter( + SparkEngineAdapter, patch_get_data_objects=False + ) mocker.patch.object(adapter, "_get_data_objects", return_value=[]) table = exp.to_table( diff --git a/tests/core/engine_adapter/test_starrocks.py b/tests/core/engine_adapter/test_starrocks.py index db0b1cc4ae..72cf798f52 100644 --- a/tests/core/engine_adapter/test_starrocks.py +++ b/tests/core/engine_adapter/test_starrocks.py @@ -20,24 +20,22 @@ import typing as t import pytest +from pytest_mock.plugin import MockerFixture from sqlglot import expressions as exp from sqlglot import parse_one -from pytest_mock.plugin import MockerFixture -from sqlmesh.core.engine_adapter.shared import DataObjectType -from sqlmesh.utils.errors import SQLMeshError -from tests.core.engine_adapter import to_sql_calls +from sqlmesh.core.dialect import parse from sqlmesh.core.engine_adapter.base import EngineAdapter -from sqlmesh.core.engine_adapter.starrocks import StarRocksEngineAdapter from sqlmesh.core.engine_adapter.duckdb import DuckDBEngineAdapter -from sqlmesh.core.dialect import parse -from sqlmesh.core.model import load_sql_based_model, SqlModel -from sqlmesh.core.snapshot.definition import ( - DeployabilityIndex, - Snapshot, - SnapshotChangeCategory, -) -from sqlmesh.core.snapshot.evaluator import _adjust_physical_properties_for_engine +from sqlmesh.core.engine_adapter.shared import DataObjectType +from sqlmesh.core.engine_adapter.starrocks import StarRocksEngineAdapter +from sqlmesh.core.model import SqlModel, load_sql_based_model +from sqlmesh.core.snapshot.definition import (DeployabilityIndex, Snapshot, + SnapshotChangeCategory) +from sqlmesh.core.snapshot.evaluator import \ + _adjust_physical_properties_for_engine +from sqlmesh.utils.errors import SQLMeshError +from tests.core.engine_adapter import to_sql_calls pytestmark = [pytest.mark.starrocks, pytest.mark.engine] @@ -84,7 +82,9 @@ def test_create_schema_without_if_exists( "CREATE SCHEMA `test_schema`", ] - def test_drop_schema(self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter]): + def test_drop_schema( + self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter] + ): """Test DROP DATABASE statement generation.""" adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) adapter.drop_schema("test_schema") @@ -111,7 +111,9 @@ def test_get_data_object_materialized_view_is_distinguished_from_view( """ import pandas as pd - adapter = make_mocked_engine_adapter(StarRocksEngineAdapter, patch_get_data_objects=False) + adapter = make_mocked_engine_adapter( + StarRocksEngineAdapter, patch_get_data_objects=False + ) # information_schema.tables output (MV appears as 'view') # fetchdf is called twice: @@ -137,7 +139,9 @@ def test_get_data_object_materialized_view_is_distinguished_from_view( def fetchdf_side_effect(query: exp.Expression, *_: t.Any, **__: t.Any): query_sql = query.sql(dialect="starrocks").lower() requested = [ - name for name in known_names if f"'{name}'" in query_sql or f"`{name}`" in query_sql + name + for name in known_names + if f"'{name}'" in query_sql or f"`{name}`" in query_sql ] if "information_schema.materialized_views" in query_sql: df = mv_df @@ -158,7 +162,9 @@ def fetchdf_side_effect(query: exp.Expression, *_: t.Any, **__: t.Any): assert v1 is not None assert v1.type == DataObjectType.VIEW - mv2_objects = adapter.get_data_objects(schema_name="test_db", object_names={"mv2"}) + mv2_objects = adapter.get_data_objects( + schema_name="test_db", object_names={"mv2"} + ) assert len(mv2_objects) == 1 assert mv2_objects[0].name.lower() == "mv2" assert mv2_objects[0].type == DataObjectType.MATERIALIZED_VIEW @@ -229,7 +235,9 @@ def test_create_table_like_does_not_call_columns( """ adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) columns_mock = mocker.patch.object( - adapter, "columns", side_effect=AssertionError("columns() should not be called") + adapter, + "columns", + side_effect=AssertionError("columns() should not be called"), ) adapter.create_table_like("target_table", "source_table") @@ -255,15 +263,21 @@ def test_rename_table( # Test 1: Simple table names (no database qualifier) adapter.rename_table("old_table", "new_table") - adapter.cursor.execute.assert_called_with("ALTER TABLE `old_table` RENAME `new_table`") + adapter.cursor.execute.assert_called_with( + "ALTER TABLE `old_table` RENAME `new_table`" + ) # Test 2: Database-qualified names - RENAME only uses table name adapter.cursor.execute.reset_mock() adapter.rename_table("db.old_table", "db.new_table") # StarRocks RENAME clause requires unqualified table name - adapter.cursor.execute.assert_called_with("ALTER TABLE `db`.`old_table` RENAME `new_table`") + adapter.cursor.execute.assert_called_with( + "ALTER TABLE `db`.`old_table` RENAME `new_table`" + ) - def test_delete_from(self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter]): + def test_delete_from( + self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter] + ): """Test DELETE statement generation.""" adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) adapter.delete_from(exp.to_table("test_table"), "id = 1") @@ -282,7 +296,9 @@ def test_create_index( # StarRocks skips index creation - verify no execute call was made adapter.cursor.execute.assert_not_called() - def test_create_view(self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter]): + def test_create_view( + self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter] + ): """Test CREATE VIEW statement generation.""" adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) adapter.create_view("test_view", parse_one("SELECT a FROM tbl")) @@ -430,7 +446,9 @@ def test_create_materialized_view_with_audits_immediate_refresh_raises( ) # Fail-fast: nothing should have been dropped or created. - assert all("CREATE MATERIALIZED VIEW" not in sql for sql in to_sql_calls(adapter)) + assert all( + "CREATE MATERIALIZED VIEW" not in sql for sql in to_sql_calls(adapter) + ) assert all("DROP MATERIALIZED VIEW" not in sql for sql in to_sql_calls(adapter)) def test_create_materialized_view_with_audits_missing_refresh_moment_raises( @@ -581,7 +599,9 @@ def test_delete_with_multiple_between( adapter.delete_from( exp.to_table("test_table"), - parse_one("dt BETWEEN '2024-01-01' AND '2024-12-31' AND id BETWEEN 1 AND 100"), + parse_one( + "dt BETWEEN '2024-01-01' AND '2024-12-31' AND id BETWEEN 1 AND 100" + ), ) sql = to_sql_calls(adapter)[0] @@ -943,7 +963,9 @@ def test_column_reordering_for_key( customer_id_pos = col_defs.find("`customer_id`") assert order_id_pos < event_date_pos, "order_id must appear before event_date" - assert event_date_pos < customer_id_pos, "event_date must appear before customer_id" + assert ( + event_date_pos < customer_id_pos + ), "event_date must appear before customer_id" # ============================================================================= @@ -958,7 +980,11 @@ class TestPartitionPropertyBuilding: # Expression partitioning - single column ("'dt'", "PARTITION BY `dt`", "PARTITION BY (`dt`)"), # Expression partitioning - multi-column - ("(year, month)", "PARTITION BY `year`, `month`", "PARTITION BY (`year`, `month`)"), + ( + "(year, month)", + "PARTITION BY `year`, `month`", + "PARTITION BY (`year`, `month`)", + ), # Expression partitioning - multi-column with func ( "(date_trunc('day', dt), region)", @@ -1091,7 +1117,10 @@ def test_partition_by_alias( ) sql = to_sql_calls(adapter)[0] - assert "PARTITION BY (`year`, `month`)" in sql or "PARTITION BY `year`, `month`" in sql + assert ( + "PARTITION BY (`year`, `month`)" in sql + or "PARTITION BY `year`, `month`" in sql + ) def test_partitioned_by_as_model_parameter( self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter] @@ -1252,9 +1281,15 @@ def test_distributed_by_string_forms( "dist_struct,expected_clause", [ # Structured: HASH with quoted kind - ("(kind='HASH', expressions=id, buckets=32)", "DISTRIBUTED BY HASH (`id`) BUCKETS 32"), + ( + "(kind='HASH', expressions=id, buckets=32)", + "DISTRIBUTED BY HASH (`id`) BUCKETS 32", + ), # Structured: HASH with unquoted kind (Column) - ("(kind=HASH, expressions=id, buckets=10)", "DISTRIBUTED BY HASH (`id`) BUCKETS 10"), + ( + "(kind=HASH, expressions=id, buckets=10)", + "DISTRIBUTED BY HASH (`id`) BUCKETS 10", + ), # Structured: HASH multi-column ( "(kind='HASH', expressions=(a, b), buckets=16)", @@ -1675,7 +1710,9 @@ def test_refresh_scheme_invalid_prefix( self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter] ): adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) - model = self._build_mv_model("refresh_scheme = 'SCHEDULE EVERY (INTERVAL 5 MINUTE)'") + model = self._build_mv_model( + "refresh_scheme = 'SCHEDULE EVERY (INTERVAL 5 MINUTE)'" + ) with pytest.raises(SQLMeshError, match="refresh_scheme"): self._create_simple_mv(adapter, model) @@ -1837,7 +1874,9 @@ def test_build_create_comment_column_exp( adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) table = exp.to_table(table_name) - sql = adapter._build_create_comment_column_exp(table, column_name, comment, "TABLE") + sql = adapter._build_create_comment_column_exp( + table, column_name, comment, "TABLE" + ) assert sql == expected_sql # Should NOT contain column type @@ -1857,7 +1896,9 @@ def test_build_create_comment_column_exp_truncation( table = exp.to_table("test_table") long_comment = "y" * 500 # Longer than MAX_COLUMN_COMMENT_LENGTH (255) - sql = adapter._build_create_comment_column_exp(table, "test_col", long_comment, "TABLE") + sql = adapter._build_create_comment_column_exp( + table, "test_col", long_comment, "TABLE" + ) # The comment should be truncated to 255 characters expected_truncated = "y" * 255 @@ -2004,7 +2045,9 @@ def test_create_table_comprehensive( ), exp.EQ( this=exp.Column(this="expressions"), - expression=exp.Tuple(expressions=[exp.to_column("customer_id")]), + expression=exp.Tuple( + expressions=[exp.to_column("customer_id")] + ), ), exp.EQ( this=exp.Column(this="buckets"), @@ -2042,14 +2085,15 @@ class TestIncrementalRequiresPrimaryKey: """ def _adjust(self, adapter: EngineAdapter, model: SqlModel) -> t.Dict[str, t.Any]: - return _adjust_physical_properties_for_engine(adapter, model, model.physical_properties) + return _adjust_physical_properties_for_engine( + adapter, model, model.physical_properties + ) def test_incremental_model_without_primary_key_raises( self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter] ) -> None: adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) - model = _load_sql_model( - """ + model = _load_sql_model(""" MODEL ( name test_schema.inc_no_pk, kind INCREMENTAL_BY_TIME_RANGE (time_column event_date), @@ -2057,8 +2101,7 @@ def test_incremental_model_without_primary_key_raises( columns (id INT, event_date DATE) ); SELECT id, event_date FROM src WHERE event_date BETWEEN @start_ds AND @end_ds; - """ - ) + """) with pytest.raises(SQLMeshError, match="requires a PRIMARY KEY"): self._adjust(adapter, model) @@ -2066,8 +2109,7 @@ def test_incremental_model_with_primary_key_is_allowed( self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter] ) -> None: adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) - model = _load_sql_model( - """ + model = _load_sql_model(""" MODEL ( name test_schema.inc_pk, kind INCREMENTAL_BY_TIME_RANGE (time_column event_date), @@ -2076,8 +2118,7 @@ def test_incremental_model_with_primary_key_is_allowed( physical_properties (primary_key = (id, event_date)) ); SELECT id, event_date FROM src WHERE event_date BETWEEN @start_ds AND @end_ds; - """ - ) + """) assert "primary_key" in self._adjust(adapter, model) def test_incremental_by_unique_key_is_promoted_to_primary_key( @@ -2086,8 +2127,7 @@ def test_incremental_by_unique_key_is_promoted_to_primary_key( # INCREMENTAL_BY_UNIQUE_KEY auto-promotes the unique_key to a PRIMARY KEY (a multi-column # key becomes a tuple) rather than requiring one to be declared explicitly. adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) - model = _load_sql_model( - """ + model = _load_sql_model(""" MODEL ( name test_schema.inc_by_uk, kind INCREMENTAL_BY_UNIQUE_KEY (unique_key (id, event_date)), @@ -2095,8 +2135,7 @@ def test_incremental_by_unique_key_is_promoted_to_primary_key( columns (id INT, event_date DATE) ); SELECT id, event_date FROM src; - """ - ) + """) primary_key = self._adjust(adapter, model)["primary_key"] assert isinstance(primary_key, exp.Tuple) assert [c.name for c in primary_key.expressions] == ["id", "event_date"] @@ -2107,8 +2146,7 @@ def test_append_only_incremental_does_not_require_primary_key( # Append-only INCREMENTAL_UNMANAGED (insert_overwrite=False) only does INSERT, so it does # not need a PRIMARY KEY table. adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) - model = _load_sql_model( - """ + model = _load_sql_model(""" MODEL ( name test_schema.inc_append, kind INCREMENTAL_UNMANAGED, @@ -2116,8 +2154,7 @@ def test_append_only_incremental_does_not_require_primary_key( columns (id INT, event_date DATE) ); SELECT id, event_date FROM src; - """ - ) + """) assert "primary_key" not in self._adjust(adapter, model) def test_non_starrocks_incremental_is_unaffected( @@ -2125,8 +2162,7 @@ def test_non_starrocks_incremental_is_unaffected( ) -> None: # Engines without the PRIMARY KEY requirement inherit the base no-op and never raise. adapter = make_mocked_engine_adapter(DuckDBEngineAdapter) - model = _load_sql_model( - """ + model = _load_sql_model(""" MODEL ( name test_schema.inc_duckdb, kind INCREMENTAL_BY_TIME_RANGE (time_column event_date), @@ -2134,8 +2170,7 @@ def test_non_starrocks_incremental_is_unaffected( columns (id INT, event_date DATE) ); SELECT id, event_date FROM src WHERE event_date BETWEEN @start_ds AND @end_ds; - """ - ) + """) assert "primary_key" not in self._adjust(adapter, model) @@ -2185,8 +2220,7 @@ def test_single_managed_model_ref_is_resolved_to_physical_name( make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter], ) -> None: """excluded_trigger_tables referencing a managed model is resolved to physical db.table.""" - base_model = _load_sql_model( - """ + base_model = _load_sql_model(""" MODEL ( name starrocks.test_1_model, kind FULL, @@ -2194,14 +2228,12 @@ def test_single_managed_model_ref_is_resolved_to_physical_name( columns (a INT) ); SELECT 1 AS a; - """ - ) + """) base_snapshot = self._make_snapshot(base_model) physical_name = exp.to_table(base_snapshot.table_name()) expected_physical = f"{physical_name.db}.{physical_name.name}" - mv_model = _load_sql_model( - """ + mv_model = _load_sql_model(""" MODEL ( name starrocks.test_mv, kind VIEW (materialized true), @@ -2213,8 +2245,7 @@ def test_single_managed_model_ref_is_resolved_to_physical_name( ) ); SELECT a FROM starrocks.test_1_model; - """ - ) + """) adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) snapshots = {base_snapshot.name: base_snapshot} @@ -2229,8 +2260,7 @@ def test_single_managed_model_ref_in_excluded_refresh_tables( make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter], ) -> None: """excluded_refresh_tables referencing a managed model is also resolved.""" - base_model = _load_sql_model( - """ + base_model = _load_sql_model(""" MODEL ( name starrocks.source_model, kind FULL, @@ -2238,14 +2268,12 @@ def test_single_managed_model_ref_in_excluded_refresh_tables( columns (a INT) ); SELECT 1 AS a; - """ - ) + """) base_snapshot = self._make_snapshot(base_model) physical_name = exp.to_table(base_snapshot.table_name()) expected_physical = f"{physical_name.db}.{physical_name.name}" - mv_model = _load_sql_model( - """ + mv_model = _load_sql_model(""" MODEL ( name starrocks.test_mv2, kind VIEW (materialized true), @@ -2257,8 +2285,7 @@ def test_single_managed_model_ref_in_excluded_refresh_tables( ) ); SELECT a FROM starrocks.source_model; - """ - ) + """) adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) snapshots = {base_snapshot.name: base_snapshot} @@ -2271,8 +2298,7 @@ def test_unmanaged_source_is_left_as_is( make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter], ) -> None: """A raw source that is not a managed snapshot passes through unchanged.""" - mv_model = _load_sql_model( - """ + mv_model = _load_sql_model(""" MODEL ( name starrocks.test_mv3, kind VIEW (materialized true), @@ -2284,8 +2310,7 @@ def test_unmanaged_source_is_left_as_is( ) ); SELECT 1 AS a; - """ - ) + """) adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) ddl = self._build_mv_with_excluded_tables(adapter, mv_model, snapshots={}) @@ -2297,8 +2322,7 @@ def test_unmanaged_external_catalog_ref_keeps_catalog( make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter], ) -> None: """A three-part external-catalog reference is preserved in full (catalog not stripped).""" - mv_model = _load_sql_model( - """ + mv_model = _load_sql_model(""" MODEL ( name starrocks.test_mv_ext, kind VIEW (materialized true), @@ -2310,8 +2334,7 @@ def test_unmanaged_external_catalog_ref_keeps_catalog( ) ); SELECT 1 AS a; - """ - ) + """) adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) ddl = self._build_mv_with_excluded_tables(adapter, mv_model, snapshots={}) @@ -2323,8 +2346,7 @@ def test_mixed_list_managed_and_unmanaged( make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter], ) -> None: """A comma-separated list: managed model resolved, unmanaged source left as-is.""" - base_model = _load_sql_model( - """ + base_model = _load_sql_model(""" MODEL ( name starrocks.managed_model, kind FULL, @@ -2332,14 +2354,12 @@ def test_mixed_list_managed_and_unmanaged( columns (a INT) ); SELECT 1 AS a; - """ - ) + """) base_snapshot = self._make_snapshot(base_model) physical_name = exp.to_table(base_snapshot.table_name()) expected_physical = f"{physical_name.db}.{physical_name.name}" - mv_model = _load_sql_model( - """ + mv_model = _load_sql_model(""" MODEL ( name starrocks.test_mv4, kind VIEW (materialized true), @@ -2351,8 +2371,7 @@ def test_mixed_list_managed_and_unmanaged( ) ); SELECT 1 AS a; - """ - ) + """) adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) snapshots = {base_snapshot.name: base_snapshot} @@ -2367,20 +2386,19 @@ def test_mixed_list_managed_and_unmanaged( match = re.search(r"'excluded_trigger_tables'='([^']*)'", ddl) assert match is not None, "excluded_trigger_tables property not found in DDL" prop_value = match.group(1) - assert "starrocks.managed_model" not in prop_value, ( - "Logical model name leaked into excluded_trigger_tables property value" - ) - assert expected_physical in prop_value, ( - "Physical table name not found in excluded_trigger_tables property value" - ) + assert ( + "starrocks.managed_model" not in prop_value + ), "Logical model name leaked into excluded_trigger_tables property value" + assert ( + expected_physical in prop_value + ), "Physical table name not found in excluded_trigger_tables property value" def test_no_snapshots_passes_value_through( self, make_mocked_engine_adapter: t.Callable[..., StarRocksEngineAdapter], ) -> None: """When no snapshots are provided, the value passes through unchanged.""" - mv_model = _load_sql_model( - """ + mv_model = _load_sql_model(""" MODEL ( name starrocks.test_mv5, kind VIEW (materialized true), @@ -2392,8 +2410,7 @@ def test_no_snapshots_passes_value_through( ) ); SELECT 1 AS a; - """ - ) + """) adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) ddl = self._build_mv_with_excluded_tables(adapter, mv_model, snapshots={}) @@ -2410,8 +2427,7 @@ def test_dev_plan_non_deployable_snapshot_resolves_to_dev_table( (``__dev``-suffixed) table selected by ``to_table_mapping``. The property value must reference this dev physical name, NOT the production physical name. """ - base_model = _load_sql_model( - """ + base_model = _load_sql_model(""" MODEL ( name starrocks.upstream, kind FULL, @@ -2419,8 +2435,7 @@ def test_dev_plan_non_deployable_snapshot_resolves_to_dev_table( columns (a INT) ); SELECT 1 AS a; - """ - ) + """) base_snapshot = self._make_snapshot(base_model) prod_physical_name = exp.to_table(base_snapshot.table_name(is_deployable=True)) dev_physical_name = exp.to_table(base_snapshot.table_name(is_deployable=False)) @@ -2428,8 +2443,7 @@ def test_dev_plan_non_deployable_snapshot_resolves_to_dev_table( prod_expected = f"{prod_physical_name.db}.{prod_physical_name.name}" dev_expected = f"{dev_physical_name.db}.{dev_physical_name.name}" - mv_model = _load_sql_model( - """ + mv_model = _load_sql_model(""" MODEL ( name starrocks.test_mv_dev, kind VIEW (materialized true), @@ -2441,8 +2455,7 @@ def test_dev_plan_non_deployable_snapshot_resolves_to_dev_table( ) ); SELECT a FROM starrocks.upstream; - """ - ) + """) adapter = make_mocked_engine_adapter(StarRocksEngineAdapter) snapshots = {base_snapshot.name: base_snapshot} diff --git a/tests/core/engine_adapter/test_trino.py b/tests/core/engine_adapter/test_trino.py index 1bfe82b858..0fc04a791e 100644 --- a/tests/core/engine_adapter/test_trino.py +++ b/tests/core/engine_adapter/test_trino.py @@ -7,10 +7,10 @@ import sqlmesh.core.dialect as d from sqlmesh.core.config.connection import TrinoConnectionConfig +from sqlmesh.core.dialect import schema_ from sqlmesh.core.engine_adapter import TrinoEngineAdapter from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.model.definition import SqlModel -from sqlmesh.core.dialect import schema_ from sqlmesh.utils.date import to_ds from sqlmesh.utils.errors import SQLMeshError from tests.core.engine_adapter import to_sql_calls @@ -53,7 +53,9 @@ def test_set_current_catalog(trino_mocked_engine_adapter: TrinoEngineAdapter): @pytest.mark.parametrize("storage_type", ["iceberg", "delta_lake"]) def test_get_catalog_type( - trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: MockerFixture, storage_type: str + trino_mocked_engine_adapter: TrinoEngineAdapter, + mocker: MockerFixture, + storage_type: str, ): adapter = trino_mocked_engine_adapter mocker.patch( @@ -87,7 +89,8 @@ def mock_fetchone(sql): return ("hive",) mocker.patch( - "sqlmesh.core.engine_adapter.trino.TrinoEngineAdapter.fetchone", side_effect=mock_fetchone + "sqlmesh.core.engine_adapter.trino.TrinoEngineAdapter.fetchone", + side_effect=mock_fetchone, ) fetchone_mock = t.cast(MagicMock, adapter.fetchone) # to make mypy happy @@ -106,7 +109,9 @@ def mock_fetchone(sql): @pytest.mark.parametrize("storage_type", ["hive", "delta_lake"]) def test_partitioned_by_hive_delta( - trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: MockerFixture, storage_type: str + trino_mocked_engine_adapter: TrinoEngineAdapter, + mocker: MockerFixture, + storage_type: str, ): adapter = trino_mocked_engine_adapter @@ -121,7 +126,9 @@ def test_partitioned_by_hive_delta( "colb": exp.DataType.build("TEXT"), } - adapter.create_table("test_table", columns_to_types, partitioned_by=[exp.to_column("colb")]) + adapter.create_table( + "test_table", columns_to_types, partitioned_by=[exp.to_column("colb")] + ) adapter.ctas("test_table", parse_one("select 1"), partitioned_by=[exp.to_column("colb")]) # type: ignore @@ -147,7 +154,9 @@ def test_partitioned_by_iceberg( "colb": exp.DataType.build("TEXT"), } - adapter.create_table("test_table", columns_to_types, partitioned_by=[exp.to_column("colb")]) + adapter.create_table( + "test_table", columns_to_types, partitioned_by=[exp.to_column("colb")] + ) adapter.ctas("test_table", parse_one("select 1"), partitioned_by=[exp.to_column("colb")]) # type: ignore @@ -167,8 +176,7 @@ def test_partitioned_by_iceberg_transforms( return_value="datalake_iceberg", ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_table, partitioned_by (day(cola), truncate(colb, 8), colc), @@ -178,8 +186,7 @@ def test_partitioned_by_iceberg_transforms( ); SELECT 1::timestamp AS cola, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) adapter.create_table( @@ -216,7 +223,9 @@ def test_partitioned_by_with_multiple_catalogs_same_server( ) adapter.create_table( - "datalake.test_schema.test_table", columns_to_types, partitioned_by=[exp.to_column("colb")] + "datalake.test_schema.test_table", + columns_to_types, + partitioned_by=[exp.to_column("colb")], ) adapter.ctas( @@ -430,9 +439,9 @@ def test_timestamp_mapping(): assert config._connection_factory_with_kwargs.keywords["source"] == "my_source" adapter = config.create_engine_adapter() assert adapter.timestamp_mapping is not None - assert adapter.timestamp_mapping[exp.DataType.build("TIMESTAMP")] == exp.DataType.build( - "TIMESTAMP(6)" - ) + assert adapter.timestamp_mapping[ + exp.DataType.build("TIMESTAMP") + ] == exp.DataType.build("TIMESTAMP(6)") def test_delta_timestamps_with_custom_mapping(make_mocked_engine_adapter: t.Callable): @@ -469,7 +478,9 @@ def test_delta_timestamps_with_custom_mapping(make_mocked_engine_adapter: t.Call mapped_columns_to_types, mapped_column_names = adapter._apply_timestamp_mapping( columns_to_types ) - delta_columns_to_types = adapter._to_delta_ts(mapped_columns_to_types, mapped_column_names) + delta_columns_to_types = adapter._to_delta_ts( + mapped_columns_to_types, mapped_column_names + ) # All types were mapped, so _to_delta_ts skips them - they keep their mapped types assert delta_columns_to_types == { @@ -509,7 +520,9 @@ def test_delta_timestamps_with_partial_mapping(make_mocked_engine_adapter: t.Cal mapped_columns_to_types, mapped_column_names = adapter._apply_timestamp_mapping( columns_to_types ) - delta_columns_to_types = adapter._to_delta_ts(mapped_columns_to_types, mapped_column_names) + delta_columns_to_types = adapter._to_delta_ts( + mapped_columns_to_types, mapped_column_names + ) # TIMESTAMP is in mapping → TIMESTAMP(3), skipped by _to_delta_ts # TIMESTAMP(1) is NOT in mapping, uses default TIMESTAMP → ts6 @@ -521,15 +534,16 @@ def test_delta_timestamps_with_partial_mapping(make_mocked_engine_adapter: t.Cal } -def test_table_format(trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: MockerFixture): +def test_table_format( + trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: MockerFixture +): adapter = trino_mocked_engine_adapter mocker.patch( "sqlmesh.core.engine_adapter.trino.TrinoEngineAdapter.get_current_catalog", return_value="iceberg", ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name iceberg.test_table, kind FULL, @@ -538,8 +552,7 @@ def test_table_format(trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: M ); SELECT 1::timestamp AS cola, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) adapter.create_table( @@ -566,15 +579,16 @@ def test_table_format(trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: M ] -def test_table_location(trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: MockerFixture): +def test_table_location( + trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: MockerFixture +): adapter = trino_mocked_engine_adapter mocker.patch( "sqlmesh.core.engine_adapter.trino.TrinoEngineAdapter.get_current_catalog", return_value="iceberg", ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name iceberg.test_table, kind FULL, @@ -584,8 +598,7 @@ def test_table_location(trino_mocked_engine_adapter: TrinoEngineAdapter, mocker: ); SELECT 1::timestamp AS cola, 2::varchar as colb, 'foo' as colc; - """ - ) + """) model: SqlModel = t.cast(SqlModel, load_sql_based_model(expressions)) adapter.create_table( @@ -634,8 +647,14 @@ def test_schema_location_mapping(): assert adapter._schema_location("foo") is None assert adapter._schema_location("utils_dev") is None assert adapter._schema_location("utils") == "s3://utils-bucket/utils" - assert adapter._schema_location("staging_customers") == "s3://bucket/staging_customers_dev" - assert adapter._schema_location("staging_accounts") == "s3://bucket/staging_accounts_dev" + assert ( + adapter._schema_location("staging_customers") + == "s3://bucket/staging_customers_dev" + ) + assert ( + adapter._schema_location("staging_accounts") + == "s3://bucket/staging_accounts_dev" + ) assert ( adapter._schema_location("sqlmesh__staging_customers") == "s3://sqlmesh-internal/dev/sqlmesh__staging_customers" @@ -644,17 +663,23 @@ def test_schema_location_mapping(): adapter._schema_location("sqlmesh__staging_utils") == "s3://sqlmesh-internal/dev/sqlmesh__staging_utils" ) - assert adapter._schema_location("landing.transactions") == "s3://raw-data/landing/transactions" + assert ( + adapter._schema_location("landing.transactions") + == "s3://raw-data/landing/transactions" + ) assert ( adapter._schema_location(schema_("transactions", "landing")) == "s3://raw-data/landing/transactions" ) assert ( - adapter._schema_location('"landing"."transactions"') == "s3://raw-data/landing/transactions" + adapter._schema_location('"landing"."transactions"') + == "s3://raw-data/landing/transactions" ) -def test_create_schema_sets_location(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture): +def test_create_schema_sets_location( + make_mocked_engine_adapter: t.Callable, mocker: MockerFixture +): mocker.patch( "sqlmesh.core.engine_adapter.trino.TrinoEngineAdapter.get_catalog_type", return_value="iceberg", @@ -694,22 +719,19 @@ def test_create_schema_sets_location(make_mocked_engine_adapter: t.Callable, moc adapter.create_schema('"catalog"."staging_customers"') adapter.create_schema(schema_("transactions", "landing")) - assert ( - to_sql_calls(adapter) - == [ - 'CREATE SCHEMA IF NOT EXISTS "foo"', # no match - 'CREATE SCHEMA IF NOT EXISTS "db"."utils_dev"', # no match - 'CREATE SCHEMA IF NOT EXISTS "db"."utils"', # no match on '^utils$' because of catalog - "CREATE SCHEMA IF NOT EXISTS \"utils\" WITH (LOCATION='s3://utils-bucket/utils')", # match '^utils$' - "CREATE SCHEMA IF NOT EXISTS \"sqlmesh\" WITH (LOCATION='s3://sqlmesh-internal/dev/sqlmesh')", # match '^sqlmesh.*$' - "CREATE SCHEMA IF NOT EXISTS \"sqlmesh__staging\" WITH (LOCATION='s3://sqlmesh-internal/dev/sqlmesh__staging')", # match '^sqlmesh.*$' - 'CREATE SCHEMA IF NOT EXISTS "sqlmesh"."snapshots" WITH (LOCATION=\'s3://sqlmesh-internal/dev/snapshots\')', # match '^sqlmesh.*$' on the catalog - "CREATE SCHEMA IF NOT EXISTS \"staging_foo\" WITH (LOCATION='s3://bucket/staging_foo_dev')", # match '^staging.*$' - 'CREATE SCHEMA IF NOT EXISTS "iceberg"."staging_bar" WITH (LOCATION=\'s3://iceberg-catalog/foo_staging_bar\')', # match '^iceberg\.staging.*$' - 'CREATE SCHEMA IF NOT EXISTS "catalog"."staging_customers"', # no match - 'CREATE SCHEMA IF NOT EXISTS "landing"."transactions" WITH (LOCATION=\'s3://raw-data/landing/transactions\')', # match '^landing\..*$' - ] - ) + assert to_sql_calls(adapter) == [ + 'CREATE SCHEMA IF NOT EXISTS "foo"', # no match + 'CREATE SCHEMA IF NOT EXISTS "db"."utils_dev"', # no match + 'CREATE SCHEMA IF NOT EXISTS "db"."utils"', # no match on '^utils$' because of catalog + "CREATE SCHEMA IF NOT EXISTS \"utils\" WITH (LOCATION='s3://utils-bucket/utils')", # match '^utils$' + "CREATE SCHEMA IF NOT EXISTS \"sqlmesh\" WITH (LOCATION='s3://sqlmesh-internal/dev/sqlmesh')", # match '^sqlmesh.*$' + "CREATE SCHEMA IF NOT EXISTS \"sqlmesh__staging\" WITH (LOCATION='s3://sqlmesh-internal/dev/sqlmesh__staging')", # match '^sqlmesh.*$' + 'CREATE SCHEMA IF NOT EXISTS "sqlmesh"."snapshots" WITH (LOCATION=\'s3://sqlmesh-internal/dev/snapshots\')', # match '^sqlmesh.*$' on the catalog + "CREATE SCHEMA IF NOT EXISTS \"staging_foo\" WITH (LOCATION='s3://bucket/staging_foo_dev')", # match '^staging.*$' + 'CREATE SCHEMA IF NOT EXISTS "iceberg"."staging_bar" WITH (LOCATION=\'s3://iceberg-catalog/foo_staging_bar\')', # match '^iceberg\.staging.*$' + 'CREATE SCHEMA IF NOT EXISTS "catalog"."staging_customers"', # no match + 'CREATE SCHEMA IF NOT EXISTS "landing"."transactions" WITH (LOCATION=\'s3://raw-data/landing/transactions\')', # match '^landing\..*$' + ] def test_session_authorization(trino_mocked_engine_adapter: TrinoEngineAdapter): @@ -827,7 +849,10 @@ def test_insert_overwrite_time_partition_hive( end="2022-01-02", time_column="b", time_formatter=lambda x, _: exp.Literal.string(to_ds(x)), - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("STRING"), + }, ) assert to_sql_calls(adapter) == [ @@ -865,7 +890,10 @@ def test_insert_overwrite_time_partition_iceberg( end="2022-01-02", time_column="b", time_formatter=lambda x, _: exp.Literal.string(to_ds(x)), - target_columns_to_types={"a": exp.DataType.build("INT"), "b": exp.DataType.build("STRING")}, + target_columns_to_types={ + "a": exp.DataType.build("INT"), + "b": exp.DataType.build("STRING"), + }, ) assert to_sql_calls(adapter) == [ @@ -874,7 +902,9 @@ def test_insert_overwrite_time_partition_iceberg( ] -def test_delta_timestamps_with_non_timestamp_columns(make_mocked_engine_adapter: t.Callable): +def test_delta_timestamps_with_non_timestamp_columns( + make_mocked_engine_adapter: t.Callable, +): """Test that _apply_timestamp_mapping + _to_delta_ts handles non-timestamp columns.""" config = TrinoConnectionConfig( user="user", @@ -904,7 +934,9 @@ def test_delta_timestamps_with_non_timestamp_columns(make_mocked_engine_adapter: mapped_columns_to_types, mapped_column_names = adapter._apply_timestamp_mapping( columns_to_types ) - delta_columns_to_types = adapter._to_delta_ts(mapped_columns_to_types, mapped_column_names) + delta_columns_to_types = adapter._to_delta_ts( + mapped_columns_to_types, mapped_column_names + ) # TIMESTAMP is in mapping → TIMESTAMP(3), skipped by _to_delta_ts # TIMESTAMP(1) is NOT in mapping (exact match), uses default TIMESTAMP → ts6 diff --git a/tests/core/integration/test_audits.py b/tests/core/integration/test_audits.py index 457974fdac..d227152d8b 100644 --- a/tests/core/integration/test_audits.py +++ b/tests/core/integration/test_audits.py @@ -1,23 +1,19 @@ from __future__ import annotations import typing as t +from pathlib import Path from textwrap import dedent + import pytest -from pathlib import Path import time_machine -from sqlglot import exp from IPython.utils.capture import capture_output +from sqlglot import exp -from sqlmesh.core.config import ( - Config, - ModelDefaultsConfig, -) +from sqlmesh.core.config import Config, ModelDefaultsConfig from sqlmesh.core.context import Context -from sqlmesh.utils.errors import ( - PlanError, -) -from tests.utils.test_helpers import use_terminal_console +from sqlmesh.utils.errors import PlanError from tests.utils.test_filesystem import create_temp_file +from tests.utils.test_helpers import use_terminal_console pytestmark = pytest.mark.slow diff --git a/tests/core/integration/test_auto_restatement.py b/tests/core/integration/test_auto_restatement.py index 1bda373a8f..b568929dcb 100644 --- a/tests/core/integration/test_auto_restatement.py +++ b/tests/core/integration/test_auto_restatement.py @@ -1,6 +1,7 @@ from __future__ import annotations import typing as t + import pandas as pd # noqa: TID253 import pytest import time_machine @@ -8,9 +9,7 @@ from sqlmesh.core import dialect as d from sqlmesh.core.macros import macro -from sqlmesh.core.model import ( - load_sql_based_model, -) +from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.plan import SnapshotIntervals from sqlmesh.utils.date import to_timestamp @@ -32,11 +31,16 @@ def record_intervals( if evaluator.runtime_stage == "evaluating": evaluator.engine_adapter.insert_append( "_test_auto_restatement_intervals", - pd.DataFrame({"name": [name.name], "start_ds": [start.name], "end_ds": [end.name]}), + pd.DataFrame( + { + "name": [name.name], + "start_ds": [start.name], + "end_ds": [end.name], + } + ), ) - new_model_expr = d.parse( - """ + new_model_expr = d.parse(""" MODEL ( name memory.sushi.new_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -50,13 +54,11 @@ def record_intervals( @record_intervals('new_model', @start_ds, @end_ds); SELECT '2023-01-07' AS ds, 1 AS a; - """ - ) + """) new_model = load_sql_based_model(new_model_expr) context.upsert_model(new_model) - new_model_downstream_expr = d.parse( - """ + new_model_downstream_expr = d.parse(""" MODEL ( name memory.sushi.new_model_downstream, kind INCREMENTAL_BY_TIME_RANGE ( @@ -68,8 +70,7 @@ def record_intervals( @record_intervals('new_model_downstream', @start_ts, @end_ts); SELECT * FROM memory.sushi.new_model; - """ - ) + """) new_model_downstream = load_sql_based_model(new_model_downstream_expr) context.upsert_model(new_model_downstream) @@ -125,8 +126,7 @@ def test_run_auto_restatement_plan_preview(init_and_plan_context: t.Callable): context, init_plan = init_and_plan_context("examples/sushi") context.apply(init_plan) - new_model_expr = d.parse( - """ + new_model_expr = d.parse(""" MODEL ( name memory.sushi.new_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -137,8 +137,7 @@ def test_run_auto_restatement_plan_preview(init_and_plan_context: t.Callable): ); SELECT '2023-01-07' AS ds, 1 AS a; - """ - ) + """) new_model = load_sql_based_model(new_model_expr) context.upsert_model(new_model) snapshot = context.get_snapshot(new_model.name) @@ -182,8 +181,7 @@ def fail_auto_restatement(evaluator, start: exp.Expr, **kwargs: t.Any) -> None: if evaluator.runtime_stage == "evaluating" and start.name != "2023-01-01": raise Exception("Failed") - new_model_expr = d.parse( - """ + new_model_expr = d.parse(""" MODEL ( name memory.sushi.new_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -197,8 +195,7 @@ def fail_auto_restatement(evaluator, start: exp.Expr, **kwargs: t.Any) -> None: @fail_auto_restatement(@start_ds); SELECT '2023-01-07' AS ds, 1 AS a; - """ - ) + """) new_model = load_sql_based_model(new_model_expr) context.upsert_model(new_model) diff --git a/tests/core/integration/test_aux_commands.py b/tests/core/integration/test_aux_commands.py index 7de585576d..d2e069365d 100644 --- a/tests/core/integration/test_aux_commands.py +++ b/tests/core/integration/test_aux_commands.py @@ -1,33 +1,27 @@ from __future__ import annotations import typing as t +from pathlib import Path from unittest.mock import patch + import pytest -from pathlib import Path -from sqlmesh.core.config.naming import NameInferenceConfig -from sqlmesh.core.model.common import ParsableSql import time_machine from pytest_mock.plugin import MockerFixture -from sqlmesh.core.config import ( - Config, - GatewayConfig, - ModelDefaultsConfig, - DuckDBConnectionConfig, -) +from sqlmesh.core.config import (Config, DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig) from sqlmesh.core.config.janitor import JanitorConfig +from sqlmesh.core.config.naming import NameInferenceConfig from sqlmesh.core.context import Context -from sqlmesh.core.model import ( - SqlModel, -) -from sqlmesh.utils.errors import ( - SQLMeshError, -) +from sqlmesh.core.model import SqlModel +from sqlmesh.core.model.common import ParsableSql from sqlmesh.utils.date import now +from sqlmesh.utils.errors import SQLMeshError from tests.conftest import DuckDBMetadata -from tests.utils.test_helpers import use_terminal_console +from tests.core.integration.utils import (add_projection_to_model, + apply_to_environment) from tests.utils.test_filesystem import create_temp_file -from tests.core.integration.utils import add_projection_to_model, apply_to_environment +from tests.utils.test_helpers import use_terminal_console pytestmark = pytest.mark.slow @@ -69,7 +63,9 @@ def test_table_name(init_and_plan_context: t.Callable): # Make a forward-only change context.upsert_model(model, stamp="forward_only") - context.plan("dev_b", auto_apply=True, no_prompts=True, skip_tests=True, forward_only=True) + context.plan( + "dev_b", auto_apply=True, no_prompts=True, skip_tests=True, forward_only=True + ) forward_only_snapshot = context.get_snapshot("sushi.waiter_revenue_by_day") assert forward_only_snapshot.version == snapshot.version @@ -133,7 +129,9 @@ def setup_scenario(): ctx, model1_snapshot = setup_scenario() # - Check that the snapshot record exists in the state sync - state_snapshot = ctx.state_sync.state_sync.get_snapshots([model1_snapshot.snapshot_id]) + state_snapshot = ctx.state_sync.state_sync.get_snapshots( + [model1_snapshot.snapshot_id] + ) assert state_snapshot # - Run the janitor again, this time it should succeed @@ -141,7 +139,9 @@ def setup_scenario(): ctx._run_janitor(ignore_ttl=True) # - Check that the snapshot record does not exist in the state sync anymore - state_snapshot = ctx.state_sync.state_sync.get_snapshots([model1_snapshot.snapshot_id]) + state_snapshot = ctx.state_sync.state_sync.get_snapshots( + [model1_snapshot.snapshot_id] + ) assert not state_snapshot # Case 2: Assume that the view cleanup yields an error, the enviroment @@ -166,10 +166,14 @@ def setup_scenario(): assert not ctx.state_sync.get_environment("dev") -def test_janitor_aggregates_failures_into_single_error(mocker: MockerFixture, tmp_path: Path): +def test_janitor_aggregates_failures_into_single_error( + mocker: MockerFixture, tmp_path: Path +): models_dir = tmp_path / "models" models_dir.mkdir() - (models_dir / "model1.sql").write_text("MODEL(name test.model1, kind FULL); SELECT 1 AS col") + (models_dir / "model1.sql").write_text( + "MODEL(name test.model1, kind FULL); SELECT 1 AS col" + ) ctx = Context( paths=[tmp_path], @@ -196,7 +200,9 @@ def test_janitor_warn_on_delete_failure_downgrades_aggregated_error( ): models_dir = tmp_path / "models" models_dir.mkdir() - (models_dir / "model1.sql").write_text("MODEL(name test.model1, kind FULL); SELECT 1 AS col") + (models_dir / "model1.sql").write_text( + "MODEL(name test.model1, kind FULL); SELECT 1 AS col" + ) ctx = Context( paths=[tmp_path], @@ -229,7 +235,9 @@ def test_janitor_force_delete_removes_environment_state_despite_drop_failure( ): models_dir = tmp_path / "models" models_dir.mkdir() - (models_dir / "model1.sql").write_text("MODEL(name test.model1, kind FULL); SELECT 1 AS col") + (models_dir / "model1.sql").write_text( + "MODEL(name test.model1, kind FULL); SELECT 1 AS col" + ) ctx = Context( paths=[tmp_path], @@ -401,8 +409,12 @@ def test_destroy(copy_to_temp_path): context.fetchdf(f"SELECT * FROM db_1.sqlmesh.{table_name}") # The actual tables as well - context.engine_adapters["second"].fetchdf(f"SELECT * FROM db_2.second_schema.model_one") - context.engine_adapters["second"].fetchdf(f"SELECT * FROM db_2.second_schema.model_two") + context.engine_adapters["second"].fetchdf( + f"SELECT * FROM db_2.second_schema.model_one" + ) + context.engine_adapters["second"].fetchdf( + f"SELECT * FROM db_2.second_schema.model_two" + ) context.fetchdf(f"SELECT * FROM db_1.first_schema.model_one") context.fetchdf(f"SELECT * FROM db_1.first_schema.model_two") @@ -414,7 +426,8 @@ def test_destroy(copy_to_temp_path): # Ensure all tables have been removed for table_name in state_tables: with pytest.raises( - Exception, match=f"Catalog Error: Table with name {table_name} does not exist!" + Exception, + match=f"Catalog Error: Table with name {table_name} does not exist!", ): context.fetchdf(f"SELECT * FROM db_1.sqlmesh.{table_name}") @@ -431,11 +444,15 @@ def test_destroy(copy_to_temp_path): with pytest.raises( Exception, match=r"Catalog Error: Table with name.*model_two.*does not exist" ): - context.engine_adapters["second"].fetchdf("SELECT * FROM db_2.second_schema.model_two") + context.engine_adapters["second"].fetchdf( + "SELECT * FROM db_2.second_schema.model_two" + ) with pytest.raises( Exception, match=r"Catalog Error: Table with name.*model_one.*does not exist" ): - context.engine_adapters["second"].fetchdf("SELECT * FROM db_2.second_schema.model_one") + context.engine_adapters["second"].fetchdf( + "SELECT * FROM db_2.second_schema.model_one" + ) # Ensure the cache has been removed assert not cache_path.exists() @@ -443,7 +460,9 @@ def test_destroy(copy_to_temp_path): @use_terminal_console def test_render_path_instead_of_model(tmp_path: Path): - create_temp_file(tmp_path, Path("models/test.sql"), "MODEL (name test_model); SELECT 1 AS col") + create_temp_file( + tmp_path, Path("models/test.sql"), "MODEL (name test_model); SELECT 1 AS col" + ) ctx = Context(paths=tmp_path, config=Config()) # Case 1: Fail gracefully when the user is passing in a path instead of a model name @@ -455,7 +474,9 @@ def test_render_path_instead_of_model(tmp_path: Path): ctx.render(test_model) # Case 2: Fail gracefully when the model name is not found - with pytest.raises(SQLMeshError, match="Cannot find model with name 'incorrect_model'"): + with pytest.raises( + SQLMeshError, match="Cannot find model with name 'incorrect_model'" + ): ctx.render("incorrect_model") # Case 3: Render the model successfully @@ -492,7 +513,9 @@ def test_evaluate_uncategorized_snapshot(init_and_plan_context: t.Callable): # Downstream model references the new projection downstream_model = context.get_model("sushi.top_waiters") - context.upsert_model(add_projection_to_model(t.cast(SqlModel, downstream_model), literal=False)) + context.upsert_model( + add_projection_to_model(t.cast(SqlModel, downstream_model), literal=False) + ) df = context.evaluate( "sushi.top_waiters", start="2023-01-05", end="2023-01-06", execution_time=now() diff --git a/tests/core/integration/test_change_scenarios.py b/tests/core/integration/test_change_scenarios.py index fb1762220f..61f2ad8bfc 100644 --- a/tests/core/integration/test_change_scenarios.py +++ b/tests/core/integration/test_change_scenarios.py @@ -1,58 +1,41 @@ from __future__ import annotations -import typing as t import json +import re +import typing as t from datetime import timedelta +from pathlib import Path from unittest import mock + import pandas as pd # noqa: TID253 import pytest -from pathlib import Path -from sqlmesh.core.model.common import ParsableSql import time_machine from sqlglot.expressions import DataType -import re from sqlmesh.cli.project_init import init_example_project from sqlmesh.core import constants as c from sqlmesh.core import dialect as d -from sqlmesh.core.config import ( - AutoCategorizationMode, - Config, - GatewayConfig, - ModelDefaultsConfig, - DuckDBConnectionConfig, -) -from sqlmesh.core.context import Context +from sqlmesh.core.config import (AutoCategorizationMode, Config, + DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig) from sqlmesh.core.config.categorizer import CategorizerConfig -from sqlmesh.core.model import ( - FullKind, - ModelKind, - ModelKindName, - SqlModel, - PythonModel, - ViewKind, - load_sql_based_model, -) +from sqlmesh.core.context import Context +from sqlmesh.core.model import (FullKind, ModelKind, ModelKindName, + PythonModel, SqlModel, ViewKind, + load_sql_based_model) +from sqlmesh.core.model.common import ParsableSql from sqlmesh.core.model.kind import model_kind_type_from_name from sqlmesh.core.plan import Plan, SnapshotIntervals -from sqlmesh.core.snapshot import ( - SnapshotChangeCategory, -) +from sqlmesh.core.snapshot import SnapshotChangeCategory from sqlmesh.utils.date import now, to_timestamp -from sqlmesh.utils.errors import ( - SQLMeshError, -) -from tests.core.integration.utils import ( - apply_to_environment, - add_projection_to_model, - initial_add, - change_data_type, - validate_apply_basics, - change_model_kind, - validate_model_kind_change, - validate_query_change, - validate_plan_changes, -) +from sqlmesh.utils.errors import SQLMeshError +from tests.core.integration.utils import (add_projection_to_model, + apply_to_environment, + change_data_type, change_model_kind, + initial_add, validate_apply_basics, + validate_model_kind_change, + validate_plan_changes, + validate_query_change) pytestmark = pytest.mark.slow @@ -70,7 +53,9 @@ def test_auto_categorization(sushi_context: Context): "sushi.waiter_as_customer_by_day", raise_if_missing=True ).fingerprint - model = t.cast(SqlModel, sushi_context.get_model("sushi.customers", raise_if_missing=True)) + model = t.cast( + SqlModel, sushi_context.get_model("sushi.customers", raise_if_missing=True) + ) sushi_context.upsert_model( "sushi.customers", query_=ParsableSql(sql=model.query.select("'foo' AS foo").sql(dialect=model.dialect)), # type: ignore @@ -90,7 +75,9 @@ def test_auto_categorization(sushi_context: Context): != fingerprint ) assert ( - sushi_context.get_snapshot("sushi.waiter_as_customer_by_day", raise_if_missing=True).version + sushi_context.get_snapshot( + "sushi.waiter_as_customer_by_day", raise_if_missing=True + ).version == version ) @@ -98,7 +85,9 @@ def test_auto_categorization(sushi_context: Context): @time_machine.travel("2023-01-08 15:00:00 UTC") def test_breaking_only_impacts_immediate_children(init_and_plan_context: t.Callable): context, _ = init_and_plan_context("examples/sushi") - context.upsert_model(context.get_model("sushi.top_waiters").copy(update={"kind": FullKind()})) + context.upsert_model( + context.get_model("sushi.top_waiters").copy(update={"kind": FullKind()}) + ) context.plan("prod", skip_tests=True, auto_apply=True, no_prompts=True) breaking_model = context.get_model("sushi.orders") @@ -108,8 +97,12 @@ def test_breaking_only_impacts_immediate_children(init_and_plan_context: t.Calla non_breaking_model = context.get_model("sushi.waiter_revenue_by_day") context.upsert_model(add_projection_to_model(t.cast(SqlModel, non_breaking_model))) - non_breaking_snapshot = context.get_snapshot(non_breaking_model, raise_if_missing=True) - top_waiter_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + non_breaking_snapshot = context.get_snapshot( + non_breaking_model, raise_if_missing=True + ) + top_waiter_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan_builder = context.plan_builder("dev", skip_tests=True, enable_preview=False) plan_builder.set_choice(breaking_snapshot, SnapshotChangeCategory.BREAKING) @@ -127,7 +120,9 @@ def test_breaking_only_impacts_immediate_children(init_and_plan_context: t.Calla == SnapshotChangeCategory.INDIRECT_NON_BREAKING ) assert plan.start == to_timestamp("2023-01-01") - assert not any(i.snapshot_id == top_waiter_snapshot.snapshot_id for i in plan.missing_intervals) + assert not any( + i.snapshot_id == top_waiter_snapshot.snapshot_id for i in plan.missing_intervals + ) context.apply(plan) assert ( @@ -150,7 +145,12 @@ def test_breaking_only_impacts_immediate_children(init_and_plan_context: t.Calla @pytest.mark.parametrize( "context_fixture", - ["sushi_context", "sushi_dbt_context", "sushi_test_dbt_context", "sushi_no_default_catalog"], + [ + "sushi_context", + "sushi_dbt_context", + "sushi_test_dbt_context", + "sushi_no_default_catalog", + ], ) def test_model_add(context_fixture: Context, request): initial_add(request.getfixturevalue(context_fixture), "dev") @@ -171,11 +171,16 @@ def _validate_plan(context, plan): assert not plan.missing_intervals def _validate_apply(context): - assert not sushi_context.get_snapshot("sushi.top_waiters", raise_if_missing=False) + assert not sushi_context.get_snapshot( + "sushi.top_waiters", raise_if_missing=False + ) assert sushi_context.state_reader.get_snapshots([top_waiters_snapshot_id]) env = sushi_context.state_reader.get_environment(environment) assert env - assert all(snapshot.name != '"memory"."sushi"."top_waiters"' for snapshot in env.snapshots) + assert all( + snapshot.name != '"memory"."sushi"."top_waiters"' + for snapshot in env.snapshots + ) apply_to_environment( sushi_context, @@ -189,13 +194,17 @@ def _validate_apply(context): def test_non_breaking_change(sushi_context: Context): environment = "dev" initial_add(sushi_context, environment) - validate_query_change(sushi_context, environment, SnapshotChangeCategory.NON_BREAKING, False) + validate_query_change( + sushi_context, environment, SnapshotChangeCategory.NON_BREAKING, False + ) def test_breaking_change(sushi_context: Context): environment = "dev" initial_add(sushi_context, environment) - validate_query_change(sushi_context, environment, SnapshotChangeCategory.BREAKING, False) + validate_query_change( + sushi_context, environment, SnapshotChangeCategory.BREAKING, False + ) def test_logical_change(sushi_context: Context): @@ -211,7 +220,9 @@ def test_logical_change(sushi_context: Context): DataType.Type.DOUBLE, DataType.Type.FLOAT, ) - apply_to_environment(sushi_context, environment, SnapshotChangeCategory.NON_BREAKING) + apply_to_environment( + sushi_context, environment, SnapshotChangeCategory.NON_BREAKING + ) change_data_type( sushi_context, @@ -219,7 +230,9 @@ def test_logical_change(sushi_context: Context): DataType.Type.FLOAT, DataType.Type.DOUBLE, ) - apply_to_environment(sushi_context, environment, SnapshotChangeCategory.NON_BREAKING) + apply_to_environment( + sushi_context, environment, SnapshotChangeCategory.NON_BREAKING + ) assert ( sushi_context.get_snapshot("sushi.items", raise_if_missing=True).version @@ -234,13 +247,19 @@ def test_logical_change(sushi_context: Context): (ModelKindName.FULL, ModelKindName.INCREMENTAL_BY_TIME_RANGE), ], ) -def test_model_kind_change(from_: ModelKindName, to: ModelKindName, sushi_context: Context): +def test_model_kind_change( + from_: ModelKindName, to: ModelKindName, sushi_context: Context +): environment = f"test_model_kind_change__{from_.value.lower()}__{to.value.lower()}" - incremental_snapshot = sushi_context.get_snapshot("sushi.items", raise_if_missing=True).copy() + incremental_snapshot = sushi_context.get_snapshot( + "sushi.items", raise_if_missing=True + ).copy() if from_ != ModelKindName.INCREMENTAL_BY_TIME_RANGE: change_model_kind(sushi_context, from_) - apply_to_environment(sushi_context, environment, SnapshotChangeCategory.NON_BREAKING) + apply_to_environment( + sushi_context, environment, SnapshotChangeCategory.NON_BREAKING + ) if to == ModelKindName.INCREMENTAL_BY_TIME_RANGE: sushi_context.upsert_model(incremental_snapshot.model) @@ -291,17 +310,23 @@ def test_environment_promotion(sushi_context: Context): initial_add(sushi_context, "dev") # Simulate prod "ahead" - change_data_type(sushi_context, "sushi.items", DataType.Type.DOUBLE, DataType.Type.FLOAT) + change_data_type( + sushi_context, "sushi.items", DataType.Type.DOUBLE, DataType.Type.FLOAT + ) apply_to_environment(sushi_context, "prod", SnapshotChangeCategory.BREAKING) # Simulate rebase apply_to_environment(sushi_context, "dev", SnapshotChangeCategory.BREAKING) # Make changes in dev - change_data_type(sushi_context, "sushi.items", DataType.Type.FLOAT, DataType.Type.DECIMAL) + change_data_type( + sushi_context, "sushi.items", DataType.Type.FLOAT, DataType.Type.DECIMAL + ) apply_to_environment(sushi_context, "dev", SnapshotChangeCategory.NON_BREAKING) - change_data_type(sushi_context, "sushi.top_waiters", DataType.Type.DOUBLE, DataType.Type.INT) + change_data_type( + sushi_context, "sushi.top_waiters", DataType.Type.DOUBLE, DataType.Type.INT + ) apply_to_environment(sushi_context, "dev", SnapshotChangeCategory.BREAKING) change_data_type( @@ -319,7 +344,9 @@ def test_environment_promotion(sushi_context: Context): # Promote to prod def _validate_plan(context, plan): - sushi_items_snapshot = context.get_snapshot("sushi.items", raise_if_missing=True) + sushi_items_snapshot = context.get_snapshot( + "sushi.items", raise_if_missing=True + ) sushi_top_waiters_snapshot = context.get_snapshot( "sushi.top_waiters", raise_if_missing=True ) @@ -328,17 +355,21 @@ def _validate_plan(context, plan): ) assert ( - plan.context_diff.modified_snapshots[sushi_items_snapshot.name][0].change_category + plan.context_diff.modified_snapshots[sushi_items_snapshot.name][ + 0 + ].change_category == SnapshotChangeCategory.NON_BREAKING ) assert ( - plan.context_diff.modified_snapshots[sushi_top_waiters_snapshot.name][0].change_category + plan.context_diff.modified_snapshots[sushi_top_waiters_snapshot.name][ + 0 + ].change_category == SnapshotChangeCategory.BREAKING ) assert ( - plan.context_diff.modified_snapshots[sushi_customer_revenue_by_day_snapshot.name][ - 0 - ].change_category + plan.context_diff.modified_snapshots[ + sushi_customer_revenue_by_day_snapshot.name + ][0].change_category == SnapshotChangeCategory.NON_BREAKING ) assert plan.context_diff.snapshots[ @@ -372,7 +403,9 @@ def test_no_override(sushi_context: Context) -> None: plan_builder = sushi_context.plan_builder("prod") plan = plan_builder.build() - sushi_items_snapshot = sushi_context.get_snapshot("sushi.items", raise_if_missing=True) + sushi_items_snapshot = sushi_context.get_snapshot( + "sushi.items", raise_if_missing=True + ) sushi_order_items_snapshot = sushi_context.get_snapshot( "sushi.order_items", raise_if_missing=True ) @@ -382,7 +415,9 @@ def test_no_override(sushi_context: Context) -> None: items = plan.context_diff.snapshots[sushi_items_snapshot.snapshot_id] order_items = plan.context_diff.snapshots[sushi_order_items_snapshot.snapshot_id] - waiter_revenue = plan.context_diff.snapshots[sushi_water_revenue_by_day_snapshot.snapshot_id] + waiter_revenue = plan.context_diff.snapshots[ + sushi_water_revenue_by_day_snapshot.snapshot_id + ] plan_builder.set_choice(items, SnapshotChangeCategory.BREAKING).set_choice( order_items, SnapshotChangeCategory.NON_BREAKING @@ -424,7 +459,9 @@ def test_revert( expected: SnapshotChangeCategory, ): environment = "prod" - original_snapshot_id = sushi_context.get_snapshot("sushi.items", raise_if_missing=True) + original_snapshot_id = sushi_context.get_snapshot( + "sushi.items", raise_if_missing=True + ) types = (DataType.Type.DOUBLE, DataType.Type.FLOAT, DataType.Type.DECIMAL) assert len(change_categories) < len(types) @@ -433,13 +470,18 @@ def test_revert( change_data_type(sushi_context, "sushi.items", *types[i : i + 2]) apply_to_environment(sushi_context, environment, category) assert ( - sushi_context.get_snapshot("sushi.items", raise_if_missing=True) != original_snapshot_id + sushi_context.get_snapshot("sushi.items", raise_if_missing=True) + != original_snapshot_id ) - change_data_type(sushi_context, "sushi.items", types[len(change_categories)], types[0]) + change_data_type( + sushi_context, "sushi.items", types[len(change_categories)], types[0] + ) def _validate_plan(_, plan): - snapshot = next(s for s in plan.snapshots.values() if s.name == '"memory"."sushi"."items"') + snapshot = next( + s for s in plan.snapshots.values() if s.name == '"memory"."sushi"."items"' + ) assert snapshot.change_category == expected assert not plan.missing_intervals @@ -449,12 +491,17 @@ def _validate_plan(_, plan): change_categories[-1], plan_validators=[_validate_plan], ) - assert sushi_context.get_snapshot("sushi.items", raise_if_missing=True) == original_snapshot_id + assert ( + sushi_context.get_snapshot("sushi.items", raise_if_missing=True) + == original_snapshot_id + ) def test_revert_after_downstream_change(sushi_context: Context): environment = "prod" - change_data_type(sushi_context, "sushi.items", DataType.Type.DOUBLE, DataType.Type.FLOAT) + change_data_type( + sushi_context, "sushi.items", DataType.Type.DOUBLE, DataType.Type.FLOAT + ) apply_to_environment(sushi_context, environment, SnapshotChangeCategory.BREAKING) change_data_type( @@ -463,12 +510,18 @@ def test_revert_after_downstream_change(sushi_context: Context): DataType.Type.DOUBLE, DataType.Type.FLOAT, ) - apply_to_environment(sushi_context, environment, SnapshotChangeCategory.NON_BREAKING) + apply_to_environment( + sushi_context, environment, SnapshotChangeCategory.NON_BREAKING + ) - change_data_type(sushi_context, "sushi.items", DataType.Type.FLOAT, DataType.Type.DOUBLE) + change_data_type( + sushi_context, "sushi.items", DataType.Type.FLOAT, DataType.Type.DOUBLE + ) def _validate_plan(_, plan): - snapshot = next(s for s in plan.snapshots.values() if s.name == '"memory"."sushi"."items"') + snapshot = next( + s for s in plan.snapshots.values() if s.name == '"memory"."sushi"."items"' + ) assert snapshot.change_category == SnapshotChangeCategory.BREAKING assert plan.missing_intervals @@ -481,7 +534,9 @@ def _validate_plan(_, plan): @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_context: t.Callable): +def test_indirect_non_breaking_change_after_forward_only_in_dev( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") # Make sure that the most downstream model is a materialized model. model = context.get_model("sushi.top_waiters") @@ -492,7 +547,9 @@ def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_co # Make sushi.orders a forward-only model. model = context.get_model("sushi.orders") updated_model_kind = model.kind.copy(update={"forward_only": True}) - model = model.copy(update={"stamp": "force new version", "kind": updated_model_kind}) + model = model.copy( + update={"stamp": "force new version", "kind": updated_model_kind} + ) context.upsert_model(model) snapshot = context.get_snapshot(model, raise_if_missing=True) @@ -513,7 +570,9 @@ def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_co # Make a non-breaking change to a model. model = context.get_model("sushi.top_waiters") context.upsert_model(add_projection_to_model(t.cast(SqlModel, model))) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=False).build() assert len(plan.new_snapshots) == 1 @@ -546,12 +605,16 @@ def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_co waiter_revenue_by_day_snapshot = context.get_snapshot( "sushi.waiter_revenue_by_day", raise_if_missing=True ) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=False).build() assert len(plan.new_snapshots) == 2 assert ( - plan.context_diff.snapshots[waiter_revenue_by_day_snapshot.snapshot_id].change_category + plan.context_diff.snapshots[ + waiter_revenue_by_day_snapshot.snapshot_id + ].change_category == SnapshotChangeCategory.NON_BREAKING ) assert ( @@ -630,7 +693,9 @@ def test_plan_repairs_unrenderable_snapshot_state( # Manually corrupt the snapshot's query raw_snapshot = context.state_sync.state_sync.engine_adapter.fetchone( f"SELECT snapshot FROM sqlmesh._snapshots WHERE name = '{target_snapshot.name}' AND identifier = '{target_snapshot.identifier}'" - )[0] # type: ignore + )[ + 0 + ] # type: ignore parsed_snapshot = json.loads(raw_snapshot) parsed_snapshot["node"]["query"] = "SELECT @missing_macro()" context.state_sync.state_sync.engine_adapter.update_table( @@ -640,9 +705,9 @@ def test_plan_repairs_unrenderable_snapshot_state( ) context.clear_caches() - target_snapshot_in_state = context.state_sync.get_snapshots([target_snapshot.snapshot_id])[ - target_snapshot.snapshot_id - ] + target_snapshot_in_state = context.state_sync.get_snapshots( + [target_snapshot.snapshot_id] + )[target_snapshot.snapshot_id] with pytest.raises(Exception): target_snapshot_in_state.model.render_query_or_raise() @@ -654,7 +719,9 @@ def test_plan_repairs_unrenderable_snapshot_state( plan_builder = context.plan_builder("prod", forward_only=forward_only) plan = plan_builder.build() if not forward_only: - assert target_snapshot.snapshot_id in {i.snapshot_id for i in plan.missing_intervals} + assert target_snapshot.snapshot_id in { + i.snapshot_id for i in plan.missing_intervals + } assert plan.directly_modified == {target_snapshot.snapshot_id} plan_builder.set_choice(target_snapshot, SnapshotChangeCategory.NON_BREAKING) plan = plan_builder.build() @@ -663,14 +730,16 @@ def test_plan_repairs_unrenderable_snapshot_state( context.clear_caches() assert context.get_snapshot(target_snapshot.name).model.render_query_or_raise() - target_snapshot_in_state = context.state_sync.get_snapshots([target_snapshot.snapshot_id])[ - target_snapshot.snapshot_id - ] + target_snapshot_in_state = context.state_sync.get_snapshots( + [target_snapshot.snapshot_id] + )[target_snapshot.snapshot_id] assert target_snapshot_in_state.model.render_query_or_raise() @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_no_backfill_for_model_downstream_of_metadata_change(init_and_plan_context: t.Callable): +def test_no_backfill_for_model_downstream_of_metadata_change( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") # Make sushi.waiter_revenue_by_day a forward-only model. @@ -696,9 +765,13 @@ def test_no_backfill_for_model_downstream_of_metadata_change(init_and_plan_conte @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_plan_set_choice_is_reflected_in_missing_intervals(init_and_plan_context: t.Callable): +def test_plan_set_choice_is_reflected_in_missing_intervals( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") - context.upsert_model(context.get_model("sushi.top_waiters").copy(update={"kind": FullKind()})) + context.upsert_model( + context.get_model("sushi.top_waiters").copy(update={"kind": FullKind()}) + ) context.plan("prod", skip_tests=True, no_prompts=True, auto_apply=True) model_name = "sushi.waiter_revenue_by_day" @@ -706,7 +779,9 @@ def test_plan_set_choice_is_reflected_in_missing_intervals(init_and_plan_context model = context.get_model(model_name) context.upsert_model(add_projection_to_model(t.cast(SqlModel, model))) snapshot = context.get_snapshot(model, raise_if_missing=True) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan_builder = context.plan_builder("dev", skip_tests=True) plan = plan_builder.build() @@ -737,7 +812,8 @@ def test_plan_set_choice_is_reflected_in_missing_intervals(init_and_plan_context # Change the category to BREAKING plan = plan_builder.set_choice( - plan.context_diff.snapshots[snapshot.snapshot_id], SnapshotChangeCategory.BREAKING + plan.context_diff.snapshots[snapshot.snapshot_id], + SnapshotChangeCategory.BREAKING, ).build() assert ( plan.context_diff.snapshots[snapshot.snapshot_id].change_category @@ -776,7 +852,8 @@ def test_plan_set_choice_is_reflected_in_missing_intervals(init_and_plan_context # Change the category back to NON_BREAKING plan = plan_builder.set_choice( - plan.context_diff.snapshots[snapshot.snapshot_id], SnapshotChangeCategory.NON_BREAKING + plan.context_diff.snapshots[snapshot.snapshot_id], + SnapshotChangeCategory.NON_BREAKING, ).build() assert ( plan.context_diff.snapshots[snapshot.snapshot_id].change_category @@ -890,7 +967,10 @@ def test_plan_production_environment_statements(tmp_path: Path): assert environment_statements[0].before_all == before_all assert environment_statements[0].after_all == after_all assert environment_statements[0].python_env.keys() == {"__sqlmesh__vars__"} - assert environment_statements[0].python_env["__sqlmesh__vars__"].payload == "{'var_5': 5}" + assert ( + environment_statements[0].python_env["__sqlmesh__vars__"].payload + == "{'var_5': 5}" + ) should_create = ctx.fetchdf("select * from should_create").to_dict() assert should_create["before_all"][0] == "before_all" @@ -975,19 +1055,27 @@ def test_full_model_change_with_plan_start_not_matching_model_start( context.upsert_model(model, kind=model_kind_type_from_name("FULL")()) # type: ignore # Apply the change with --skip-backfill first and no plan start - context.plan("dev", skip_tests=True, skip_backfill=True, no_prompts=True, auto_apply=True) + context.plan( + "dev", skip_tests=True, skip_backfill=True, no_prompts=True, auto_apply=True + ) # Apply the plan again but this time don't skip backfill and set start # to be later than the model start - context.plan("dev", skip_tests=True, no_prompts=True, auto_apply=True, start="1 day ago") + context.plan( + "dev", skip_tests=True, no_prompts=True, auto_apply=True, start="1 day ago" + ) # Check that the number of rows is not 0 - row_num = context.engine_adapter.fetchone(f"SELECT COUNT(*) FROM sushi__dev.top_waiters")[0] + row_num = context.engine_adapter.fetchone( + f"SELECT COUNT(*) FROM sushi__dev.top_waiters" + )[0] assert row_num > 0 @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_hourly_model_with_lookback_no_backfill_in_dev(init_and_plan_context: t.Callable): +def test_hourly_model_with_lookback_no_backfill_in_dev( + init_and_plan_context: t.Callable, +): context, plan = init_and_plan_context("examples/sushi") model_name = "sushi.waiter_revenue_by_day" @@ -1007,11 +1095,15 @@ def test_hourly_model_with_lookback_no_backfill_in_dev(init_and_plan_context: t. context.apply(plan) top_waiters_model = context.get_model("sushi.top_waiters") - top_waiters_model = add_projection_to_model(t.cast(SqlModel, top_waiters_model), literal=True) + top_waiters_model = add_projection_to_model( + t.cast(SqlModel, top_waiters_model), literal=True + ) context.upsert_model(top_waiters_model) context.get_snapshot(model, raise_if_missing=True) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) with time_machine.travel(now() + timedelta(hours=2)): plan = context.plan_builder("dev", skip_tests=True).build() @@ -1109,7 +1201,9 @@ def test_plan_environment_statements_doesnt_cause_extra_diff(tmp_path: Path): @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_plan_snapshot_table_exists_for_promoted_snapshot(init_and_plan_context: t.Callable): +def test_plan_snapshot_table_exists_for_promoted_snapshot( + init_and_plan_context: t.Callable, +): context, plan = init_and_plan_context("examples/sushi") context.apply(plan) @@ -1119,7 +1213,9 @@ def test_plan_snapshot_table_exists_for_promoted_snapshot(init_and_plan_context: context.plan("dev", auto_apply=True, no_prompts=True, skip_tests=True) # Drop the views and make sure SQLMesh recreates them later - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) context.engine_adapter.drop_view(top_waiters_snapshot.table_name()) context.engine_adapter.drop_view(top_waiters_snapshot.table_name(False)) @@ -1155,7 +1251,9 @@ def test_plan_twice_with_star_macro_yields_no_diff(tmp_path: Path): db_path = str(tmp_path / "db.db") config = Config( - gateways={"main": GatewayConfig(connection=DuckDBConnectionConfig(database=db_path))}, + gateways={ + "main": GatewayConfig(connection=DuckDBConnectionConfig(database=db_path)) + }, model_defaults=ModelDefaultsConfig(dialect="duckdb"), ) context = Context(paths=tmp_path, config=config) @@ -1212,7 +1310,9 @@ def execute( context: Context context, _ = init_and_plan_context("examples/sushi") - with open(context.path / "models" / "python_view_model.py", mode="w", encoding="utf8") as f: + with open( + context.path / "models" / "python_view_model.py", mode="w", encoding="utf8" + ) as f: f.write(python_model_file) # monkey-patch PythonModel to default to kind: View again @@ -1233,7 +1333,9 @@ def execute( # check that run() still works even though we have a Python model with kind: View in the state snapshot_ids = [s for s in plan.directly_modified if "python_view_model" in s.name] - snapshot_from_state = list(context.state_sync.get_snapshots(snapshot_ids).values())[0] + snapshot_from_state = list(context.state_sync.get_snapshots(snapshot_ids).values())[ + 0 + ] assert snapshot_from_state.model.kind.name == ModelKindName.VIEW assert snapshot_from_state.model.source_type == "python" context.run() @@ -1255,7 +1357,10 @@ def execute( assert len(plan.directly_modified) == 1 snapshot_id = list(plan.directly_modified)[0] assert snapshot_id.name == '"memory"."sushi"."python_view_model"' - assert plan.modified_snapshots[snapshot_id].change_category == SnapshotChangeCategory.BREAKING + assert ( + plan.modified_snapshots[snapshot_id].change_category + == SnapshotChangeCategory.BREAKING + ) context.apply(plan) @@ -1329,14 +1434,18 @@ def test_rebase_two_changed_parents( # Make change A and deploy it to dev_a context.upsert_model(initial_model_a.name, stamp="1") plan_builder = context.plan_builder("dev_a", skip_tests=True) - plan_builder.set_choice(context.get_snapshot(initial_model_a.name), parent_a_category) + plan_builder.set_choice( + context.get_snapshot(initial_model_a.name), parent_a_category + ) context.apply(plan_builder.build()) # Make change B and deploy it to dev_b context.upsert_model(initial_model_a) context.upsert_model(initial_model_b.name, stamp="1") plan_builder = context.plan_builder("dev_b", skip_tests=True) - plan_builder.set_choice(context.get_snapshot(initial_model_b.name), parent_b_category) + plan_builder.set_choice( + context.get_snapshot(initial_model_b.name), parent_b_category + ) context.apply(plan_builder.build()) # Deploy change A to prod @@ -1349,10 +1458,14 @@ def test_rebase_two_changed_parents( plan = context.plan_builder("prod", skip_tests=True).build() # Validate the category of child snapshots - direct_child_snapshot = plan.snapshots[context.get_snapshot("sushi.order_items").snapshot_id] + direct_child_snapshot = plan.snapshots[ + context.get_snapshot("sushi.order_items").snapshot_id + ] assert direct_child_snapshot.change_category == expected_child_category - indirect_child_snapshot = plan.snapshots[context.get_snapshot("sushi.top_waiters").snapshot_id] + indirect_child_snapshot = plan.snapshots[ + context.get_snapshot("sushi.top_waiters").snapshot_id + ] assert indirect_child_snapshot.change_category == expected_child_category @@ -1382,13 +1495,14 @@ def test_unaligned_start_snapshots(context_fixture: Context, request): @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_unaligned_start_snapshot_with_non_deployable_downstream(init_and_plan_context: t.Callable): +def test_unaligned_start_snapshot_with_non_deployable_downstream( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") downstream_model_name = "memory.sushi.customer_max_revenue" - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name {downstream_model_name}, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -1401,8 +1515,7 @@ def test_unaligned_start_snapshot_with_non_deployable_downstream(init_and_plan_c customer_id, MAX(revenue) AS max_revenue FROM memory.sushi.customer_revenue_lifetime GROUP BY 1; - """ - ) + """) downstream_model = load_sql_based_model(expressions) assert downstream_model.forward_only @@ -1410,7 +1523,9 @@ def test_unaligned_start_snapshot_with_non_deployable_downstream(init_and_plan_c context.plan(auto_apply=True, no_prompts=True) - customer_revenue_lifetime_model = context.get_model("sushi.customer_revenue_lifetime") + customer_revenue_lifetime_model = context.get_model( + "sushi.customer_revenue_lifetime" + ) kwargs = { **customer_revenue_lifetime_model.dict(), "name": "memory.sushi.customer_revenue_lifetime_new", @@ -1454,7 +1569,9 @@ def test_indirect_non_breaking_view_is_updated_with_new_table_references( # Check the downstream view and make sure it's still queryable assert context.get_model("sushi.top_waiters").kind.is_view - row_num = context.engine_adapter.fetchone(f"SELECT COUNT(*) FROM sushi.top_waiters")[0] + row_num = context.engine_adapter.fetchone( + f"SELECT COUNT(*) FROM sushi.top_waiters" + )[0] assert row_num > 0 @@ -1463,8 +1580,7 @@ def test_annotated_self_referential_model(init_and_plan_context: t.Callable): context, _ = init_and_plan_context("examples/sushi") # Projections are fully annotated in the query but columns were not specified explicitly - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name memory.sushi.test_self_ref, kind FULL, @@ -1472,8 +1588,7 @@ def test_annotated_self_referential_model(init_and_plan_context: t.Callable): ); SELECT 1::INT AS one FROM memory.sushi.test_self_ref; - """ - ) + """) model = load_sql_based_model(expressions) assert model.depends_on_self context.upsert_model(model) @@ -1488,8 +1603,7 @@ def test_annotated_self_referential_model(init_and_plan_context: t.Callable): def test_creating_stage_for_first_batch_only(init_and_plan_context: t.Callable): context, _ = init_and_plan_context("examples/sushi") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.sushi.test_batch_size, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -1506,12 +1620,14 @@ def test_creating_stage_for_first_batch_only(init_and_plan_context: t.Callable): SELECT 1::INT AS one; @IF(@runtime_stage = 'creating', INSERT INTO test_schema.creating_counter (a) VALUES (1)); - """ - ) + """) model = load_sql_based_model(expressions) context.upsert_model(model) context.plan("prod", skip_tests=True, no_prompts=True, auto_apply=True) assert ( - context.engine_adapter.fetchone("SELECT COUNT(*) FROM test_schema.creating_counter")[0] == 1 + context.engine_adapter.fetchone( + "SELECT COUNT(*) FROM test_schema.creating_counter" + )[0] + == 1 ) diff --git a/tests/core/integration/test_config.py b/tests/core/integration/test_config.py index 5d571cd7c5..18f1c8bb66 100644 --- a/tests/core/integration/test_config.py +++ b/tests/core/integration/test_config.py @@ -1,39 +1,31 @@ from __future__ import annotations +import logging import typing as t +from pathlib import Path from unittest.mock import patch -import logging + import pytest +from IPython.utils.capture import capture_output from pytest import MonkeyPatch -from pathlib import Path from pytest_mock.plugin import MockerFixture from sqlglot import exp -from IPython.utils.capture import capture_output -from sqlmesh.core.config import ( - Config, - GatewayConfig, - ModelDefaultsConfig, - DuckDBConnectionConfig, - TableNamingConvention, - AutoCategorizationMode, -) +from sqlmesh.core.config import (AutoCategorizationMode, Config, + DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig, TableNamingConvention) from sqlmesh.core.config.common import EnvironmentSuffixTarget -from sqlmesh.core.context import Context from sqlmesh.core.config.plan import PlanConfig +from sqlmesh.core.context import Context from sqlmesh.core.engine_adapter import DuckDBEngineAdapter from sqlmesh.core.model import SqlModel from sqlmesh.core.model.common import ParsableSql -from sqlmesh.core.snapshot import ( - SnapshotChangeCategory, -) -from sqlmesh.utils.errors import ( - ConfigError, -) +from sqlmesh.core.snapshot import SnapshotChangeCategory +from sqlmesh.utils.errors import ConfigError from tests.conftest import DuckDBMetadata -from tests.utils.test_helpers import use_terminal_console -from tests.utils.test_filesystem import create_temp_file from tests.core.integration.utils import apply_to_environment, initial_add +from tests.utils.test_filesystem import create_temp_file +from tests.utils.test_helpers import use_terminal_console pytestmark = pytest.mark.slow @@ -53,7 +45,9 @@ def test_missing_connection_config(): # Case 3: Specifying a default_connection or connection in the gateway should work ctx = Context(config=Config(default_connection=DuckDBConnectionConfig())) ctx = Context( - config=Config(gateways={"default": GatewayConfig(connection=DuckDBConnectionConfig())}) + config=Config( + gateways={"default": GatewayConfig(connection=DuckDBConnectionConfig())} + ) ) @@ -63,7 +57,10 @@ def test_physical_table_naming_strategy_table_only(copy_to_temp_path: t.Callable config="table_only_naming_config", ) - assert sushi_context.config.physical_table_naming_convention == TableNamingConvention.TABLE_ONLY + assert ( + sushi_context.config.physical_table_naming_convention + == TableNamingConvention.TABLE_ONLY + ) sushi_context.plan(auto_apply=True) adapter = sushi_context.engine_adapter @@ -94,7 +91,10 @@ def test_physical_table_naming_strategy_hash_md5(copy_to_temp_path: t.Callable): config="hash_md5_naming_config", ) - assert sushi_context.config.physical_table_naming_convention == TableNamingConvention.HASH_MD5 + assert ( + sushi_context.config.physical_table_naming_convention + == TableNamingConvention.HASH_MD5 + ) sushi_context.plan(auto_apply=True) adapter = sushi_context.engine_adapter @@ -138,11 +138,15 @@ def test_environment_suffix_target_table(init_and_plan_context: t.Callable): # Make sure no new schemas are created assert set(metadata.schemas) - starting_schemas == {"raw"} dev_views = { - x for x in metadata.qualified_views if x.db in environments_schemas and "__dev" in x.name + x + for x in metadata.qualified_views + if x.db in environments_schemas and "__dev" in x.name } # Make sure that there is a view with `__dev` for each view that exists in prod assert len(dev_views) == len(prod_views) - assert {x.name.replace("__dev", "") for x in dev_views} - {x.name for x in prod_views} == set() + assert {x.name.replace("__dev", "") for x in dev_views} - { + x.name for x in prod_views + } == set() context.invalidate_environment("dev") context._run_janitor() views_after_janitor = metadata.qualified_views @@ -159,12 +163,16 @@ def test_environment_suffix_target_table(init_and_plan_context: t.Callable): } == set() -def test_environment_suffix_target_catalog(tmp_path: Path, monkeypatch: MonkeyPatch) -> None: +def test_environment_suffix_target_catalog( + tmp_path: Path, monkeypatch: MonkeyPatch +) -> None: monkeypatch.chdir(tmp_path) config = Config( model_defaults=ModelDefaultsConfig(dialect="duckdb"), - default_connection=DuckDBConnectionConfig(catalogs={"main_warehouse": ":memory:"}), + default_connection=DuckDBConnectionConfig( + catalogs={"main_warehouse": ":memory:"} + ), environment_suffix_target=EnvironmentSuffixTarget.CATALOG, ) @@ -219,18 +227,27 @@ def test_environment_suffix_target_catalog(tmp_path: Path, monkeypatch: MonkeyPa # dev should be overridden to go to a catalogs called 'main_warehouse__dev' and 'memory__dev' ctx.plan(environment="dev", include_unmodified=True, auto_apply=True) assert ( - ctx.engine_adapter.fetchone("select * from main_warehouse__dev.example_schema.test_model")[ + ctx.engine_adapter.fetchone( + "select * from main_warehouse__dev.example_schema.test_model" + )[ 0 ] # type: ignore == "1" ) assert ( - ctx.engine_adapter.fetchone("select * from memory__dev.example_fqn_schema.test_model_fqn")[ + ctx.engine_adapter.fetchone( + "select * from memory__dev.example_fqn_schema.test_model_fqn" + )[ 0 ] # type: ignore == "1" ) - assert metadata.catalogs == {"main_warehouse", "main_warehouse__dev", "memory", "memory__dev"} + assert metadata.catalogs == { + "main_warehouse", + "main_warehouse__dev", + "memory", + "memory__dev", + } # schemas in dev envs should match prod and not have a suffix assert metadata.schemas_in_catalog("main_warehouse") == [ @@ -272,12 +289,22 @@ def test_environment_suffix_target_catalog(tmp_path: Path, monkeypatch: MonkeyPa def test_environment_catalog_mapping(init_and_plan_context: t.Callable): environments_schemas = {"raw", "sushi"} - def get_prod_dev_views(metadata: DuckDBMetadata) -> t.Tuple[t.Set[exp.Table], t.Set[exp.Table]]: + def get_prod_dev_views( + metadata: DuckDBMetadata, + ) -> t.Tuple[t.Set[exp.Table], t.Set[exp.Table]]: views = metadata.qualified_views prod_views = { - x for x in views if x.catalog == "prod_catalog" if x.db in environments_schemas + x + for x in views + if x.catalog == "prod_catalog" + if x.db in environments_schemas + } + dev_views = { + x + for x in views + if x.catalog == "dev_catalog" + if x.db in environments_schemas } - dev_views = {x for x in views if x.catalog == "dev_catalog" if x.db in environments_schemas} return prod_views, dev_views def get_default_catalog_and_non_tables( @@ -391,7 +418,9 @@ def plan_with_output(ctx: Context, environment: str): logger = logging.getLogger("sqlmesh.core.state_sync.db.facade") create_temp_file( - tmp_path, models_dir / "a.sql", "MODEL (name test.a, kind FULL); SELECT 1 AS col" + tmp_path, + models_dir / "a.sql", + "MODEL (name test.a, kind FULL); SELECT 1 AS col", ) config = Config(plan=PlanConfig(always_recreate_environment=True)) @@ -404,7 +433,9 @@ def plan_with_output(ctx: Context, environment: str): # Case 2: Prod does not exist, so dev is updated create_temp_file( - tmp_path, models_dir / "a.sql", "MODEL (name test.a, kind FULL); SELECT 5 AS col" + tmp_path, + models_dir / "a.sql", + "MODEL (name test.a, kind FULL); SELECT 5 AS col", ) output = plan_with_output(ctx, "dev") @@ -416,7 +447,9 @@ def plan_with_output(ctx: Context, environment: str): # Case 4: Dev is updated with a breaking change. Prod exists now so plan comparisons moving forward should be against prod create_temp_file( - tmp_path, models_dir / "a.sql", "MODEL (name test.a, kind FULL); SELECT 10 AS col" + tmp_path, + models_dir / "a.sql", + "MODEL (name test.a, kind FULL); SELECT 10 AS col", ) ctx.load() @@ -452,17 +485,14 @@ def plan_with_output(ctx: Context, environment: str): assert "Differences from the `prod` environment" in output.stdout stdout_rstrip = "\n".join([line.rstrip() for line in output.stdout.split("\n")]) - assert ( - """MODEL ( + assert """MODEL ( name test.a, + owner test, kind FULL ) SELECT - 5 AS col -+ 10 AS col""" - in stdout_rstrip - ) ++ 10 AS col""" in stdout_rstrip # Case 6: Ensure that target environment and create_from environment are not the same output = plan_with_output(ctx, "prod") @@ -491,9 +521,7 @@ def test_before_all_after_all_execution_order(tmp_path: Path, mocker: MockerFixt f.write(model) # before_all statement that creates a table that the above model depends on - before_all_statement = ( - "CREATE TABLE IF NOT EXISTS before_all_created_table AS SELECT 1 AS id, 'test' AS value" - ) + before_all_statement = "CREATE TABLE IF NOT EXISTS before_all_created_table AS SELECT 1 AS id, 'test' AS value" # after_all that depends on the model after_all_statement = "CREATE TABLE IF NOT EXISTS after_all_created_table AS SELECT id, value FROM test_schema.model_that_depends_on_before_all" @@ -509,7 +537,11 @@ def test_before_all_after_all_execution_order(tmp_path: Path, mocker: MockerFixt original_duckdb_execute = DuckDBEngineAdapter.execute def track_duckdb_execute(self, expression, **kwargs): - sql = expression if isinstance(expression, str) else expression.sql(dialect="duckdb") + sql = ( + expression + if isinstance(expression, str) + else expression.sql(dialect="duckdb") + ) state_tables = [ "_snapshots", "_environments", @@ -555,7 +587,9 @@ def test_auto_categorization(sushi_context: Context): "sushi.waiter_as_customer_by_day", raise_if_missing=True ).fingerprint - model = t.cast(SqlModel, sushi_context.get_model("sushi.customers", raise_if_missing=True)) + model = t.cast( + SqlModel, sushi_context.get_model("sushi.customers", raise_if_missing=True) + ) sushi_context.upsert_model( "sushi.customers", query_=ParsableSql(sql=model.query.select("'foo' AS foo").sql(dialect=model.dialect)), # type: ignore @@ -575,6 +609,8 @@ def test_auto_categorization(sushi_context: Context): != fingerprint ) assert ( - sushi_context.get_snapshot("sushi.waiter_as_customer_by_day", raise_if_missing=True).version + sushi_context.get_snapshot( + "sushi.waiter_as_customer_by_day", raise_if_missing=True + ).version == version ) diff --git a/tests/core/integration/test_cron.py b/tests/core/integration/test_cron.py index fa327ac36f..bf0d6ea011 100644 --- a/tests/core/integration/test_cron.py +++ b/tests/core/integration/test_cron.py @@ -1,14 +1,12 @@ from __future__ import annotations import typing as t + import pytest import time_machine from sqlmesh.core import dialect as d -from sqlmesh.core.model import ( - SqlModel, - load_sql_based_model, -) +from sqlmesh.core.model import SqlModel, load_sql_based_model from sqlmesh.core.plan import SnapshotIntervals from sqlmesh.utils.date import to_timestamp from tests.core.integration.utils import add_projection_to_model @@ -59,7 +57,9 @@ def test_cron_not_aligned_with_day_boundary( plan = context.plan_builder("prod", skip_tests=True).build() context.apply(plan) - waiter_revenue_by_day_snapshot = context.get_snapshot(model.name, raise_if_missing=True) + waiter_revenue_by_day_snapshot = context.get_snapshot( + model.name, raise_if_missing=True + ) assert waiter_revenue_by_day_snapshot.intervals == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-07")) ] @@ -84,7 +84,9 @@ def test_cron_not_aligned_with_day_boundary( @time_machine.travel("2023-01-08 00:00:00 UTC") -def test_cron_not_aligned_with_day_boundary_new_model(init_and_plan_context: t.Callable): +def test_cron_not_aligned_with_day_boundary_new_model( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") existing_model = context.get_model("sushi.waiter_revenue_by_day") @@ -101,9 +103,7 @@ def test_cron_not_aligned_with_day_boundary_new_model(init_and_plan_context: t.C # Add a new model and make a change to a forward-only model. # The cron of the new model is not aligned with the day boundary. - new_model = load_sql_based_model( - d.parse( - """ + new_model = load_sql_based_model(d.parse(""" MODEL ( name memory.sushi.new_model, kind FULL, @@ -112,12 +112,12 @@ def test_cron_not_aligned_with_day_boundary_new_model(init_and_plan_context: t.C ); SELECT 1 AS one; - """ - ) - ) + """)) context.upsert_model(new_model) - existing_model = add_projection_to_model(t.cast(SqlModel, existing_model), literal=True) + existing_model = add_projection_to_model( + t.cast(SqlModel, existing_model), literal=True + ) context.upsert_model(existing_model) plan = context.plan_builder("dev", skip_tests=True, enable_preview=True).build() @@ -156,18 +156,26 @@ def test_parent_cron_after_child(init_and_plan_context: t.Callable): plan = context.plan_builder("prod", skip_tests=True).build() context.apply(plan) - waiter_revenue_by_day_snapshot = context.get_snapshot(model.name, raise_if_missing=True) + waiter_revenue_by_day_snapshot = context.get_snapshot( + model.name, raise_if_missing=True + ) assert waiter_revenue_by_day_snapshot.intervals == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-07")) ] top_waiters_model = context.get_model("sushi.top_waiters") - top_waiters_model = add_projection_to_model(t.cast(SqlModel, top_waiters_model), literal=True) + top_waiters_model = add_projection_to_model( + t.cast(SqlModel, top_waiters_model), literal=True + ) context.upsert_model(top_waiters_model) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) - with time_machine.travel("2023-01-08 23:55:00 UTC"): # Past parent's cron, but before child's + with time_machine.travel( + "2023-01-08 23:55:00 UTC" + ): # Past parent's cron, but before child's plan = context.plan_builder("dev", skip_tests=True).build() # Make sure the waiter_revenue_by_day model is not backfilled. assert plan.missing_intervals == [ @@ -218,7 +226,9 @@ def assert_intervals(plan, intervals): # now we're ready 8AM UTC == midnight PST with time_machine.travel("2025-03-08 08:00:00 UTC"): plan = context.plan_builder("prod", skip_tests=True).build() - assert_intervals(plan, [(to_timestamp("2025-03-07"), to_timestamp("2025-03-08"))]) + assert_intervals( + plan, [(to_timestamp("2025-03-07"), to_timestamp("2025-03-08"))] + ) with time_machine.travel("2025-03-09 07:00:00 UTC"): plan = context.plan_builder("prod", skip_tests=True).build() diff --git a/tests/core/integration/test_dbt.py b/tests/core/integration/test_dbt.py index 6f23acb97e..0b37399a7d 100644 --- a/tests/core/integration/test_dbt.py +++ b/tests/core/integration/test_dbt.py @@ -1,18 +1,14 @@ from __future__ import annotations import typing as t + import pytest -from sqlmesh.core.model.common import ParsableSql import time_machine from sqlmesh.core.context import Context -from sqlmesh.core.model import ( - IncrementalUnmanagedKind, -) -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - SnapshotChangeCategory, -) +from sqlmesh.core.model import IncrementalUnmanagedKind +from sqlmesh.core.model.common import ParsableSql +from sqlmesh.core.snapshot import DeployabilityIndex, SnapshotChangeCategory if t.TYPE_CHECKING: pass @@ -35,10 +31,19 @@ def test_dbt_select_star_is_directly_modified(sushi_test_dbt_context: Context): plan = context.plan_builder("dev", skip_tests=True).build() assert plan.directly_modified == {snapshot_a_id, snapshot_b_id} - assert {i.snapshot_id for i in plan.missing_intervals} == {snapshot_a_id, snapshot_b_id} - - assert plan.snapshots[snapshot_a_id].change_category == SnapshotChangeCategory.NON_BREAKING - assert plan.snapshots[snapshot_b_id].change_category == SnapshotChangeCategory.NON_BREAKING + assert {i.snapshot_id for i in plan.missing_intervals} == { + snapshot_a_id, + snapshot_b_id, + } + + assert ( + plan.snapshots[snapshot_a_id].change_category + == SnapshotChangeCategory.NON_BREAKING + ) + assert ( + plan.snapshots[snapshot_b_id].change_category + == SnapshotChangeCategory.NON_BREAKING + ) @time_machine.travel("2023-01-08 15:00:00 UTC") @@ -46,7 +51,9 @@ def test_dbt_is_incremental_table_is_missing(sushi_test_dbt_context: Context): context = sushi_test_dbt_context model = context.get_model("sushi.waiter_revenue_by_day_v2") - model = model.copy(update={"kind": IncrementalUnmanagedKind(), "start": "2023-01-01"}) + model = model.copy( + update={"kind": IncrementalUnmanagedKind(), "start": "2023-01-01"} + ) context.upsert_model(model) context._standalone_audits["sushi.test_top_waiters"].start = "2023-01-01" @@ -105,7 +112,8 @@ def test_dbt_requirements(sushi_dbt_context: Context): @time_machine.travel("2023-01-08 15:00:00 UTC") def test_dbt_dialect_with_normalization_strategy(init_and_plan_context: t.Callable): context, _ = init_and_plan_context( - "tests/fixtures/dbt/sushi_test", config="test_config_with_normalization_strategy" + "tests/fixtures/dbt/sushi_test", + config="test_config_with_normalization_strategy", ) assert context.default_dialect == "duckdb,normalization_strategy=LOWERCASE" @@ -113,11 +121,14 @@ def test_dbt_dialect_with_normalization_strategy(init_and_plan_context: t.Callab @time_machine.travel("2023-01-08 15:00:00 UTC") def test_dbt_before_all_with_var_ref_source(init_and_plan_context: t.Callable): _, plan = init_and_plan_context( - "tests/fixtures/dbt/sushi_test", config="test_config_with_normalization_strategy" + "tests/fixtures/dbt/sushi_test", + config="test_config_with_normalization_strategy", ) environment_statements = plan.to_evaluatable().environment_statements assert environment_statements - rendered_statements = [e.render_before_all(dialect="duckdb") for e in environment_statements] + rendered_statements = [ + e.render_before_all(dialect="duckdb") for e in environment_statements + ] assert rendered_statements[0] == [ "CREATE TABLE IF NOT EXISTS analytic_stats (physical_table TEXT, evaluation_time TEXT)", "CREATE TABLE IF NOT EXISTS to_be_executed_last (col TEXT)", diff --git a/tests/core/integration/test_dev_only_vde.py b/tests/core/integration/test_dev_only_vde.py index 611e207771..6f1fa9b52b 100644 --- a/tests/core/integration/test_dev_only_vde.py +++ b/tests/core/integration/test_dev_only_vde.py @@ -1,23 +1,17 @@ from __future__ import annotations import typing as t + import pytest -from sqlmesh.core.model.common import ParsableSql import time_machine from sqlmesh.core import dialect as d from sqlmesh.core.config.common import VirtualEnvironmentMode -from sqlmesh.core.model import ( - FullKind, - IncrementalUnmanagedKind, - SqlModel, - ViewKind, - load_sql_based_model, -) +from sqlmesh.core.model import (FullKind, IncrementalUnmanagedKind, SqlModel, + ViewKind, load_sql_based_model) +from sqlmesh.core.model.common import ParsableSql from sqlmesh.core.plan import SnapshotIntervals -from sqlmesh.core.snapshot import ( - SnapshotChangeCategory, -) +from sqlmesh.core.snapshot import SnapshotChangeCategory from sqlmesh.utils.date import to_date, to_timestamp from tests.core.integration.utils import add_projection_to_model @@ -44,7 +38,9 @@ def test_virtual_environment_mode_dev_only(init_and_plan_context: t.Callable): model = original_model.copy( update={ "query_": ParsableSql( - sql=original_model.query.order_by("waiter_id").sql(dialect=original_model.dialect) + sql=original_model.query.order_by("waiter_id").sql( + dialect=original_model.dialect + ) ) } ) @@ -64,7 +60,9 @@ def test_virtual_environment_mode_dev_only(init_and_plan_context: t.Callable): intervals=[(to_timestamp("2023-01-07"), to_timestamp("2023-01-08"))], ), ] - assert plan_dev.context_diff.snapshots[context.get_snapshot(model.name).snapshot_id].intervals + assert plan_dev.context_diff.snapshots[ + context.get_snapshot(model.name).snapshot_id + ].intervals assert plan_dev.context_diff.snapshots[ context.get_snapshot("sushi.top_waiters").snapshot_id ].intervals @@ -112,7 +110,9 @@ def test_virtual_environment_mode_dev_only(init_and_plan_context: t.Callable): @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_virtual_environment_mode_dev_only_model_kind_change(init_and_plan_context: t.Callable): +def test_virtual_environment_mode_dev_only_model_kind_change( + init_and_plan_context: t.Callable, +): context, plan = init_and_plan_context( "examples/sushi", config="test_config_virtual_environment_mode_dev_only" ) @@ -188,8 +188,7 @@ def test_virtual_environment_mode_dev_only_model_kind_change_incremental( ) forward_only_model_name = "memory.sushi.test_forward_only_model" - forward_only_model_expressions = d.parse( - f""" + forward_only_model_expressions = d.parse(f""" MODEL ( name {forward_only_model_name}, kind INCREMENTAL_BY_TIME_RANGE ( @@ -199,8 +198,7 @@ def test_virtual_environment_mode_dev_only_model_kind_change_incremental( ); SELECT '2023-01-01' AS ds, 'value' AS value; - """ - ) + """) forward_only_model = load_sql_based_model(forward_only_model_expressions) forward_only_model = forward_only_model.copy( update={"virtual_environment_mode": VirtualEnvironmentMode.DEV_ONLY} @@ -221,7 +219,9 @@ def test_virtual_environment_mode_dev_only_model_kind_change_incremental( context.get_snapshot(model.name).snapshot_id ].intervals context.apply(prod_plan) - data_objects = context.engine_adapter.get_data_objects("sushi", {"test_forward_only_model"}) + data_objects = context.engine_adapter.get_data_objects( + "sushi", {"test_forward_only_model"} + ) assert len(data_objects) == 1 assert data_objects[0].type == "view" @@ -234,7 +234,9 @@ def test_virtual_environment_mode_dev_only_model_kind_change_incremental( context.get_snapshot(model.name).snapshot_id ].intervals context.apply(prod_plan) - data_objects = context.engine_adapter.get_data_objects("sushi", {"test_forward_only_model"}) + data_objects = context.engine_adapter.get_data_objects( + "sushi", {"test_forward_only_model"} + ) assert len(data_objects) == 1 assert data_objects[0].type == "table" @@ -293,9 +295,13 @@ def test_virtual_environment_mode_dev_only_model_kind_change_manual_categorizati model = context.get_model("sushi.top_waiters") model = model.copy(update={"kind": FullKind()}) context.upsert_model(model) - dev_plan_builder = context.plan_builder("dev", skip_tests=True, no_auto_categorization=True) + dev_plan_builder = context.plan_builder( + "dev", skip_tests=True, no_auto_categorization=True + ) dev_plan_builder.set_choice( - dev_plan_builder._context_diff.snapshots[context.get_snapshot(model.name).snapshot_id], + dev_plan_builder._context_diff.snapshots[ + context.get_snapshot(model.name).snapshot_id + ], SnapshotChangeCategory.NON_BREAKING, ) dev_plan = dev_plan_builder.build() @@ -342,9 +348,15 @@ def test_virtual_environment_mode_dev_only_seed_model_change( assert len(plan.missing_intervals) == 2 context.apply(plan) - actual_seed_df_in_dev = context.fetchdf("SELECT * FROM sushi__dev.waiter_names WHERE id = 123") - assert actual_seed_df_in_dev.to_dict("records") == [{"id": 123, "name": "New Test Name"}] - actual_seed_df_in_prod = context.fetchdf("SELECT * FROM sushi.waiter_names WHERE id = 123") + actual_seed_df_in_dev = context.fetchdf( + "SELECT * FROM sushi__dev.waiter_names WHERE id = 123" + ) + assert actual_seed_df_in_dev.to_dict("records") == [ + {"id": 123, "name": "New Test Name"} + ] + actual_seed_df_in_prod = context.fetchdf( + "SELECT * FROM sushi.waiter_names WHERE id = 123" + ) assert actual_seed_df_in_prod.empty plan = context.plan_builder("prod").build() @@ -353,8 +365,12 @@ def test_virtual_environment_mode_dev_only_seed_model_change( assert plan.missing_intervals[0].snapshot_id == seed_model_snapshot.snapshot_id context.apply(plan) - actual_seed_df_in_prod = context.fetchdf("SELECT * FROM sushi.waiter_names WHERE id = 123") - assert actual_seed_df_in_prod.to_dict("records") == [{"id": 123, "name": "New Test Name"}] + actual_seed_df_in_prod = context.fetchdf( + "SELECT * FROM sushi.waiter_names WHERE id = 123" + ) + assert actual_seed_df_in_prod.to_dict("records") == [ + {"id": 123, "name": "New Test Name"} + ] @time_machine.travel("2023-01-08 15:00:00 UTC") @@ -462,7 +478,9 @@ def test_virtual_environment_mode_dev_only_seed_model_change_schema( } context.upsert_model(SqlModel.parse_obj(downstream_model_kwargs)) - context.plan("dev", auto_apply=True, no_prompts=True, skip_tests=True, enable_preview=True) + context.plan( + "dev", auto_apply=True, no_prompts=True, skip_tests=True, enable_preview=True + ) assert ( context.engine_adapter.fetchone( @@ -474,4 +492,6 @@ def test_virtual_environment_mode_dev_only_seed_model_change_schema( # Deploy to prod context.clear_caches() context.plan("prod", auto_apply=True, no_prompts=True, skip_tests=True) - assert "new_column" in context.engine_adapter.columns("sushi.waiter_as_customer_by_day") + assert "new_column" in context.engine_adapter.columns( + "sushi.waiter_as_customer_by_day" + ) diff --git a/tests/core/integration/test_forward_only.py b/tests/core/integration/test_forward_only.py index 2dddf18efd..deb26d343d 100644 --- a/tests/core/integration/test_forward_only.py +++ b/tests/core/integration/test_forward_only.py @@ -1,23 +1,18 @@ from __future__ import annotations import typing as t + import numpy as np # noqa: TID253 import pandas as pd # noqa: TID253 import pytest import time_machine from sqlmesh.core import dialect as d -from sqlmesh.core.context import Context from sqlmesh.core.config.categorizer import CategorizerConfig -from sqlmesh.core.model import ( - FullKind, - SqlModel, - load_sql_based_model, -) +from sqlmesh.core.context import Context +from sqlmesh.core.model import FullKind, SqlModel, load_sql_based_model from sqlmesh.core.plan import SnapshotIntervals -from sqlmesh.core.snapshot import ( - SnapshotChangeCategory, -) +from sqlmesh.core.snapshot import SnapshotChangeCategory from sqlmesh.utils.date import to_datetime, to_timestamp from tests.core.integration.utils import add_projection_to_model @@ -33,9 +28,13 @@ def test_forward_only_plan_with_effective_date(context_fixture: Context, request context = request.getfixturevalue(context_fixture) model_name = "sushi.waiter_revenue_by_day" model = context.get_model(model_name) - context.upsert_model(add_projection_to_model(t.cast(SqlModel, model)), start="2023-01-01") + context.upsert_model( + add_projection_to_model(t.cast(SqlModel, model)), start="2023-01-01" + ) snapshot = context.get_snapshot(model, raise_if_missing=True) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan_builder = context.plan_builder("dev", skip_tests=True, forward_only=True) plan = plan_builder.build() @@ -159,7 +158,8 @@ def test_forward_only_plan_with_effective_date(context_fixture: Context, request "SELECT DISTINCT event_date FROM sushi.waiter_revenue_by_day WHERE one IS NOT NULL ORDER BY event_date" ) assert prod_df["event_date"].tolist() == [ - pd.to_datetime(x) for x in ["2023-01-04", "2023-01-05", "2023-01-06", "2023-01-07"] + pd.to_datetime(x) + for x in ["2023-01-04", "2023-01-05", "2023-01-06", "2023-01-07"] ] @@ -177,7 +177,9 @@ def test_forward_only_model_regular_plan(init_and_plan_context: t.Callable): context.upsert_model(model) snapshot = context.get_snapshot(model, raise_if_missing=True) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=False).build() assert len(plan.new_snapshots) == 2 @@ -273,7 +275,9 @@ def test_forward_only_model_regular_plan(init_and_plan_context: t.Callable): @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_forward_only_model_regular_plan_preview_enabled(init_and_plan_context: t.Callable): +def test_forward_only_model_regular_plan_preview_enabled( + init_and_plan_context: t.Callable, +): context, plan = init_and_plan_context("examples/sushi") context.apply(plan) @@ -286,7 +290,9 @@ def test_forward_only_model_regular_plan_preview_enabled(init_and_plan_context: context.upsert_model(model) snapshot = context.get_snapshot(model, raise_if_missing=True) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=True).build() assert len(plan.new_snapshots) == 2 @@ -320,12 +326,13 @@ def test_forward_only_model_regular_plan_preview_enabled(init_and_plan_context: @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_forward_only_model_restate_full_history_in_dev(init_and_plan_context: t.Callable): +def test_forward_only_model_restate_full_history_in_dev( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") model_name = "memory.sushi.customer_max_revenue" - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name {model_name}, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -338,8 +345,7 @@ def test_forward_only_model_restate_full_history_in_dev(init_and_plan_context: t customer_id, MAX(revenue) AS max_revenue FROM memory.sushi.customer_revenue_lifetime GROUP BY 1; - """ - ) + """) model = load_sql_based_model(expressions) assert model.forward_only @@ -379,7 +385,9 @@ def test_forward_only_model_restate_full_history_in_dev(init_and_plan_context: t assert df["cnt"][0] == 1 # Apply a restatement plan in dev - plan = context.plan("dev", restate_models=[model.name], auto_apply=True, enable_preview=False) + plan = context.plan( + "dev", restate_models=[model.name], auto_apply=True, enable_preview=False + ) assert len(plan.missing_intervals) == 1 # Check that the dummy value is not present @@ -429,11 +437,15 @@ def test_full_history_restatement_model_regular_plan_preview_enabled( == SnapshotChangeCategory.INDIRECT_NON_BREAKING ) assert ( - plan.context_diff.snapshots[active_customers_snapshot.snapshot_id].change_category + plan.context_diff.snapshots[ + active_customers_snapshot.snapshot_id + ].change_category == SnapshotChangeCategory.INDIRECT_NON_BREAKING ) assert ( - plan.context_diff.snapshots[waiter_as_customer_snapshot.snapshot_id].change_category + plan.context_diff.snapshots[ + waiter_as_customer_snapshot.snapshot_id + ].change_category == SnapshotChangeCategory.INDIRECT_NON_BREAKING ) assert all(s.is_forward_only for s in plan.new_snapshots) @@ -452,7 +464,9 @@ def test_full_history_restatement_model_regular_plan_preview_enabled( @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_metadata_changed_regular_plan_preview_enabled(init_and_plan_context: t.Callable): +def test_metadata_changed_regular_plan_preview_enabled( + init_and_plan_context: t.Callable, +): context, plan = init_and_plan_context("examples/sushi") context.apply(plan) @@ -463,7 +477,9 @@ def test_metadata_changed_regular_plan_preview_enabled(init_and_plan_context: t. context.upsert_model(model) snapshot = context.get_snapshot(model, raise_if_missing=True) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=True).build() assert len(plan.new_snapshots) == 2 @@ -480,13 +496,13 @@ def test_metadata_changed_regular_plan_preview_enabled(init_and_plan_context: t. @time_machine.travel("2023-01-08 00:00:00 UTC") -def test_forward_only_preview_child_that_runs_before_parent(init_and_plan_context: t.Callable): +def test_forward_only_preview_child_that_runs_before_parent( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") # This model runs at minute 30 of every hour - upstream_model = load_sql_based_model( - d.parse( - """ + upstream_model = load_sql_based_model(d.parse(""" MODEL ( name memory.sushi.upstream_model, kind FULL, @@ -495,15 +511,11 @@ def test_forward_only_preview_child_that_runs_before_parent(init_and_plan_contex ); SELECT 1 AS a; - """ - ) - ) + """)) context.upsert_model(upstream_model) # This model runs at minute 0 of every hour, so it runs before the upstream model - downstream_model = load_sql_based_model( - d.parse( - """ + downstream_model = load_sql_based_model(d.parse(""" MODEL ( name memory.sushi.downstream_model, kind INCREMENTAL_BY_TIME_RANGE( @@ -515,9 +527,7 @@ def test_forward_only_preview_child_that_runs_before_parent(init_and_plan_contex ); SELECT a, '2023-01-06' AS event_date FROM memory.sushi.upstream_model; - """ - ) - ) + """)) context.upsert_model(downstream_model) context.plan("prod", skip_tests=True, auto_apply=True) @@ -529,7 +539,9 @@ def test_forward_only_preview_child_that_runs_before_parent(init_and_plan_contex # Now it's time for the upstream model to run but it hasn't run yet with time_machine.travel("2023-01-08 00:35:00 UTC"): # Make a change to the downstream model. - downstream_model = add_projection_to_model(t.cast(SqlModel, downstream_model), literal=True) + downstream_model = add_projection_to_model( + t.cast(SqlModel, downstream_model), literal=True + ) context.upsert_model(downstream_model) # The plan should only backfill the downstream model despite upstream missing intervals @@ -540,7 +552,10 @@ def test_forward_only_preview_child_that_runs_before_parent(init_and_plan_contex downstream_model.name, raise_if_missing=True ).snapshot_id, intervals=[ - (to_timestamp("2023-01-07 23:00:00"), to_timestamp("2023-01-08 00:00:00")) + ( + to_timestamp("2023-01-07 23:00:00"), + to_timestamp("2023-01-08 00:00:00"), + ) ], ), ] @@ -565,7 +580,9 @@ def test_forward_only_monthly_model(init_and_plan_context: t.Callable): plan = context.plan_builder("prod", skip_tests=True).build() context.apply(plan) - waiter_revenue_by_day_snapshot = context.get_snapshot(model.name, raise_if_missing=True) + waiter_revenue_by_day_snapshot = context.get_snapshot( + model.name, raise_if_missing=True + ) assert waiter_revenue_by_day_snapshot.intervals == [ (to_timestamp("2022-01-01"), to_timestamp("2023-01-01")) ] @@ -601,7 +618,9 @@ def test_forward_only_parent_created_in_dev_child_created_in_prod( waiter_revenue_by_day_model = add_projection_to_model( t.cast(SqlModel, waiter_revenue_by_day_model) ) - forward_only_kind = waiter_revenue_by_day_model.kind.copy(update={"forward_only": True}) + forward_only_kind = waiter_revenue_by_day_model.kind.copy( + update={"forward_only": True} + ) waiter_revenue_by_day_model = waiter_revenue_by_day_model.copy( update={"kind": forward_only_kind} ) @@ -610,12 +629,16 @@ def test_forward_only_parent_created_in_dev_child_created_in_prod( waiter_revenue_by_day_snapshot = context.get_snapshot( waiter_revenue_by_day_model, raise_if_missing=True ) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=False).build() assert len(plan.new_snapshots) == 2 assert ( - plan.context_diff.snapshots[waiter_revenue_by_day_snapshot.snapshot_id].change_category + plan.context_diff.snapshots[ + waiter_revenue_by_day_snapshot.snapshot_id + ].change_category == SnapshotChangeCategory.NON_BREAKING ) assert ( @@ -630,10 +653,14 @@ def test_forward_only_parent_created_in_dev_child_created_in_prod( # Update the child to refer to a newly added column. top_waiters_model = context.get_model("sushi.top_waiters") - top_waiters_model = add_projection_to_model(t.cast(SqlModel, top_waiters_model), literal=False) + top_waiters_model = add_projection_to_model( + t.cast(SqlModel, top_waiters_model), literal=False + ) context.upsert_model(top_waiters_model) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("prod", skip_tests=True, enable_preview=False).build() assert len(plan.new_snapshots) == 1 @@ -658,7 +685,9 @@ def test_forward_only_view_migration( context.upsert_model(model) # Apply a forward-only plan - context.plan("prod", skip_tests=True, no_prompts=True, auto_apply=True, forward_only=True) + context.plan( + "prod", skip_tests=True, no_prompts=True, auto_apply=True, forward_only=True + ) # Make sure that the new column got reflected in the view schema df = context.fetchdf("SELECT one FROM sushi.top_waiters LIMIT 1") @@ -669,7 +698,9 @@ def test_forward_only_view_migration( def test_new_forward_only_model(init_and_plan_context: t.Callable): context, _ = init_and_plan_context("examples/sushi") - context.plan("dev", skip_tests=True, no_prompts=True, auto_apply=True, enable_preview=False) + context.plan( + "dev", skip_tests=True, no_prompts=True, auto_apply=True, enable_preview=False + ) snapshot = context.get_snapshot("sushi.marketing") @@ -697,12 +728,16 @@ def test_non_breaking_change_after_forward_only_in_dev( waiter_revenue_by_day_snapshot = context.get_snapshot( "sushi.waiter_revenue_by_day", raise_if_missing=True ) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, forward_only=True).build() assert len(plan.new_snapshots) == 2 assert ( - plan.context_diff.snapshots[waiter_revenue_by_day_snapshot.snapshot_id].change_category + plan.context_diff.snapshots[ + waiter_revenue_by_day_snapshot.snapshot_id + ].change_category == SnapshotChangeCategory.NON_BREAKING ) assert ( @@ -729,8 +764,12 @@ def test_non_breaking_change_after_forward_only_in_dev( # Make a non-breaking change to a model downstream. model = context.get_model("sushi.top_waiters") # Select 'one' column from the updated upstream model. - context.upsert_model(add_projection_to_model(t.cast(SqlModel, model), literal=False)) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + context.upsert_model( + add_projection_to_model(t.cast(SqlModel, model), literal=False) + ) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True).build() assert len(plan.new_snapshots) == 1 @@ -797,7 +836,9 @@ def test_non_breaking_change_after_forward_only_in_dev( @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_context: t.Callable): +def test_indirect_non_breaking_change_after_forward_only_in_dev( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") # Make sure that the most downstream model is a materialized model. model = context.get_model("sushi.top_waiters") @@ -808,7 +849,9 @@ def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_co # Make sushi.orders a forward-only model. model = context.get_model("sushi.orders") updated_model_kind = model.kind.copy(update={"forward_only": True}) - model = model.copy(update={"stamp": "force new version", "kind": updated_model_kind}) + model = model.copy( + update={"stamp": "force new version", "kind": updated_model_kind} + ) context.upsert_model(model) snapshot = context.get_snapshot(model, raise_if_missing=True) @@ -829,7 +872,9 @@ def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_co # Make a non-breaking change to a model. model = context.get_model("sushi.top_waiters") context.upsert_model(add_projection_to_model(t.cast(SqlModel, model))) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=False).build() assert len(plan.new_snapshots) == 1 @@ -862,12 +907,16 @@ def test_indirect_non_breaking_change_after_forward_only_in_dev(init_and_plan_co waiter_revenue_by_day_snapshot = context.get_snapshot( "sushi.waiter_revenue_by_day", raise_if_missing=True ) - top_waiters_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + top_waiters_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder("dev", skip_tests=True, enable_preview=False).build() assert len(plan.new_snapshots) == 2 assert ( - plan.context_diff.snapshots[waiter_revenue_by_day_snapshot.snapshot_id].change_category + plan.context_diff.snapshots[ + waiter_revenue_by_day_snapshot.snapshot_id + ].change_category == SnapshotChangeCategory.NON_BREAKING ) assert ( @@ -946,7 +995,9 @@ def test_changes_downstream_of_indirect_non_breaking_snapshot_without_intervals( plan_builder = context.plan_builder( "dev", skip_backfill=True, skip_tests=True, no_auto_categorization=True ) - plan_builder.set_choice(context.get_snapshot(model), SnapshotChangeCategory.BREAKING) + plan_builder.set_choice( + context.get_snapshot(model), SnapshotChangeCategory.BREAKING + ) context.apply(plan_builder.build()) # Now make a non-breaking change to the same snapshot. @@ -955,12 +1006,16 @@ def test_changes_downstream_of_indirect_non_breaking_snapshot_without_intervals( plan_builder = context.plan_builder( "dev", skip_backfill=True, skip_tests=True, no_auto_categorization=True ) - plan_builder.set_choice(context.get_snapshot(model), SnapshotChangeCategory.NON_BREAKING) + plan_builder.set_choice( + context.get_snapshot(model), SnapshotChangeCategory.NON_BREAKING + ) context.apply(plan_builder.build()) # Now make a change to a model downstream of the above model. downstream_model = context.get_model("sushi.top_waiters") - downstream_model = downstream_model.copy(update={"stamp": "yet another new version"}) + downstream_model = downstream_model.copy( + update={"stamp": "yet another new version"} + ) context.upsert_model(downstream_model) plan = context.plan_builder("dev", skip_tests=True).build() @@ -969,11 +1024,15 @@ def test_changes_downstream_of_indirect_non_breaking_snapshot_without_intervals( assert not deployability_index.is_representative( context.get_snapshot("sushi.waiter_revenue_by_day") ) - assert not deployability_index.is_deployable(context.get_snapshot("sushi.top_waiters")) + assert not deployability_index.is_deployable( + context.get_snapshot("sushi.top_waiters") + ) @time_machine.travel("2023-01-08 15:00:00 UTC", tick=True) -def test_metadata_change_after_forward_only_results_in_migration(init_and_plan_context: t.Callable): +def test_metadata_change_after_forward_only_results_in_migration( + init_and_plan_context: t.Callable, +): context, plan = init_and_plan_context("examples/sushi") context.apply(plan) @@ -991,7 +1050,9 @@ def test_metadata_change_after_forward_only_results_in_migration(init_and_plan_c context.upsert_model(model) plan = context.plan("dev", skip_tests=True, auto_apply=True, no_prompts=True) assert len(plan.new_snapshots) == 2 - assert all(s.change_category == SnapshotChangeCategory.METADATA for s in plan.new_snapshots) + assert all( + s.change_category == SnapshotChangeCategory.METADATA for s in plan.new_snapshots + ) # Deploy the latest change to prod context.plan("prod", skip_tests=True, auto_apply=True, no_prompts=True) @@ -1002,7 +1063,9 @@ def test_metadata_change_after_forward_only_results_in_migration(init_and_plan_c @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_indirect_non_breaking_downstream_of_forward_only(init_and_plan_context: t.Callable): +def test_indirect_non_breaking_downstream_of_forward_only( + init_and_plan_context: t.Callable, +): context, plan = init_and_plan_context("examples/sushi") context.apply(plan) @@ -1013,13 +1076,19 @@ def test_indirect_non_breaking_downstream_of_forward_only(init_and_plan_context: update={"stamp": "force new version", "kind": updated_model_kind} ) context.upsert_model(forward_only_model) - forward_only_snapshot = context.get_snapshot(forward_only_model, raise_if_missing=True) + forward_only_snapshot = context.get_snapshot( + forward_only_model, raise_if_missing=True + ) non_breaking_model = context.get_model("sushi.waiter_revenue_by_day") non_breaking_model = non_breaking_model.copy(update={"start": "2023-01-01"}) context.upsert_model(add_projection_to_model(t.cast(SqlModel, non_breaking_model))) - non_breaking_snapshot = context.get_snapshot(non_breaking_model, raise_if_missing=True) - top_waiter_snapshot = context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + non_breaking_snapshot = context.get_snapshot( + non_breaking_model, raise_if_missing=True + ) + top_waiter_snapshot = context.get_snapshot( + "sushi.top_waiters", raise_if_missing=True + ) plan = context.plan_builder( "dev", @@ -1039,9 +1108,15 @@ def test_indirect_non_breaking_downstream_of_forward_only(init_and_plan_context: plan.context_diff.snapshots[top_waiter_snapshot.snapshot_id].change_category == SnapshotChangeCategory.INDIRECT_NON_BREAKING ) - assert plan.context_diff.snapshots[forward_only_snapshot.snapshot_id].is_forward_only - assert not plan.context_diff.snapshots[non_breaking_snapshot.snapshot_id].is_forward_only - assert not plan.context_diff.snapshots[top_waiter_snapshot.snapshot_id].is_forward_only + assert plan.context_diff.snapshots[ + forward_only_snapshot.snapshot_id + ].is_forward_only + assert not plan.context_diff.snapshots[ + non_breaking_snapshot.snapshot_id + ].is_forward_only + assert not plan.context_diff.snapshots[ + top_waiter_snapshot.snapshot_id + ].is_forward_only assert plan.start == to_timestamp("2023-01-01") assert plan.missing_intervals == [ @@ -1124,8 +1199,7 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot( # Forward-only parent forward_only_model_name = "memory.sushi.test_forward_only_model" - forward_only_model_expressions = d.parse( - f""" + forward_only_model_expressions = d.parse(f""" MODEL ( name {forward_only_model_name}, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1135,39 +1209,34 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot( ); SELECT '2023-01-01' AS ds, 'value' AS value; - """ - ) + """) forward_only_model = load_sql_based_model(forward_only_model_expressions) assert forward_only_model.forward_only context.upsert_model(forward_only_model) # FULL downstream model. full_downstream_model_name = "memory.sushi.test_full_downstream_model" - full_downstream_model_expressions = d.parse( - f""" + full_downstream_model_expressions = d.parse(f""" MODEL ( name {full_downstream_model_name}, kind FULL, ); SELECT ds, value FROM {forward_only_model_name}; - """ - ) + """) full_downstream_model = load_sql_based_model(full_downstream_model_expressions) context.upsert_model(full_downstream_model) # VIEW downstream of the previous FULL model. view_downstream_model_name = "memory.sushi.test_view_downstream_model" - view_downstream_model_expressions = d.parse( - f""" + view_downstream_model_expressions = d.parse(f""" MODEL ( name {view_downstream_model_name}, kind VIEW, ); SELECT ds, value FROM {full_downstream_model_name}; - """ - ) + """) view_downstream_model = load_sql_based_model(view_downstream_model_expressions) context.upsert_model(view_downstream_model) @@ -1176,10 +1245,18 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot( # Make a change to the forward-only model and apply it in dev. context.upsert_model(add_projection_to_model(t.cast(SqlModel, forward_only_model))) - forward_only_model_snapshot_id = context.get_snapshot(forward_only_model_name).snapshot_id - full_downstream_model_snapshot_id = context.get_snapshot(full_downstream_model_name).snapshot_id - view_downstream_model_snapshot_id = context.get_snapshot(view_downstream_model_name).snapshot_id - dev_plan = context.plan("dev", auto_apply=True, no_prompts=True, enable_preview=False) + forward_only_model_snapshot_id = context.get_snapshot( + forward_only_model_name + ).snapshot_id + full_downstream_model_snapshot_id = context.get_snapshot( + full_downstream_model_name + ).snapshot_id + view_downstream_model_snapshot_id = context.get_snapshot( + view_downstream_model_name + ).snapshot_id + dev_plan = context.plan( + "dev", auto_apply=True, no_prompts=True, enable_preview=False + ) assert ( dev_plan.snapshots[forward_only_model_snapshot_id].change_category == SnapshotChangeCategory.NON_BREAKING @@ -1195,20 +1272,24 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot( assert not dev_plan.missing_intervals # Make a follow-up breaking change to the downstream full model. - new_full_downstream_model_expressions = d.parse( - f""" + new_full_downstream_model_expressions = d.parse(f""" MODEL ( name {full_downstream_model_name}, kind FULL, ); SELECT ds, 'new_value' AS value FROM {forward_only_model_name}; - """ + """) + new_full_downstream_model = load_sql_based_model( + new_full_downstream_model_expressions ) - new_full_downstream_model = load_sql_based_model(new_full_downstream_model_expressions) context.upsert_model(new_full_downstream_model) - full_downstream_model_snapshot_id = context.get_snapshot(full_downstream_model_name).snapshot_id - view_downstream_model_snapshot_id = context.get_snapshot(view_downstream_model_name).snapshot_id + full_downstream_model_snapshot_id = context.get_snapshot( + full_downstream_model_name + ).snapshot_id + view_downstream_model_snapshot_id = context.get_snapshot( + view_downstream_model_name + ).snapshot_id dev_plan = context.plan( "dev", categorizer_config=CategorizerConfig.all_full(), @@ -1225,8 +1306,12 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot( == SnapshotChangeCategory.INDIRECT_BREAKING ) assert len(dev_plan.missing_intervals) == 2 - assert dev_plan.missing_intervals[0].snapshot_id == full_downstream_model_snapshot_id - assert dev_plan.missing_intervals[1].snapshot_id == view_downstream_model_snapshot_id + assert ( + dev_plan.missing_intervals[0].snapshot_id == full_downstream_model_snapshot_id + ) + assert ( + dev_plan.missing_intervals[1].snapshot_id == view_downstream_model_snapshot_id + ) # Check that the representative view hasn't been created yet. assert not context.engine_adapter.table_exists( @@ -1235,12 +1320,22 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot( # Now promote the very first change to prod without promoting the 2nd breaking change. context.upsert_model(full_downstream_model) - context.plan(auto_apply=True, no_prompts=True, categorizer_config=CategorizerConfig.all_full()) + context.plan( + auto_apply=True, + no_prompts=True, + categorizer_config=CategorizerConfig.all_full(), + ) # Finally, make a non-breaking change to the full model in the same dev environment. - context.upsert_model(add_projection_to_model(t.cast(SqlModel, new_full_downstream_model))) - full_downstream_model_snapshot_id = context.get_snapshot(full_downstream_model_name).snapshot_id - view_downstream_model_snapshot_id = context.get_snapshot(view_downstream_model_name).snapshot_id + context.upsert_model( + add_projection_to_model(t.cast(SqlModel, new_full_downstream_model)) + ) + full_downstream_model_snapshot_id = context.get_snapshot( + full_downstream_model_name + ).snapshot_id + view_downstream_model_snapshot_id = context.get_snapshot( + view_downstream_model_name + ).snapshot_id dev_plan = context.plan( "dev", categorizer_config=CategorizerConfig.all_full(), @@ -1272,8 +1367,7 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot_migration( ): context, _ = init_and_plan_context("examples/sushi") - forward_only_model_expr = d.parse( - """ + forward_only_model_expr = d.parse(""" MODEL ( name memory.sushi.forward_only_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1284,34 +1378,29 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot_migration( ); SELECT '2023-01-07' AS ds, 1 AS a; - """ - ) + """) forward_only_model = load_sql_based_model(forward_only_model_expr) context.upsert_model(forward_only_model) - downstream_view_a_expr = d.parse( - """ + downstream_view_a_expr = d.parse(""" MODEL ( name memory.sushi.downstream_view_a, kind VIEW, ); SELECT a from memory.sushi.forward_only_model; - """ - ) + """) downstream_view_a = load_sql_based_model(downstream_view_a_expr) context.upsert_model(downstream_view_a) - downstream_view_b_expr = d.parse( - """ + downstream_view_b_expr = d.parse(""" MODEL ( name memory.sushi.downstream_view_b, kind VIEW, ); SELECT a from memory.sushi.downstream_view_a; - """ - ) + """) downstream_view_b = load_sql_based_model(downstream_view_b_expr) context.upsert_model(downstream_view_b) @@ -1325,9 +1414,9 @@ def test_indirect_non_breaking_view_model_non_representative_snapshot_migration( context.plan(auto_apply=True, no_prompts=True, skip_tests=True) # Make sure the downstrean indirect non-breaking view is available in prod - count = context.engine_adapter.fetchone("SELECT COUNT(*) FROM memory.sushi.downstream_view_b")[ - 0 - ] + count = context.engine_adapter.fetchone( + "SELECT COUNT(*) FROM memory.sushi.downstream_view_b" + )[0] assert count > 0 @@ -1336,8 +1425,7 @@ def test_new_forward_only_model_concurrent_versions(init_and_plan_context: t.Cal context, plan = init_and_plan_context("examples/sushi") context.apply(plan) - new_model_expr = d.parse( - """ + new_model_expr = d.parse(""" MODEL ( name memory.sushi.new_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1348,8 +1436,7 @@ def test_new_forward_only_model_concurrent_versions(init_and_plan_context: t.Cal ); SELECT '2023-01-07' AS ds, 1 AS a; - """ - ) + """) new_model = load_sql_based_model(new_model_expr) # Add the first version of the model and apply it to dev_a. @@ -1364,8 +1451,7 @@ def test_new_forward_only_model_concurrent_versions(init_and_plan_context: t.Cal context.apply(plan_a) - new_model_alt_expr = d.parse( - """ + new_model_alt_expr = d.parse(""" MODEL ( name memory.sushi.new_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1376,8 +1462,7 @@ def test_new_forward_only_model_concurrent_versions(init_and_plan_context: t.Cal ); SELECT '2023-01-07' AS ds, 1 AS b; - """ - ) + """) new_model_alt = load_sql_based_model(new_model_alt_expr) # Add the second version of the model but don't apply it yet @@ -1436,8 +1521,7 @@ def test_new_forward_only_model_same_dev_environment(init_and_plan_context: t.Ca context, plan = init_and_plan_context("examples/sushi") context.apply(plan) - new_model_expr = d.parse( - """ + new_model_expr = d.parse(""" MODEL ( name memory.sushi.new_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1448,8 +1532,7 @@ def test_new_forward_only_model_same_dev_environment(init_and_plan_context: t.Ca ); SELECT '2023-01-07' AS ds, 1 AS a; - """ - ) + """) new_model = load_sql_based_model(new_model_expr) # Add the first version of the model and apply it to dev. @@ -1467,8 +1550,7 @@ def test_new_forward_only_model_same_dev_environment(init_and_plan_context: t.Ca df = context.fetchdf("SELECT * FROM memory.sushi__dev.new_model") assert df.to_dict() == {"ds": {0: "2023-01-07"}, "a": {0: 1}} - new_model_alt_expr = d.parse( - """ + new_model_alt_expr = d.parse(""" MODEL ( name memory.sushi.new_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1479,8 +1561,7 @@ def test_new_forward_only_model_same_dev_environment(init_and_plan_context: t.Ca ); SELECT '2023-01-07' AS ds, 1 AS b; - """ - ) + """) new_model_alt = load_sql_based_model(new_model_alt_expr) # Add the second version of the model and apply it to the same environment. @@ -1493,5 +1574,7 @@ def test_new_forward_only_model_same_dev_environment(init_and_plan_context: t.Ca context.apply(plan_b) - df = context.fetchdf("SELECT * FROM memory.sushi__dev.new_model").replace({np.nan: None}) + df = context.fetchdf("SELECT * FROM memory.sushi__dev.new_model").replace( + {np.nan: None} + ) assert df.to_dict() == {"ds": {0: "2023-01-07"}, "b": {0: 1}} diff --git a/tests/core/integration/test_model_kinds.py b/tests/core/integration/test_model_kinds.py index 108fd1cb01..7e33161c4f 100644 --- a/tests/core/integration/test_model_kinds.py +++ b/tests/core/integration/test_model_kinds.py @@ -3,36 +3,29 @@ import typing as t from collections import Counter from datetime import timedelta +from pathlib import Path from unittest import mock + import pandas as pd # noqa: TID253 import pytest -from pathlib import Path import time_machine from pytest_mock.plugin import MockerFixture from sqlglot import exp from sqlmesh import CustomMaterialization from sqlmesh.core import dialect as d -from sqlmesh.core.config import ( - Config, - ModelDefaultsConfig, - DuckDBConnectionConfig, - GatewayConfig, -) +from sqlmesh.core.config import (Config, DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig) +from sqlmesh.core.config.categorizer import CategorizerConfig from sqlmesh.core.console import Console from sqlmesh.core.context import Context -from sqlmesh.core.config.categorizer import CategorizerConfig -from sqlmesh.core.model import ( - Model, - SqlModel, - CustomKind, - load_sql_based_model, -) +from sqlmesh.core.model import (CustomKind, Model, SqlModel, + load_sql_based_model) from sqlmesh.core.plan import SnapshotIntervals +from sqlmesh.utils import CorrelationId from sqlmesh.utils.date import to_date, to_timestamp from sqlmesh.utils.pydantic import validate_string from tests.conftest import SushiDataValidator -from sqlmesh.utils import CorrelationId from tests.utils.test_filesystem import create_temp_file if t.TYPE_CHECKING: @@ -49,8 +42,7 @@ def test_incremental_by_partition(init_and_plan_context: t.Callable): source_name = "raw.test_incremental_by_partition" model_name = "memory.sushi.test_incremental_by_partition" - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name {model_name}, kind INCREMENTAL_BY_PARTITION (disable_restatement false), @@ -60,8 +52,7 @@ def test_incremental_by_partition(init_and_plan_context: t.Callable): ); SELECT key, value FROM {source_name}; - """ - ) + """) model = load_sql_based_model(expressions) context.upsert_model(model) @@ -209,7 +200,9 @@ def insert( def test_incremental_time_self_reference( - mocker: MockerFixture, sushi_context: Context, sushi_data_validator: SushiDataValidator + mocker: MockerFixture, + sushi_context: Context, + sushi_data_validator: SushiDataValidator, ): start_ts = to_timestamp("1 week ago") start_date, end_date = to_date("1 week ago"), to_date("yesterday") @@ -225,9 +218,14 @@ def test_incremental_time_self_reference( "SELECT MAX(event_date) FROM sushi.customer_revenue_lifetime" ) assert df.iloc[0, 0] == pd.to_datetime(end_date) - results = sushi_data_validator.validate("sushi.customer_revenue_lifetime", start_date, end_date) + results = sushi_data_validator.validate( + "sushi.customer_revenue_lifetime", start_date, end_date + ) plan = sushi_context.plan_builder( - restate_models=["sushi.customer_revenue_lifetime", "sushi.customer_revenue_by_day"], + restate_models=[ + "sushi.customer_revenue_lifetime", + "sushi.customer_revenue_by_day", + ], start=start_date, end="5 days ago", ).build() @@ -242,20 +240,47 @@ def test_incremental_time_self_reference( SnapshotIntervals( snapshot_id=revenue_lifeteime_snapshot.snapshot_id, intervals=[ - (to_timestamp(to_date("7 days ago")), to_timestamp(to_date("6 days ago"))), - (to_timestamp(to_date("6 days ago")), to_timestamp(to_date("5 days ago"))), - (to_timestamp(to_date("5 days ago")), to_timestamp(to_date("4 days ago"))), - (to_timestamp(to_date("4 days ago")), to_timestamp(to_date("3 days ago"))), - (to_timestamp(to_date("3 days ago")), to_timestamp(to_date("2 days ago"))), - (to_timestamp(to_date("2 days ago")), to_timestamp(to_date("1 days ago"))), - (to_timestamp(to_date("1 day ago")), to_timestamp(to_date("today"))), + ( + to_timestamp(to_date("7 days ago")), + to_timestamp(to_date("6 days ago")), + ), + ( + to_timestamp(to_date("6 days ago")), + to_timestamp(to_date("5 days ago")), + ), + ( + to_timestamp(to_date("5 days ago")), + to_timestamp(to_date("4 days ago")), + ), + ( + to_timestamp(to_date("4 days ago")), + to_timestamp(to_date("3 days ago")), + ), + ( + to_timestamp(to_date("3 days ago")), + to_timestamp(to_date("2 days ago")), + ), + ( + to_timestamp(to_date("2 days ago")), + to_timestamp(to_date("1 days ago")), + ), + ( + to_timestamp(to_date("1 day ago")), + to_timestamp(to_date("today")), + ), ], ), SnapshotIntervals( snapshot_id=revenue_by_day_snapshot.snapshot_id, intervals=[ - (to_timestamp(to_date("7 days ago")), to_timestamp(to_date("6 days ago"))), - (to_timestamp(to_date("6 days ago")), to_timestamp(to_date("5 days ago"))), + ( + to_timestamp(to_date("7 days ago")), + to_timestamp(to_date("6 days ago")), + ), + ( + to_timestamp(to_date("6 days ago")), + to_timestamp(to_date("5 days ago")), + ), ], ), ], @@ -268,8 +293,12 @@ def test_incremental_time_self_reference( ) # Validate that we made 7 calls to the customer_revenue_lifetime snapshot and 1 call to the customer_revenue_by_day snapshot assert num_batch_calls == { - sushi_context.get_snapshot("sushi.customer_revenue_lifetime", raise_if_missing=True): 7, - sushi_context.get_snapshot("sushi.customer_revenue_by_day", raise_if_missing=True): 1, + sushi_context.get_snapshot( + "sushi.customer_revenue_lifetime", raise_if_missing=True + ): 7, + sushi_context.get_snapshot( + "sushi.customer_revenue_by_day", raise_if_missing=True + ): 1, } # Validate that the results are the same as before the restate assert results == sushi_data_validator.validate( @@ -560,7 +589,9 @@ def test_incremental_by_time_model_ignore_additive_change(tmp_path: Path): (models_dir / "test_model.sql").write_text(initial_model) context = Context(paths=[tmp_path], config=config) - context.engine_adapter.execute("ALTER TABLE source_table ADD COLUMN new_column INT") + context.engine_adapter.execute( + "ALTER TABLE source_table ADD COLUMN new_column INT" + ) context.plan("prod", auto_apply=True, no_prompts=True) # Verify data loading continued to work @@ -1961,7 +1992,9 @@ def test_incremental_by_time_model_ignore_destructive_change_unit_test(tmp_path: with time_machine.travel("2023-01-10 00:00:00 UTC"): context = Context(paths=[tmp_path], config=config) - context.engine_adapter.execute("INSERT INTO source_table VALUES (2, NULL, 3, '2023-01-09')") + context.engine_adapter.execute( + "INSERT INTO source_table VALUES (2, NULL, 3, '2023-01-09')" + ) context.run() test_result = context.test() updated_df = context.fetchdf('SELECT * FROM "default"."test_model"') @@ -2127,7 +2160,9 @@ def test_incremental_by_time_model_ignore_additive_change_unit_test(tmp_path: Pa with time_machine.travel("2023-01-10 00:00:00 UTC"): context = Context(paths=[tmp_path], config=config) - context.engine_adapter.execute("INSERT INTO source_table VALUES (2, NULL, 3, '2023-01-09')") + context.engine_adapter.execute( + "INSERT INTO source_table VALUES (2, NULL, 3, '2023-01-09')" + ) context.run() test_result = context.test() updated_df = context.fetchdf('SELECT * FROM "default"."test_model"') @@ -2223,7 +2258,9 @@ def test_scd_type_2_full_restatement_no_start_date(init_and_plan_context: t.Call raw_products_v2_model = load_sql_based_model(raw_products_v2) context.upsert_model(raw_products_v2_model) context.plan( - auto_apply=True, no_prompts=True, categorizer_config=CategorizerConfig.all_full() + auto_apply=True, + no_prompts=True, + categorizer_config=CategorizerConfig.all_full(), ) context.run() @@ -2249,7 +2286,9 @@ def test_scd_type_2_full_restatement_no_start_date(init_and_plan_context: t.Call raw_products_v3_model = load_sql_based_model(raw_products_v3) context.upsert_model(raw_products_v3_model) context.plan( - auto_apply=True, no_prompts=True, categorizer_config=CategorizerConfig.all_full() + auto_apply=True, + no_prompts=True, + categorizer_config=CategorizerConfig.all_full(), ) context.run() data_after_second_change = context.engine_adapter.fetchdf(query) @@ -2322,7 +2361,9 @@ def _correlation_id_in_sqls(correlation_id: CorrelationId, mock_logger): f"MODEL (name test.a, kind FULL); SELECT {i} AS col", ) - with mock.patch("sqlmesh.core.engine_adapter.base.EngineAdapter._log_sql") as mock_logger: + with mock.patch( + "sqlmesh.core.engine_adapter.base.EngineAdapter._log_sql" + ) as mock_logger: ctx.load() plan = ctx.plan(auto_apply=True, no_prompts=True) @@ -2449,7 +2490,9 @@ def test_scd_type_2_regular_run_with_offset(init_and_plan_context: t.Callable): raw_employee_status_v2_model = load_sql_based_model(raw_employee_status_v2) context.upsert_model(raw_employee_status_v2_model) context.plan( - auto_apply=True, no_prompts=True, categorizer_config=CategorizerConfig.all_full() + auto_apply=True, + no_prompts=True, + categorizer_config=CategorizerConfig.all_full(), ) # The 7th hour of the day the run is kicked off for the SCD Type 2 model @@ -2486,7 +2529,9 @@ def test_scd_type_2_regular_run_with_offset(init_and_plan_context: t.Callable): raw_employee_status_v2_model = load_sql_based_model(raw_employee_status_v2) context.upsert_model(raw_employee_status_v2_model) context.plan( - auto_apply=True, no_prompts=True, categorizer_config=CategorizerConfig.all_full() + auto_apply=True, + no_prompts=True, + categorizer_config=CategorizerConfig.all_full(), ) # A day later the run is kicked off for the SCD Type 2 model again @@ -2576,9 +2621,9 @@ def test_seed_model_metadata_update_does_not_trigger_backfill(tmp_path: Path): assert plan.missing_intervals # prove data loaded - assert ctx.engine_adapter.fetchall("select id, name from memory.test.source_data") == [ - (1, "test") - ] + assert ctx.engine_adapter.fetchall( + "select id, name from memory.test.source_data" + ) == [(1, "test")] # prove no diff ctx.load() @@ -2644,7 +2689,9 @@ def test_seed_model_metadata_update_does_not_trigger_backfill(tmp_path: Path): assert plan.missing_intervals # prove backfilled data loaded - assert ctx.engine_adapter.fetchall("select id, name from memory.test.source_data") == [ + assert ctx.engine_adapter.fetchall( + "select id, name from memory.test.source_data" + ) == [ (1, "test"), (2, "updated"), ] @@ -2681,6 +2728,8 @@ def test_seed_model_promote_to_prod_after_dev( context.apply(plan) assert ( - context.engine_adapter.fetchone("SELECT COUNT(*) FROM sushi.waiter_names WHERE id = 10")[0] + context.engine_adapter.fetchone( + "SELECT COUNT(*) FROM sushi.waiter_names WHERE id = 10" + )[0] == 1 ) diff --git a/tests/core/integration/test_multi_repo.py b/tests/core/integration/test_multi_repo.py index 035c8bfda8..c03ab61b38 100644 --- a/tests/core/integration/test_multi_repo.py +++ b/tests/core/integration/test_multi_repo.py @@ -1,39 +1,34 @@ from __future__ import annotations -from unittest.mock import patch -from textwrap import dedent import os -import pytest from pathlib import Path -from sqlmesh.core.console import ( - get_console, -) -from sqlmesh.core.config.naming import NameInferenceConfig -from sqlmesh.core.model.common import ParsableSql -from sqlmesh.utils.concurrency import NodeExecutionFailedError +from textwrap import dedent +from unittest.mock import patch + +import pytest from sqlmesh.core import constants as c -from sqlmesh.core.config import ( - Config, - GatewayConfig, - ModelDefaultsConfig, - DuckDBConnectionConfig, -) +from sqlmesh.core.config import (Config, DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig) +from sqlmesh.core.config.naming import NameInferenceConfig from sqlmesh.core.console import get_console from sqlmesh.core.context import Context +from sqlmesh.core.model.common import ParsableSql from sqlmesh.utils import yaml +from sqlmesh.utils.concurrency import NodeExecutionFailedError from sqlmesh.utils.date import now from tests.conftest import DuckDBMetadata -from tests.utils.test_helpers import use_terminal_console from tests.core.integration.utils import validate_apply_basics - +from tests.utils.test_helpers import use_terminal_console pytestmark = pytest.mark.slow @use_terminal_console def test_multi(mocker): - context = Context(paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory") + context = Context( + paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory" + ) with patch.object(get_console(), "log_warning") as mock_logger: context.plan_builder(environment="dev") @@ -77,7 +72,9 @@ def test_multi(mocker): context.upsert_model( model.copy( update={ - "query_": ParsableSql(sql=model.query.select("'c' AS c").sql(dialect=model.dialect)) + "query_": ParsableSql( + sql=model.query.select("'c' AS c").sql(dialect=model.dialect) + ) } ) ) @@ -143,16 +140,22 @@ def make_config(project: str, default_gateway: str) -> Config: project=project, gateways={ "repo-one": GatewayConfig( - connection=DuckDBConnectionConfig(database=str(tmp_path / "repo_one.duckdb")), + connection=DuckDBConnectionConfig( + database=str(tmp_path / "repo_one.duckdb") + ), variables={"owner": "repo-one-variable"}, ), "repo-two": GatewayConfig( - connection=DuckDBConnectionConfig(database=str(tmp_path / "repo_two.duckdb")), + connection=DuckDBConnectionConfig( + database=str(tmp_path / "repo_two.duckdb") + ), variables={"owner": "repo-two-variable"}, ), }, default_gateway=default_gateway, - model_defaults=ModelDefaultsConfig(dialect="duckdb", gateway=default_gateway), + model_defaults=ModelDefaultsConfig( + dialect="duckdb", gateway=default_gateway + ), variables={"owner": "global-variable"}, ) @@ -172,17 +175,23 @@ def make_config(project: str, default_gateway: str) -> Config: assert repo_one_model.gateway == "repo-one" assert repo_one_model.catalog == "repo_one" assert repo_one_model.project == "repo_one" - assert context.render(repo_one_model.fqn).sql() == ("SELECT 'repo-one-variable' AS \"owner\"") + assert context.render(repo_one_model.fqn).sql() == ( + "SELECT 'repo-one-variable' AS \"owner\"" + ) assert repo_two_model.gateway == "repo-two" assert repo_two_model.catalog == "repo_two" assert repo_two_model.project == "repo_two" - assert context.render(repo_two_model.fqn).sql() == ("SELECT 'repo-two-variable' AS \"owner\"") + assert context.render(repo_two_model.fqn).sql() == ( + "SELECT 'repo-two-variable' AS \"owner\"" + ) assert override_model.gateway == "repo-two" assert override_model.catalog == "repo_two" assert override_model.project == "repo_one" - assert context.render(override_model.fqn).sql() == ("SELECT 'repo-two-variable' AS \"owner\"") + assert context.render(override_model.fqn).sql() == ( + "SELECT 'repo-two-variable' AS \"owner\"" + ) @use_terminal_console @@ -223,7 +232,9 @@ def test_multi_repo_single_project_environment_statements_update(copy_to_temp_pa # Plan with only repo_1, this should preserve repo_2's statements from state repo_1_plan = context_repo_1_only.plan_builder(environment="dev").build() context_repo_1_only.apply(repo_1_plan) - updated_statements = context_repo_1_only.state_reader.get_environment_statements("dev") + updated_statements = context_repo_1_only.state_reader.get_environment_statements( + "dev" + ) # Should still have statements from both projects assert len(updated_statements) == 2 @@ -235,14 +246,19 @@ def test_multi_repo_single_project_environment_statements_update(copy_to_temp_pa repo_1_updated = sorted_updated[0] assert repo_1_updated.project == "repo_1" assert len(repo_1_updated.before_all) == 2 - assert "CREATE TABLE IF NOT EXISTS before_1_modified" in repo_1_updated.before_all[1] + assert ( + "CREATE TABLE IF NOT EXISTS before_1_modified" in repo_1_updated.before_all[1] + ) # Verify repo_2 statements are preserved from state repo_2_preserved = sorted_updated[1] assert repo_2_preserved.project == "repo_2" assert len(repo_2_preserved.before_all) == 1 assert "CREATE TABLE IF NOT EXISTS before_2" in repo_2_preserved.before_all[0] - assert "CREATE TABLE IF NOT EXISTS after_2 AS select @dup()" in repo_2_preserved.after_all[0] + assert ( + "CREATE TABLE IF NOT EXISTS after_2 AS select @dup()" + in repo_2_preserved.after_all[0] + ) @use_terminal_console @@ -422,7 +438,9 @@ def test_multi_virtual_layer(copy_to_temp_path): 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*"): + with pytest.raises( + NodeExecutionFailedError, match=r"Execution failed for node SnapshotId*" + ): context.apply(plan) @@ -448,7 +466,9 @@ def test_multi_dbt(mocker): ] assert "store_schemas" in silver_statements.jinja_macros.root_macros analytics_table = context.fetchdf("select * from analytic_stats;") - assert sorted(analytics_table.columns) == sorted(["physical_table", "evaluation_time"]) + assert sorted(analytics_table.columns) == sorted( + ["physical_table", "evaluation_time"] + ) schema_table = context.fetchdf("select * from schema_table;") assert sorted(schema_table.all_schemas[0]) == sorted(["bronze", "silver"]) @@ -462,9 +482,15 @@ def test_multi_hybrid(mocker): assert len(plan.new_snapshots) == 5 assert context.dag.roots == {'"memory"."dbt_repo"."e"'} - assert context.dag.graph['"memory"."dbt_repo"."c"'] == {'"memory"."sqlmesh_repo"."b"'} - assert context.dag.graph['"memory"."sqlmesh_repo"."b"'] == {'"memory"."sqlmesh_repo"."a"'} - assert context.dag.graph['"memory"."sqlmesh_repo"."a"'] == {'"memory"."dbt_repo"."e"'} + assert context.dag.graph['"memory"."dbt_repo"."c"'] == { + '"memory"."sqlmesh_repo"."b"' + } + assert context.dag.graph['"memory"."sqlmesh_repo"."b"'] == { + '"memory"."sqlmesh_repo"."a"' + } + assert context.dag.graph['"memory"."sqlmesh_repo"."a"'] == { + '"memory"."dbt_repo"."e"' + } assert context.dag.downstream('"memory"."dbt_repo"."e"') == [ '"memory"."sqlmesh_repo"."a"', '"memory"."sqlmesh_repo"."b"', @@ -476,9 +502,7 @@ def test_multi_hybrid(mocker): dbt_model_c = context.get_model("dbt_repo.c") assert sqlmesh_model_a.project == "sqlmesh_repo" - sqlmesh_rendered = ( - 'SELECT "e"."col_a" AS "col_a", "e"."col_b" AS "col_b" FROM "memory"."dbt_repo"."e" AS "e"' - ) + sqlmesh_rendered = 'SELECT "e"."col_a" AS "col_a", "e"."col_b" AS "col_b" FROM "memory"."dbt_repo"."e" AS "e"' dbt_rendered = 'SELECT DISTINCT ROUND(CAST(("b"."col_a" / NULLIF(100, 0)) AS DECIMAL(16, 2)), 2) AS "rounded_col_a" FROM "memory"."sqlmesh_repo"."b" AS "b"' assert sqlmesh_model_a.render_query().sql() == sqlmesh_rendered assert dbt_model_c.render_query().sql() == dbt_rendered @@ -549,8 +573,7 @@ def test_multi_repo_local_model_overrides_prod_from_other_project(copy_to_temp_p assert prod_model_c.project == "repo_2" with open(f"{repo_1_path}/models/c.sql", "w") as f: - f.write( - dedent("""\ + f.write(dedent("""\ MODEL ( name silver.c, kind FULL @@ -558,8 +581,7 @@ def test_multi_repo_local_model_overrides_prod_from_other_project(copy_to_temp_p SELECT DISTINCT col_a, col_b FROM bronze.a - """) - ) + """)) # silver.c exists locally in repo 1 now AND in prod under repo_2 context_repo1 = Context( @@ -666,15 +688,17 @@ def test_multi_repo_create_external_models(copy_to_temp_path): if repo_2_external.exists(): contents = yaml.load(repo_2_external) external_names = [e["name"] for e in contents] - assert not any("bronze" in name and "a" in name for name in external_names), ( - f"bronze.a should not be in repo_2's external models, but found: {external_names}" - ) + assert not any( + "bronze" in name and "a" in name for name in external_names + ), f"bronze.a should not be in repo_2's external models, but found: {external_names}" # repo_1 has no external dependencies at all repo_1_external = Path(repo_1_path) / c.EXTERNAL_MODELS_YAML if repo_1_external.exists(): contents = yaml.load(repo_1_external) - assert len(contents) == 0, f"repo_1 should have no external models, got: {contents}" + assert ( + len(contents) == 0 + ), f"repo_1 should have no external models, got: {contents}" # Plan should still resolve all 5 models as internal after create_external_models context.load() diff --git a/tests/core/integration/test_plan_options.py b/tests/core/integration/test_plan_options.py index a50dc145cd..0ff9559506 100644 --- a/tests/core/integration/test_plan_options.py +++ b/tests/core/integration/test_plan_options.py @@ -1,31 +1,18 @@ from __future__ import annotations import typing as t + import pytest -from sqlmesh.core.console import ( - set_console, - get_console, - TerminalConsole, -) import time_machine from sqlmesh.core import dialect as d -from sqlmesh.core.console import get_console -from sqlmesh.core.model import ( - SqlModel, - load_sql_based_model, -) +from sqlmesh.core.console import TerminalConsole, get_console, set_console +from sqlmesh.core.model import SqlModel, load_sql_based_model from sqlmesh.core.plan import SnapshotIntervals -from sqlmesh.core.snapshot import ( - SnapshotChangeCategory, -) +from sqlmesh.core.snapshot import SnapshotChangeCategory from sqlmesh.utils.date import to_datetime, to_timestamp -from sqlmesh.utils.errors import ( - NoChangesPlanError, -) -from tests.core.integration.utils import ( - add_projection_to_model, -) +from sqlmesh.utils.errors import NoChangesPlanError +from tests.core.integration.utils import add_projection_to_model pytestmark = pytest.mark.slow @@ -44,7 +31,9 @@ def test_empty_backfill(init_and_plan_context: t.Callable): for model in context.models.values(): if model.is_seed or model.kind.is_symbolic: continue - row_num = context.engine_adapter.fetchone(f"SELECT COUNT(*) FROM {model.name}")[0] + row_num = context.engine_adapter.fetchone(f"SELECT COUNT(*) FROM {model.name}")[ + 0 + ] assert row_num == 0 plan = context.plan_builder("prod", skip_tests=True).build() @@ -64,9 +53,7 @@ def test_empty_backfill_new_model(init_and_plan_context: t.Callable): context, plan = init_and_plan_context("examples/sushi") context.apply(plan) - new_model = load_sql_based_model( - d.parse( - """ + new_model = load_sql_based_model(d.parse(""" MODEL ( name memory.sushi.new_model, kind FULL, @@ -75,9 +62,7 @@ def test_empty_backfill_new_model(init_and_plan_context: t.Callable): ); SELECT 1 AS one; - """ - ) - ) + """)) new_model_name = context.upsert_model(new_model).fqn with time_machine.travel("2023-01-09 00:00:00 UTC"): @@ -92,9 +77,9 @@ def test_empty_backfill_new_model(init_and_plan_context: t.Callable): for model in context.models.values(): if model.is_seed or model.kind.is_symbolic: continue - row_num = context.engine_adapter.fetchone(f"SELECT COUNT(*) FROM sushi__dev.new_model")[ - 0 - ] + row_num = context.engine_adapter.fetchone( + f"SELECT COUNT(*) FROM sushi__dev.new_model" + )[0] assert row_num == 0 plan = context.plan_builder("prod", skip_tests=True).build() @@ -125,7 +110,9 @@ def test_plan_explain(init_and_plan_context: t.Callable): ) context.upsert_model(waiter_revenue_by_day_model) - waiter_revenue_by_day_snapshot = context.get_snapshot(waiter_revenue_by_day_model.name) + waiter_revenue_by_day_snapshot = context.get_snapshot( + waiter_revenue_by_day_model.name + ) top_waiters_snapshot = context.get_snapshot("sushi.top_waiters") common_kwargs = dict(skip_tests=True, no_prompts=True, explain=True) @@ -137,7 +124,9 @@ def test_plan_explain(init_and_plan_context: t.Callable): context.plan("dev", **common_kwargs, forward_only=True, enable_preview=True) context.plan("prod", **common_kwargs) context.plan("prod", **common_kwargs, forward_only=True) - context.plan("prod", **common_kwargs, restate_models=[waiter_revenue_by_day_model.name]) + context.plan( + "prod", **common_kwargs, restate_models=[waiter_revenue_by_day_model.name] + ) set_console(old_console) @@ -160,8 +149,7 @@ def test_plan_ignore_cron( ): context, _ = init_and_plan_context("examples/sushi") - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name memory.sushi.test_allow_partials, kind INCREMENTAL_UNMANAGED, @@ -170,17 +158,16 @@ def test_plan_ignore_cron( ); SELECT @end_ts AS end_ts - """ - ) + """) model = load_sql_based_model(expressions) context.upsert_model(model) context.plan("prod", skip_tests=True, auto_apply=True, no_prompts=True) assert ( - context.engine_adapter.fetchone("SELECT MAX(end_ts) FROM memory.sushi.test_allow_partials")[ - 0 - ] + context.engine_adapter.fetchone( + "SELECT MAX(end_ts) FROM memory.sushi.test_allow_partials" + )[0] == "2023-01-07 23:59:59.999999" ) @@ -189,7 +176,9 @@ def test_plan_ignore_cron( ).build() assert not plan_no_ignore_cron.missing_intervals - plan = context.plan_builder("prod", run=True, ignore_cron=True, skip_tests=True).build() + plan = context.plan_builder( + "prod", run=True, ignore_cron=True, skip_tests=True + ).build() assert plan.missing_intervals == [ SnapshotIntervals( snapshot_id=context.get_snapshot(model, raise_if_missing=True).snapshot_id, @@ -201,9 +190,9 @@ def test_plan_ignore_cron( context.apply(plan) assert ( - context.engine_adapter.fetchone("SELECT MAX(end_ts) FROM memory.sushi.test_allow_partials")[ - 0 - ] + context.engine_adapter.fetchone( + "SELECT MAX(end_ts) FROM memory.sushi.test_allow_partials" + )[0] == "2023-01-08 14:59:59.999999" ) @@ -225,8 +214,12 @@ def test_plan_with_run( context.apply(plan) - snapshots = context.state_sync.state_sync.get_snapshots(context.snapshots.values()) - assert {s.name: s.intervals[0][1] for s in snapshots.values() if s.intervals} == { + snapshots = context.state_sync.state_sync.get_snapshots( + context.snapshots.values() + ) + assert { + s.name: s.intervals[0][1] for s in snapshots.values() if s.intervals + } == { '"memory"."sushi"."waiter_revenue_by_day"': to_timestamp("2023-01-09"), '"memory"."sushi"."order_items"': to_timestamp("2023-01-09"), '"memory"."sushi"."orders"': to_timestamp("2023-01-09"), @@ -340,7 +333,9 @@ def test_select_models_for_backfill(init_and_plan_context: t.Callable): assert plan.missing_intervals == [ SnapshotIntervals( - snapshot_id=context.get_snapshot("sushi.items", raise_if_missing=True).snapshot_id, + snapshot_id=context.get_snapshot( + "sushi.items", raise_if_missing=True + ).snapshot_id, intervals=expected_intervals, ), SnapshotIntervals( @@ -350,7 +345,9 @@ def test_select_models_for_backfill(init_and_plan_context: t.Callable): intervals=expected_intervals, ), SnapshotIntervals( - snapshot_id=context.get_snapshot("sushi.orders", raise_if_missing=True).snapshot_id, + snapshot_id=context.get_snapshot( + "sushi.orders", raise_if_missing=True + ).snapshot_id, intervals=expected_intervals, ), SnapshotIntervals( @@ -443,7 +440,9 @@ def test_select_unchanged_model_for_backfill(init_and_plan_context: t.Callable): assert {o.name for o in schema_objects} == {"waiter_revenue_by_day"} # Now select a model downstream from the previously modified one in order to backfill it. - plan = context.plan_builder("dev", select_models=["*top_waiters"], skip_tests=True).build() + plan = context.plan_builder( + "dev", select_models=["*top_waiters"], skip_tests=True + ).build() assert not plan.has_changes assert plan.missing_intervals == [ diff --git a/tests/core/integration/test_restatement.py b/tests/core/integration/test_restatement.py index 3694efce31..ce7fb21f26 100644 --- a/tests/core/integration/test_restatement.py +++ b/tests/core/integration/test_restatement.py @@ -1,44 +1,29 @@ from __future__ import annotations +import queue +import re +import time import typing as t +from concurrent.futures import ThreadPoolExecutor, TimeoutError +from pathlib import Path + import pandas as pd # noqa: TID253 import pytest -from pathlib import Path -from sqlmesh.core.console import ( - MarkdownConsole, - set_console, - get_console, - CaptureTerminalConsole, -) import time_machine from sqlglot import exp -import re -from concurrent.futures import ThreadPoolExecutor, TimeoutError -import time -import queue from sqlmesh.core import constants as c -from sqlmesh.core.config import ( - Config, - GatewayConfig, - ModelDefaultsConfig, - DuckDBConnectionConfig, -) +from sqlmesh.core.config import (Config, DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig) +from sqlmesh.core.console import (CaptureTerminalConsole, MarkdownConsole, + get_console, set_console) from sqlmesh.core.context import Context -from sqlmesh.core.model import ( - IncrementalByTimeRangeKind, - IncrementalUnmanagedKind, - SqlModel, -) +from sqlmesh.core.model import (IncrementalByTimeRangeKind, + IncrementalUnmanagedKind, SqlModel) from sqlmesh.core.plan import SnapshotIntervals -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotId, -) +from sqlmesh.core.snapshot import Snapshot, SnapshotId from sqlmesh.utils.date import to_timestamp -from sqlmesh.utils.errors import ( - ConflictingPlanError, -) +from sqlmesh.utils.errors import ConflictingPlanError from tests.core.integration.utils import add_projection_to_model pytestmark = pytest.mark.slow @@ -63,7 +48,10 @@ def test_restatement_plan_ignores_changes(init_and_plan_context: t.Callable): assert not plan.new_snapshots assert plan.requires_backfill assert plan.restatements == { - restated_snapshot.snapshot_id: (to_timestamp("2023-01-01"), to_timestamp("2023-01-09")) + restated_snapshot.snapshot_id: ( + to_timestamp("2023-01-01"), + to_timestamp("2023-01-09"), + ) } assert plan.missing_intervals == [ SnapshotIntervals( @@ -135,7 +123,9 @@ def test_restatement_plan_across_environments_snapshot_with_shared_version( assert not plan.missing_intervals -def test_restatement_plan_hourly_with_downstream_daily_restates_correct_intervals(tmp_path: Path): +def test_restatement_plan_hourly_with_downstream_daily_restates_correct_intervals( + tmp_path: Path, +): model_a = """ MODEL ( name test.a, @@ -189,9 +179,13 @@ def test_restatement_plan_hourly_with_downstream_daily_restates_correct_interval "ts": exp.DataType.build("timestamp"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # plan + apply @@ -199,7 +193,8 @@ def test_restatement_plan_hourly_with_downstream_daily_restates_correct_interval def _dates_in_table(table_name: str) -> t.List[str]: return [ - str(r[0]) for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") + str(r[0]) + for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") ] # verify initial state @@ -212,7 +207,9 @@ def _dates_in_table(table_name: str) -> t.List[str]: ] # restate A - engine_adapter.execute("delete from test.external_table where ts = '2024-01-01 01:30:00'") + engine_adapter.execute( + "delete from test.external_table where ts = '2024-01-01 01:30:00'" + ) ctx.plan( restate_models=["test.a"], start="2024-01-01 01:00:00", @@ -242,7 +239,9 @@ def _dates_in_table(table_name: str) -> t.List[str]: } ) engine_adapter.replace_query( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # Restate A across a day boundary with the expectation that two day intervals in B are affected @@ -321,9 +320,13 @@ def test_restatement_plan_respects_disable_restatements(tmp_path: Path): "ts": exp.DataType.build("timestamp"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # plan + apply @@ -331,7 +334,8 @@ def test_restatement_plan_respects_disable_restatements(tmp_path: Path): def _dates_in_table(table_name: str) -> t.List[str]: return [ - str(r[0]) for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") + str(r[0]) + for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") ] def get_snapshot_intervals(snapshot_id): @@ -347,8 +351,12 @@ def get_snapshot_intervals(snapshot_id): ] # restate A and expect b to be ignored - starting_b_intervals = get_snapshot_intervals(ctx.snapshots['"memory"."test"."b"'].snapshot_id) - engine_adapter.execute("delete from test.external_table where ts = '2024-01-01 01:30:00'") + starting_b_intervals = get_snapshot_intervals( + ctx.snapshots['"memory"."test"."b"'].snapshot_id + ) + engine_adapter.execute( + "delete from test.external_table where ts = '2024-01-01 01:30:00'" + ) ctx.plan( restate_models=["test.a"], start="2024-01-01", @@ -371,7 +379,9 @@ def get_snapshot_intervals(snapshot_id): ] # Verify B intervals were not touched - b_intervals = get_snapshot_intervals(ctx.snapshots['"memory"."test"."b"'].snapshot_id) + b_intervals = get_snapshot_intervals( + ctx.snapshots['"memory"."test"."b"'].snapshot_id + ) assert starting_b_intervals == b_intervals @@ -418,7 +428,13 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: { "account_id": [1001, 1002, 1003, 1004, 1005], "name": ["foo", "bar", "baz", "bing", "bong"], - "date": ["2024-01-01", "2024-01-02", "2024-01-03", "2024-01-04", "2024-01-05"], + "date": [ + "2024-01-01", + "2024-01-02", + "2024-01-03", + "2024-01-04", + "2024-01-05", + ], } ) columns_to_types = { @@ -427,15 +443,23 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: "date": exp.DataType.build("date"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # first, create the prod models ctx.plan(auto_apply=True, no_prompts=True) - assert engine_adapter.fetchone("select count(*) from test.incremental_model") == (5,) - assert engine_adapter.fetchone("select count(*) from test.downstream_of_incremental") == (5,) + assert engine_adapter.fetchone("select count(*) from test.incremental_model") == ( + 5, + ) + assert engine_adapter.fetchone( + "select count(*) from test.downstream_of_incremental" + ) == (5,) assert not engine_adapter.table_exists("test__dev.incremental_model") # then, make a dev version @@ -457,7 +481,9 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: ctx.plan(environment="dev", auto_apply=True, no_prompts=True) assert engine_adapter.table_exists("test__dev.incremental_model") - assert engine_adapter.fetchone("select count(*) from test__dev.incremental_model") == (5,) + assert engine_adapter.fetchone( + "select count(*) from test__dev.incremental_model" + ) == (5,) # drop some source data so when we restate the interval it essentially clears it which is easy to verify engine_adapter.execute("delete from test.external_table where date = '2024-01-01'") @@ -472,18 +498,24 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: auto_apply=True, no_prompts=True, ) - assert engine_adapter.fetchone("select count(*) from test.incremental_model") == (5,) + assert engine_adapter.fetchone("select count(*) from test.incremental_model") == ( + 5, + ) assert engine_adapter.fetchone( "select count(*) from test.incremental_model where date = '2024-01-01'" ) == (1,) - assert engine_adapter.fetchone("select count(*) from test__dev.incremental_model") == (4,) + assert engine_adapter.fetchone( + "select count(*) from test__dev.incremental_model" + ) == (4,) assert engine_adapter.fetchone( "select count(*) from test__dev.incremental_model where date = '2024-01-01'" ) == (0,) # prod still should not be affected by a run because the restatement only happened in dev ctx.run() - assert engine_adapter.fetchone("select count(*) from test.incremental_model") == (5,) + assert engine_adapter.fetchone("select count(*) from test.incremental_model") == ( + 5, + ) assert engine_adapter.fetchone( "select count(*) from test.incremental_model where date = '2024-01-01'" ) == (1,) @@ -499,7 +531,9 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: auto_apply=True, no_prompts=True, ) - assert engine_adapter.fetchone("select count(*) from test.incremental_model") == (3,) + assert engine_adapter.fetchone("select count(*) from test.incremental_model") == ( + 3, + ) assert engine_adapter.fetchone( "select count(*) from test.incremental_model where date = '2024-01-01'" ) == (0,) @@ -511,7 +545,9 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: ) == (1,) # dev not affected yet until `sqlmesh run` is run - assert engine_adapter.fetchone("select count(*) from test__dev.incremental_model") == (4,) + assert engine_adapter.fetchone( + "select count(*) from test__dev.incremental_model" + ) == (4,) assert engine_adapter.fetchone( "select count(*) from test__dev.incremental_model where date = '2024-01-01'" ) == (0,) @@ -524,7 +560,9 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: # the restatement plan for prod should have cleared dev intervals too, which means this `sqlmesh run` re-runs 2024-01-01 and 2024-01-02 ctx.run(environment="dev") - assert engine_adapter.fetchone("select count(*) from test__dev.incremental_model") == (3,) + assert engine_adapter.fetchone( + "select count(*) from test__dev.incremental_model" + ) == (3,) assert engine_adapter.fetchone( "select count(*) from test__dev.incremental_model where date = '2024-01-01'" ) == (0,) @@ -536,13 +574,17 @@ def test_restatement_plan_clears_correct_intervals_across_environments(tmp_path: ) == (1,) # the downstream full model should always reflect whatever the incremental model is showing - assert engine_adapter.fetchone("select count(*) from test.downstream_of_incremental") == (3,) - assert engine_adapter.fetchone("select count(*) from test__dev.downstream_of_incremental") == ( - 3, - ) + assert engine_adapter.fetchone( + "select count(*) from test.downstream_of_incremental" + ) == (3,) + assert engine_adapter.fetchone( + "select count(*) from test__dev.downstream_of_incremental" + ) == (3,) -def test_prod_restatement_plan_clears_correct_intervals_in_derived_dev_tables(tmp_path: Path): +def test_prod_restatement_plan_clears_correct_intervals_in_derived_dev_tables( + tmp_path: Path, +): """ Scenario: I have models A[hourly] <- B[daily] <- C in prod @@ -622,9 +664,13 @@ def _derived_incremental_model_def(name: str, upstream: str) -> str: "ts": exp.DataType.build("timestamp"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # plan + apply A, B, C in prod @@ -647,7 +693,8 @@ def _derived_incremental_model_def(name: str, upstream: str) -> str: def _dates_in_table(table_name: str) -> t.List[str]: return [ - str(r[0]) for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") + str(r[0]) + for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") ] # verify initial state @@ -664,7 +711,9 @@ def _dates_in_table(table_name: str) -> t.List[str]: assert not engine_adapter.table_exists(tbl) # restate A in prod - engine_adapter.execute("delete from test.external_table where ts = '2024-01-01 01:30:00'") + engine_adapter.execute( + "delete from test.external_table where ts = '2024-01-01 01:30:00'" + ) ctx.plan( restate_models=["test.a"], start="2024-01-01 01:00:00", @@ -703,7 +752,9 @@ def _dates_in_table(table_name: str) -> t.List[str]: ], f"Table {tbl} wasnt cleared" -def test_prod_restatement_plan_clears_unaligned_intervals_in_derived_dev_tables(tmp_path: Path): +def test_prod_restatement_plan_clears_unaligned_intervals_in_derived_dev_tables( + tmp_path: Path, +): """ Scenario: I have a model A[hourly] in prod @@ -772,9 +823,13 @@ def test_prod_restatement_plan_clears_unaligned_intervals_in_derived_dev_tables( "ts": exp.DataType.build("timestamp"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # plan + apply A[hourly] in prod @@ -790,7 +845,8 @@ def test_prod_restatement_plan_clears_unaligned_intervals_in_derived_dev_tables( def _dates_in_table(table_name: str) -> t.List[str]: return [ - str(r[0]) for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") + str(r[0]) + for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") ] # verify initial state @@ -803,7 +859,9 @@ def _dates_in_table(table_name: str) -> t.List[str]: ] # restate A in prod - engine_adapter.execute("delete from test.external_table where ts = '2024-01-01 01:30:00'") + engine_adapter.execute( + "delete from test.external_table where ts = '2024-01-01 01:30:00'" + ) ctx.plan( restate_models=["test.a"], start="2024-01-01 01:00:00", @@ -914,9 +972,13 @@ def test_prod_restatement_plan_causes_dev_intervals_to_be_processed_in_next_dev_ "ts": exp.DataType.build("timestamp"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # plan + apply A[hourly] in prod @@ -932,7 +994,8 @@ def test_prod_restatement_plan_causes_dev_intervals_to_be_processed_in_next_dev_ def _dates_in_table(table_name: str) -> t.List[str]: return [ - str(r[0]) for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") + str(r[0]) + for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") ] # verify initial state @@ -945,7 +1008,9 @@ def _dates_in_table(table_name: str) -> t.List[str]: ] # restate A in prod - engine_adapter.execute("delete from test.external_table where ts = '2024-01-01 01:30:00'") + engine_adapter.execute( + "delete from test.external_table where ts = '2024-01-01 01:30:00'" + ) ctx.plan( restate_models=["test.a"], start="2024-01-01 01:00:00", @@ -1049,9 +1114,13 @@ def test_prod_restatement_plan_causes_dev_intervals_to_be_widened_on_full_restat "ts": exp.DataType.build("timestamp"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # plan + apply A[daily] in prod @@ -1067,7 +1136,8 @@ def test_prod_restatement_plan_causes_dev_intervals_to_be_widened_on_full_restat def _dates_in_table(table_name: str) -> t.List[str]: return [ - str(r[0]) for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") + str(r[0]) + for r in engine_adapter.fetchall(f"select ts from {table_name} order by ts") ] # verify initial state @@ -1080,7 +1150,9 @@ def _dates_in_table(table_name: str) -> t.List[str]: ] # restate A in prod - engine_adapter.execute("delete from test.external_table where ts = '2024-01-02 01:30:00'") + engine_adapter.execute( + "delete from test.external_table where ts = '2024-01-02 01:30:00'" + ) ctx.plan( restate_models=["test.a"], start="2024-01-02 00:00:00", @@ -1184,9 +1256,13 @@ def test_prod_restatement_plan_missing_model_in_dev( "ts": exp.DataType.build("timestamp"), } external_table = exp.table_(table="external_table", db="test", quoted=True) - engine_adapter.create_table(table_name=external_table, target_columns_to_types=columns_to_types) + engine_adapter.create_table( + table_name=external_table, target_columns_to_types=columns_to_types + ) engine_adapter.insert_append( - table_name=external_table, query_or_df=df, target_columns_to_types=columns_to_types + table_name=external_table, + query_or_df=df, + target_columns_to_types=columns_to_types, ) # plan + apply A[hourly] in dev @@ -1255,7 +1331,9 @@ def test_prod_restatement_plan_includes_related_unpromoted_snapshots(tmp_path: P select a, ts from test.a """) - config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb", start="2024-01-01")) + config = Config( + model_defaults=ModelDefaultsConfig(dialect="duckdb", start="2024-01-01") + ) ctx = Context(paths=[tmp_path], config=config) def _all_snapshots() -> t.Dict[SnapshotId, Snapshot]: @@ -1342,7 +1420,9 @@ def _all_snapshots() -> t.Dict[SnapshotId, Snapshot]: all_snapshots_prior_to_restatement = _all_snapshots() assert len(all_snapshots_prior_to_restatement) == 7 - def _snapshot_instances(lst: t.Dict[SnapshotId, Snapshot], name_match: str) -> t.List[Snapshot]: + def _snapshot_instances( + lst: t.Dict[SnapshotId, Snapshot], name_match: str + ) -> t.List[Snapshot]: return [s for s_id, s in lst.items() if name_match in s_id.name] # verify initial state @@ -1360,7 +1440,9 @@ def _snapshot_instances(lst: t.Dict[SnapshotId, Snapshot], name_match: str) -> t assert len(_snapshot_instances(all_snapshots_prior_to_restatement, '"d"')) == 1 # restate A in prod - ctx.plan(environment="prod", restate_models=['"memory"."test"."a"'], auto_apply=True) + ctx.plan( + environment="prod", restate_models=['"memory"."test"."a"'], auto_apply=True + ) all_snapshots_after_restatement = _all_snapshots() @@ -1417,17 +1499,28 @@ def test_restatement_of_full_model_with_start(init_and_plan_context: t.Callable) sushi_customer_interval = restatement_plan.restatements[ context.get_snapshot("sushi.customers").snapshot_id ] - assert sushi_customer_interval == (to_timestamp("2023-01-01"), to_timestamp("2023-01-09")) + assert sushi_customer_interval == ( + to_timestamp("2023-01-01"), + to_timestamp("2023-01-09"), + ) waiter_by_day_interval = restatement_plan.restatements[ context.get_snapshot("sushi.waiter_as_customer_by_day").snapshot_id ] - assert waiter_by_day_interval == (to_timestamp("2023-01-07"), to_timestamp("2023-01-08")) + assert waiter_by_day_interval == ( + to_timestamp("2023-01-07"), + to_timestamp("2023-01-08"), + ) @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_restatement_should_not_override_environment_statements(init_and_plan_context: t.Callable): +def test_restatement_should_not_override_environment_statements( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") - context.config.before_all = ["SELECT 'test_before_all';", *context.config.before_all] + context.config.before_all = [ + "SELECT 'test_before_all';", + *context.config.before_all, + ] context.load() context.plan("prod", auto_apply=True, no_prompts=True, skip_tests=True) @@ -1447,7 +1540,9 @@ def test_restatement_should_not_override_environment_statements(init_and_plan_co @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_restatement_shouldnt_backfill_beyond_prod_intervals(init_and_plan_context: t.Callable): +def test_restatement_shouldnt_backfill_beyond_prod_intervals( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") model = context.get_model("sushi.top_waiters") @@ -1465,9 +1560,9 @@ def test_restatement_shouldnt_backfill_beyond_prod_intervals(init_and_plan_conte ) intervals_by_id = {i.snapshot_id: i for i in restatement_plan.missing_intervals} # Make sure the intervals don't go beyond the prod intervals - assert intervals_by_id[context.get_snapshot("sushi.top_waiters").snapshot_id].intervals[-1][ - 1 - ] == to_timestamp("2023-01-08 15:00:00 UTC") + assert intervals_by_id[ + context.get_snapshot("sushi.top_waiters").snapshot_id + ].intervals[-1][1] == to_timestamp("2023-01-08 15:00:00 UTC") assert intervals_by_id[ context.get_snapshot("sushi.waiter_revenue_by_day").snapshot_id ].intervals[-1][1] == to_timestamp("2023-01-08 00:00:00 UTC") @@ -1491,7 +1586,9 @@ def test_restatement_plan_interval_external_visibility(tmp_path: Path): models_dir = tmp_path / "models" models_dir.mkdir() - lock_file_path = tmp_path / "test.lock" # python model blocks while this file is present + lock_file_path = ( + tmp_path / "test.lock" + ) # python model blocks while this file is present evaluation_lock_file_path = ( tmp_path / "evaluation.lock" @@ -1537,7 +1634,9 @@ def entrypoint(evaluator: MacroEvaluator) -> str: gateways={ "": GatewayConfig( connection=DuckDBConnectionConfig(database=str(tmp_path / "db.db")), - state_connection=DuckDBConnectionConfig(database=str(tmp_path / "state.db")), + state_connection=DuckDBConnectionConfig( + database=str(tmp_path / "state.db") + ), ) }, model_defaults=ModelDefaultsConfig(dialect="duckdb", start="2024-01-01"), @@ -1596,7 +1695,9 @@ def _run_restatement_plan(tmp_path: Path, config: Config, q: queue.Queue): q.put("plan_started") plan = restatement_ctx.plan( - environment="prod", restate_models=['"db"."test"."model_a"'], auto_apply=True + environment="prod", + restate_models=['"db"."test"."model_a"'], + auto_apply=True, ) q.put("plan_completed") @@ -1607,7 +1708,9 @@ def _run_restatement_plan(tmp_path: Path, config: Config, q: queue.Queue): executor = ThreadPoolExecutor() q: queue.Queue = queue.Queue() - restatement_plan_future = executor.submit(_run_restatement_plan, tmp_path, config, q) + restatement_plan_future = executor.submit( + _run_restatement_plan, tmp_path, config, q + ) assert q.get() == "thread_started" try: @@ -1733,7 +1836,9 @@ def test_restatement_plan_detects_prod_deployment_during_restatement(tmp_path: P models_dir = tmp_path / "models" models_dir.mkdir() - lock_file_path = tmp_path / "test.lock" # python model blocks while this file is present + lock_file_path = ( + tmp_path / "test.lock" + ) # python model blocks while this file is present evaluation_lock_file_path = ( tmp_path / "evaluation.lock" @@ -1770,7 +1875,9 @@ def entrypoint(evaluator: MacroEvaluator) -> str: gateways={ "": GatewayConfig( connection=DuckDBConnectionConfig(database=str(tmp_path / "db.db")), - state_connection=DuckDBConnectionConfig(database=str(tmp_path / "state.db")), + state_connection=DuckDBConnectionConfig( + database=str(tmp_path / "state.db") + ), ) }, model_defaults=ModelDefaultsConfig(dialect="duckdb", start="2024-01-01"), @@ -1820,7 +1927,9 @@ def _run_restatement_plan(tmp_path: Path, config: Config, q: queue.Queue): expected_error = None try: restatement_ctx.plan( - environment="prod", restate_models=['"db"."test"."model_a"'], auto_apply=True + environment="prod", + restate_models=['"db"."test"."model_a"'], + auto_apply=True, ) except ConflictingPlanError as e: expected_error = e @@ -1832,7 +1941,9 @@ def _run_restatement_plan(tmp_path: Path, config: Config, q: queue.Queue): q: queue.Queue = queue.Queue() lock_file_path.touch() - restatement_plan_future = executor.submit(_run_restatement_plan, tmp_path, config, q) + restatement_plan_future = executor.submit( + _run_restatement_plan, tmp_path, config, q + ) restatement_plan_future.add_done_callback(lambda _: executor.shutdown()) assert q.get() == "thread_started" @@ -1908,8 +2019,14 @@ def test_restatement_plan_outside_parent_date_range(init_and_plan_context: t.Cal assert plan.requires_backfill assert plan.restatements == { - restated_snapshot.snapshot_id: (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")), - downstream_snapshot.snapshot_id: (to_timestamp("2023-01-01"), to_timestamp("2023-01-09")), + restated_snapshot.snapshot_id: ( + to_timestamp("2023-01-01"), + to_timestamp("2023-01-02"), + ), + downstream_snapshot.snapshot_id: ( + to_timestamp("2023-01-01"), + to_timestamp("2023-01-09"), + ), } assert plan.missing_intervals == [ SnapshotIntervals( diff --git a/tests/core/integration/test_run.py b/tests/core/integration/test_run.py index c3e6626ad0..b09210fbf4 100644 --- a/tests/core/integration/test_run.py +++ b/tests/core/integration/test_run.py @@ -1,6 +1,7 @@ from __future__ import annotations import typing as t + import pytest import time_machine from pytest_mock.plugin import MockerFixture @@ -8,11 +9,7 @@ from sqlmesh.core import constants as c from sqlmesh.core import dialect as d from sqlmesh.core.config.categorizer import CategorizerConfig -from sqlmesh.core.model import ( - SqlModel, - PythonModel, - load_sql_based_model, -) +from sqlmesh.core.model import PythonModel, SqlModel, load_sql_based_model from sqlmesh.utils.date import to_timestamp if t.TYPE_CHECKING: @@ -31,9 +28,13 @@ def test_run_with_select_models( with time_machine.travel("2023-01-09 00:00:00 UTC"): assert context.run(select_models=["*waiter_revenue_by_day"]) - snapshots = context.state_sync.state_sync.get_snapshots(context.snapshots.values()) + snapshots = context.state_sync.state_sync.get_snapshots( + context.snapshots.values() + ) # Only waiter_revenue_by_day and its parents should be backfilled up to 2023-01-09. - assert {s.name: s.intervals[0][1] for s in snapshots.values() if s.intervals} == { + assert { + s.name: s.intervals[0][1] for s in snapshots.values() if s.intervals + } == { '"memory"."sushi"."waiter_revenue_by_day"': to_timestamp("2023-01-09"), '"memory"."sushi"."order_items"': to_timestamp("2023-01-09"), '"memory"."sushi"."orders"': to_timestamp("2023-01-09"), @@ -68,11 +69,17 @@ def test_run_with_select_models_no_auto_upstream( context.plan("prod", no_prompts=True, skip_tests=True, auto_apply=True) with time_machine.travel("2023-01-09 00:00:00 UTC"): - assert context.run(select_models=["*waiter_revenue_by_day"], no_auto_upstream=True) + assert context.run( + select_models=["*waiter_revenue_by_day"], no_auto_upstream=True + ) - snapshots = context.state_sync.state_sync.get_snapshots(context.snapshots.values()) + snapshots = context.state_sync.state_sync.get_snapshots( + context.snapshots.values() + ) # Only waiter_revenue_by_day should be backfilled up to 2023-01-09. - assert {s.name: s.intervals[0][1] for s in snapshots.values() if s.intervals} == { + assert { + s.name: s.intervals[0][1] for s in snapshots.values() if s.intervals + } == { '"memory"."sushi"."waiter_revenue_by_day"': to_timestamp("2023-01-09"), '"memory"."sushi"."order_items"': to_timestamp("2023-01-08"), '"memory"."sushi"."orders"': to_timestamp("2023-01-08"), @@ -95,15 +102,16 @@ def test_run_with_select_models_no_auto_upstream( @time_machine.travel("2023-01-08 15:00:00 UTC") -def test_run_respects_excluded_transitive_dependencies(init_and_plan_context: t.Callable): +def test_run_respects_excluded_transitive_dependencies( + init_and_plan_context: t.Callable, +): context, _ = init_and_plan_context("examples/sushi") # Graph: C <- B <- A # B is a transitive dependency linking A and C # Note that the alphabetical ordering of the model names is intentional and helps # surface the problem - expressions_a = d.parse( - f""" + expressions_a = d.parse(f""" MODEL ( name memory.sushi.test_model_c, kind FULL, @@ -113,14 +121,12 @@ def test_run_respects_excluded_transitive_dependencies(init_and_plan_context: t. ); SELECT @execution_ts AS execution_ts - """ - ) + """) model_c = load_sql_based_model(expressions_a) context.upsert_model(model_c) # A VIEW model with no partials allowed and a daily cron instead of hourly. - expressions_b = d.parse( - f""" + expressions_b = d.parse(f""" MODEL ( name memory.sushi.test_model_b, kind VIEW, @@ -129,13 +135,11 @@ def test_run_respects_excluded_transitive_dependencies(init_and_plan_context: t. ); SELECT * FROM memory.sushi.test_model_c - """ - ) + """) model_b = load_sql_based_model(expressions_b) context.upsert_model(model_b) - expressions_a = d.parse( - f""" + expressions_a = d.parse(f""" MODEL ( name memory.sushi.test_model_a, kind FULL, @@ -144,16 +148,15 @@ def test_run_respects_excluded_transitive_dependencies(init_and_plan_context: t. ); SELECT * FROM memory.sushi.test_model_b - """ - ) + """) model_a = load_sql_based_model(expressions_a) context.upsert_model(model_a) context.plan("prod", skip_tests=True, auto_apply=True, no_prompts=True) assert ( - context.fetchdf("SELECT execution_ts FROM memory.sushi.test_model_c")["execution_ts"].iloc[ - 0 - ] + context.fetchdf("SELECT execution_ts FROM memory.sushi.test_model_c")[ + "execution_ts" + ].iloc[0] == "2023-01-08 15:00:00" ) @@ -211,7 +214,11 @@ def test_snapshot_triggers(init_and_plan_context: t.Callable, mocker: MockerFixt } context.upsert_model(SqlModel.parse_obj(waiter_revenue_by_day_kwargs)) - context.plan(auto_apply=True, no_prompts=True, categorizer_config=CategorizerConfig.all_full()) + context.plan( + auto_apply=True, + no_prompts=True, + categorizer_config=CategorizerConfig.all_full(), + ) scheduler = context.scheduler() @@ -241,7 +248,9 @@ def test_snapshot_triggers(init_and_plan_context: t.Callable, mocker: MockerFixt if model_name in ("orders", "order_items", "waiter_revenue_by_day"): assert auto_restatement_triggers == [model_name] elif model_name in ("customer_revenue_lifetime", "customer_revenue_by_day"): - assert sorted(auto_restatement_triggers) == sorted(["orders", "order_items"]) + assert sorted(auto_restatement_triggers) == sorted( + ["orders", "order_items"] + ) elif model_name == "top_waiters": assert auto_restatement_triggers == ["waiter_revenue_by_day"] else: diff --git a/tests/core/integration/utils.py b/tests/core/integration/utils.py index ba233080b5..96212438c8 100644 --- a/tests/core/integration/utils.py +++ b/tests/core/integration/utils.py @@ -1,7 +1,7 @@ from __future__ import annotations import typing as t -from sqlmesh.core.model.common import ParsableSql + from sqlglot import exp from sqlglot.expressions import DataType @@ -9,24 +9,15 @@ from sqlmesh.core.context import Context from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.environment import EnvironmentNamingInfo -from sqlmesh.core.model import ( - IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind, - ModelKind, - ModelKindName, - SqlModel, - TimeColumn, -) +from sqlmesh.core.model import (IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, ModelKind, + ModelKindName, SqlModel, TimeColumn) +from sqlmesh.core.model.common import ParsableSql from sqlmesh.core.model.kind import model_kind_type_from_name from sqlmesh.core.plan import Plan, PlanBuilder -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Snapshot, - SnapshotChangeCategory, - SnapshotId, - SnapshotInfoLike, - SnapshotTableInfo, -) +from sqlmesh.core.snapshot import (DeployabilityIndex, Snapshot, + SnapshotChangeCategory, SnapshotId, + SnapshotInfoLike, SnapshotTableInfo) from sqlmesh.utils.date import TimeLike @@ -81,7 +72,9 @@ def apply_to_environment( start=plan_start or start(context) if environment != c.PROD else None, forward_only=choice == SnapshotChangeCategory.FORWARD_ONLY, include_unmodified=True, - allow_destructive_models=allow_destructive_models if allow_destructive_models else [], + allow_destructive_models=( + allow_destructive_models if allow_destructive_models else [] + ), enable_preview=enable_preview, ) if environment != c.PROD: @@ -98,7 +91,9 @@ def apply_to_environment( plan = plan_builder.build() context.apply(plan) - validate_apply_basics(context, environment, plan.snapshots.values(), plan.deployability_index) + validate_apply_basics( + context, environment, plan.snapshots.values(), plan.deployability_index + ) for validator in apply_validators: validator(context) return plan @@ -119,7 +114,9 @@ def change_data_type( for data_type in data_types: if data_type.this == old_type: data_type.set("this", new_type) - context.upsert_model(model_name, query_=ParsableSql(sql=query.sql(dialect=model.dialect))) + context.upsert_model( + model_name, query_=ParsableSql(sql=query.sql(dialect=model.dialect)) + ) elif model.columns_to_types_ is not None: for k, v in model.columns_to_types_.items(): if v.this == old_type: @@ -127,7 +124,9 @@ def change_data_type( context.upsert_model(model_name, columns=model.columns_to_types_) -def validate_snapshots_in_state_sync(snapshots: t.Iterable[Snapshot], context: Context) -> None: +def validate_snapshots_in_state_sync( + snapshots: t.Iterable[Snapshot], context: Context +) -> None: snapshot_infos = map(to_snapshot_info, snapshots) state_sync_table_infos = map( to_snapshot_info, context.state_reader.get_snapshots(snapshots).values() @@ -157,7 +156,10 @@ def validate_tables( if not snapshot.is_model or snapshot.is_external: continue table_should_exist = not snapshot.is_embedded - assert adapter.table_exists(snapshot.table_name(is_deployable)) == table_should_exist + assert ( + adapter.table_exists(snapshot.table_name(is_deployable)) + == table_should_exist + ) if table_should_exist: assert select_all(snapshot.table_name(is_deployable), adapter) @@ -260,7 +262,8 @@ def validate_query_change( not_modified = [ snapshot.name for snapshot in context.snapshots.values() - if snapshot.name not in directly_modified and snapshot.name not in indirectly_modified + if snapshot.name not in directly_modified + and snapshot.name not in indirectly_modified ] if change_category == SnapshotChangeCategory.BREAKING and not logical: @@ -294,8 +297,12 @@ def _validate_apply(context): def initial_add(context: Context, environment: str): assert not context.state_reader.get_environment(environment) - plan = context.plan(environment, start=start(context), create_from="nonexistent_env") - validate_plan_changes(plan, added={x.snapshot_id for x in context.snapshots.values()}) + plan = context.plan( + environment, start=start(context), create_from="nonexistent_env" + ) + validate_plan_changes( + plan, added={x.snapshot_id for x in context.snapshots.values()} + ) context.apply(plan) validate_apply_basics(context, environment, plan.snapshots.values()) @@ -327,7 +334,9 @@ def validate_model_kind_change( "assert_item_price_above_zero", ] if kind_name == ModelKindName.INCREMENTAL_BY_TIME_RANGE: - kind: ModelKind = IncrementalByTimeRangeKind(time_column=TimeColumn(column="event_date")) + kind: ModelKind = IncrementalByTimeRangeKind( + time_column=TimeColumn(column="event_date") + ) elif kind_name == ModelKindName.INCREMENTAL_BY_UNIQUE_KEY: kind = IncrementalByUniqueKeyKind(unique_key="id") else: diff --git a/tests/core/linter/test_builtin.py b/tests/core/linter/test_builtin.py index 0ff91470ff..0b3340bf8f 100644 --- a/tests/core/linter/test_builtin.py +++ b/tests/core/linter/test_builtin.py @@ -221,14 +221,19 @@ def test_no_missing_unit_tests(tmp_path, copy_to_temp_path): assert any("is missing unit test(s)" in msg for msg in violation_messages) # Check that models with existing tests don't have violations - models_with_tests = ["customer_revenue_by_day", "customer_revenue_lifetime", "order_items"] + models_with_tests = [ + "customer_revenue_by_day", + "customer_revenue_lifetime", + "order_items", + ] for model_name in models_with_tests: model_violations = [ lint for lint in lints - if model_name in lint.violation_msg and "is missing unit test(s)" in lint.violation_msg + if model_name in lint.violation_msg + and "is missing unit test(s)" in lint.violation_msg ] - assert len(model_violations) == 0, ( - f"Model {model_name} should not have a violation since it has a test" - ) + assert ( + len(model_violations) == 0 + ), f"Model {model_name} should not have a violation since it has a test" diff --git a/tests/core/linter/test_helpers.py b/tests/core/linter/test_helpers.py index c3ba46f304..bb17df355c 100644 --- a/tests/core/linter/test_helpers.py +++ b/tests/core/linter/test_helpers.py @@ -1,9 +1,7 @@ from sqlmesh import Context -from sqlmesh.core.linter.helpers import ( - read_range_from_file, - get_range_of_model_block, - get_range_of_a_key_in_model_block, -) +from sqlmesh.core.linter.helpers import (get_range_of_a_key_in_model_block, + get_range_of_model_block, + read_range_from_file) from sqlmesh.core.model import SqlModel diff --git a/tests/core/metric/test_metric.py b/tests/core/metric/test_metric.py index 51c97fbe3d..95fc44176e 100644 --- a/tests/core/metric/test_metric.py +++ b/tests/core/metric/test_metric.py @@ -7,16 +7,14 @@ def test_load_metric_ddl(): - a = d.parse_one( - """ + a = d.parse_one(""" -- description a METRIC ( name A, expression SUM(x), owner b ); - """ - ) + """) meta = load_metric_ddl(a, dialect="") assert meta.name == "a" @@ -30,31 +28,28 @@ def test_load_invalid(): ConfigError, match=r"Only METRIC\(...\) statements are allowed. Found SELECT" ): load_metric_ddl( - d.parse_one( - """ + d.parse_one(""" SELECT 1; - """ - ), + """), dialect="", ) - with pytest.raises(ConfigError, match=r"Metric 'a' missing an aggregation or metric ref."): + with pytest.raises( + ConfigError, match=r"Metric 'a' missing an aggregation or metric ref." + ): load_metric_ddl( - d.parse_one( - """ + d.parse_one(""" METRIC ( name a, expression 1 ) - """ - ), + """), dialect="", ).to_metric({}, {}) def test_expand_metrics(): - expressions = d.parse( - """ + expressions = d.parse(""" -- description a METRIC ( name a, @@ -82,8 +77,7 @@ def test_expand_metrics(): expression c + 1, owner b ); - """ - ) + """) metas = {} for expr in expressions: @@ -114,7 +108,10 @@ def test_expand_metrics(): metric_d = metrics["d"] assert metric_d.expression.sql() == "c + 1" - assert metric_d.expanded.sql() == "SUM(model.x) AS a / COUNT(DISTINCT model.y) AS b + 1" + assert ( + metric_d.expanded.sql() + == "SUM(model.x) AS a / COUNT(DISTINCT model.y) AS b + 1" + ) assert metric_d.formula.sql() == "a / b + 1 AS d" assert metric_d.aggs == { @@ -150,7 +147,9 @@ def test_get_measure_and_dim_tables(): assert _get_measure_and_dim_tables( d.parse_one("SUM(IF(c.z = 'dim' AND b.y > 0, (a.x + a.x) + 3, 0))") ) == ("a", ("c", "b")) - assert _get_measure_and_dim_tables(d.parse_one("SUM(CASE b.y WHEN 1 THEN a.x ELSE 0 END)")) == ( + assert _get_measure_and_dim_tables( + d.parse_one("SUM(CASE b.y WHEN 1 THEN a.x ELSE 0 END)") + ) == ( "a", ("b",), ) diff --git a/tests/core/state_sync/test_export_import.py b/tests/core/state_sync/test_export_import.py index 769fa2c2fa..a9cb43f8e1 100644 --- a/tests/core/state_sync/test_export_import.py +++ b/tests/core/state_sync/test_export_import.py @@ -1,15 +1,18 @@ -import pytest +import json from pathlib import Path -from sqlmesh.core.state_sync import StateSync, EngineAdapterStateSync, CachingStateSync -from sqlmesh.core.state_sync.export_import import export_state, import_state -from sqlmesh.utils.errors import SQLMeshError -from sqlmesh.core import constants as c + +import pytest + from sqlmesh.cli.project_init import init_example_project +from sqlmesh.core import constants as c +from sqlmesh.core.config import (Config, DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig) from sqlmesh.core.context import Context from sqlmesh.core.environment import Environment -from sqlmesh.core.config import Config, GatewayConfig, DuckDBConnectionConfig, ModelDefaultsConfig - -import json +from sqlmesh.core.state_sync import (CachingStateSync, EngineAdapterStateSync, + StateSync) +from sqlmesh.core.state_sync.export_import import export_state, import_state +from sqlmesh.utils.errors import SQLMeshError @pytest.fixture @@ -17,8 +20,12 @@ def example_project_config(tmp_path: Path) -> Config: return Config( gateways={ "main": GatewayConfig( - connection=DuckDBConnectionConfig(database=str(tmp_path / "warehouse.db")), - state_connection=DuckDBConnectionConfig(database=str(tmp_path / "state.db")), + connection=DuckDBConnectionConfig( + database=str(tmp_path / "warehouse.db") + ), + state_connection=DuckDBConnectionConfig( + database=str(tmp_path / "state.db") + ), ) }, default_gateway="main", @@ -75,7 +82,9 @@ def test_export_entire_project( tmp_path: Path, example_project_config: Config, state_sync: StateSync ) -> None: init_example_project(path=tmp_path, engine_type="duckdb") - context = Context(paths=tmp_path, config=example_project_config, state_sync=state_sync) + context = Context( + paths=tmp_path, config=example_project_config, state_sync=state_sync + ) # prod plan = context.plan(auto_apply=True) @@ -128,7 +137,9 @@ def test_export_entire_project( assert len(state["snapshots"]) > 0 snapshot_names = [s["name"] for s in state["snapshots"]] assert len(snapshot_names) == 5 - assert '"warehouse"."sqlmesh_example"."full_model"' in snapshot_names # will be in here twice + assert ( + '"warehouse"."sqlmesh_example"."full_model"' in snapshot_names + ) # will be in here twice assert '"warehouse"."sqlmesh_example"."incremental_model"' in snapshot_names assert '"warehouse"."sqlmesh_example"."seed_model"' in snapshot_names assert '"warehouse"."sqlmesh_example"."new_model"' in snapshot_names @@ -138,14 +149,20 @@ def test_export_entire_project( prod = state["environments"]["prod"]["environment"] assert len(prod["snapshots"]) == 3 - prod_snapshot_ids = [s.snapshot_id for s in Environment.model_validate(prod).snapshots] + prod_snapshot_ids = [ + s.snapshot_id for s in Environment.model_validate(prod).snapshots + ] dev = state["environments"]["dev"]["environment"] assert len(dev["snapshots"]) == 4 - dev_snapshot_ids = [s.snapshot_id for s in Environment.model_validate(dev).snapshots] + dev_snapshot_ids = [ + s.snapshot_id for s in Environment.model_validate(dev).snapshots + ] full_model_id = next(s for s in dev_snapshot_ids if "full_model" in s.name) - incremental_model_id = next(s for s in dev_snapshot_ids if "incremental_model" in s.name) + incremental_model_id = next( + s for s in dev_snapshot_ids if "incremental_model" in s.name + ) seed_model_id = next(s for s in dev_snapshot_ids if "seed_model" in s.name) new_model_id = next(s for s in dev_snapshot_ids if "new_model" in s.name) @@ -160,7 +177,9 @@ def test_export_specific_environment( ) -> None: output_file = tmp_path / "state_dump.json" init_example_project(path=tmp_path, engine_type="duckdb") - context = Context(paths=tmp_path, config=example_project_config, state_sync=state_sync) + context = Context( + paths=tmp_path, config=example_project_config, state_sync=state_sync + ) # create prod context.plan(auto_apply=True) @@ -202,7 +221,9 @@ def test_export_specific_environment( assert any("full_model" in name for name in snapshot_names) assert any("incremental_model" in name for name in snapshot_names) assert any("seed_model" in name for name in snapshot_names) - dev_full_model = next(s for s in dev_state["snapshots"] if "full_model" in s["name"]) + dev_full_model = next( + s for s in dev_state["snapshots"] if "full_model" in s["name"] + ) assert len(dev_state["environments"]) == 1 assert "dev" in dev_state["environments"] @@ -218,13 +239,18 @@ def test_export_specific_environment( assert any("full_model" in name for name in snapshot_names) assert any("incremental_model" in name for name in snapshot_names) assert any("seed_model" in name for name in snapshot_names) - prod_full_model = next(s for s in prod_state["snapshots"] if "full_model" in s["name"]) + prod_full_model = next( + s for s in prod_state["snapshots"] if "full_model" in s["name"] + ) assert len(prod_state["environments"]) == 1 assert "prod" in prod_state["environments"] assert prod_state["metadata"]["importable"] - assert dev_full_model["fingerprint"]["data_hash"] != prod_full_model["fingerprint"]["data_hash"] + assert ( + dev_full_model["fingerprint"]["data_hash"] + != prod_full_model["fingerprint"]["data_hash"] + ) def test_export_local_state( @@ -232,7 +258,9 @@ def test_export_local_state( ) -> None: output_file = tmp_path / "state_dump.json" init_example_project(path=tmp_path, engine_type="duckdb") - context = Context(paths=tmp_path, config=example_project_config, state_sync=state_sync) + context = Context( + paths=tmp_path, config=example_project_config, state_sync=state_sync + ) # create prod context.plan(auto_apply=True) @@ -309,15 +337,21 @@ def test_import_invalid_file(tmp_path: Path, state_sync: StateSync) -> None: import_state(state_sync, state_file) state_file.write_text('{ "metadata": [] }') - with pytest.raises(SQLMeshError, match=r"Expecting the 'metadata' key to contain an object"): + with pytest.raises( + SQLMeshError, match=r"Expecting the 'metadata' key to contain an object" + ): import_state(state_sync, state_file) state_file.write_text('{ "metadata": {} }') - with pytest.raises(SQLMeshError, match=r"Unable to determine state file format version"): + with pytest.raises( + SQLMeshError, match=r"Unable to determine state file format version" + ): import_state(state_sync, state_file) state_file.write_text('{ "metadata": { "file_version": "blah" } }') - with pytest.raises(SQLMeshError, match=r"Unable to parse state file format version"): + with pytest.raises( + SQLMeshError, match=r"Unable to parse state file format version" + ): import_state(state_sync, state_file) state_file.write_text('{ "metadata": { "file_version": 1, "importable": false } }') @@ -325,12 +359,16 @@ def test_import_invalid_file(tmp_path: Path, state_sync: StateSync) -> None: import_state(state_sync, state_file) -def test_import_from_older_version_export_fails(tmp_path: Path, state_sync: StateSync) -> None: +def test_import_from_older_version_export_fails( + tmp_path: Path, state_sync: StateSync +) -> None: state_sync.migrate() current_version = state_sync.get_versions() major, minor = current_version.minor_sqlmesh_version - older_version = current_version.copy(update=dict(sqlmesh_version=f"{major}.{minor - 1}.0")) + older_version = current_version.copy( + update=dict(sqlmesh_version=f"{major}.{minor - 1}.0") + ) assert older_version.minor_sqlmesh_version < current_version.minor_sqlmesh_version @@ -353,12 +391,16 @@ def test_import_from_older_version_export_fails(tmp_path: Path, state_sync: Stat import_state(state_sync, state_file) -def test_import_from_newer_version_export_fails(tmp_path: Path, state_sync: StateSync) -> None: +def test_import_from_newer_version_export_fails( + tmp_path: Path, state_sync: StateSync +) -> None: state_sync.migrate() current_version = state_sync.get_versions() major, minor = current_version.minor_sqlmesh_version - newer_version = current_version.copy(update=dict(sqlmesh_version=f"{major}.{minor + 1}.0")) + newer_version = current_version.copy( + update=dict(sqlmesh_version=f"{major}.{minor + 1}.0") + ) assert newer_version.minor_sqlmesh_version > current_version.minor_sqlmesh_version @@ -386,7 +428,9 @@ def test_import_local_state_fails( ) -> None: output_file = tmp_path / "state_dump.json" init_example_project(path=tmp_path, engine_type="duckdb") - context = Context(paths=tmp_path, config=example_project_config, state_sync=state_sync) + context = Context( + paths=tmp_path, config=example_project_config, state_sync=state_sync + ) export_state(state_sync, output_file, context.snapshots) state = json.loads(output_file.read_text(encoding="utf8")) @@ -401,7 +445,9 @@ def test_import_partial( ) -> None: output_file = tmp_path / "state_dump.json" init_example_project(path=tmp_path, engine_type="duckdb") - context = Context(paths=tmp_path, config=example_project_config, state_sync=state_sync) + context = Context( + paths=tmp_path, config=example_project_config, state_sync=state_sync + ) # create prod context.plan(auto_apply=True) @@ -450,11 +496,15 @@ def test_import_partial( ).has_changes # prod has changes the 'new_model' model hasnt been applied -def test_roundtrip(tmp_path: Path, example_project_config: Config, state_sync: StateSync) -> None: +def test_roundtrip( + tmp_path: Path, example_project_config: Config, state_sync: StateSync +) -> None: state_file = tmp_path / "state_dump.json" init_example_project(path=tmp_path, engine_type="duckdb") - context = Context(paths=tmp_path, config=example_project_config, state_sync=state_sync) + context = Context( + paths=tmp_path, config=example_project_config, state_sync=state_sync + ) # populate initial state plan = context.plan(auto_apply=True) @@ -541,7 +591,9 @@ def test_roundtrip_includes_auto_restatements( SELECT 1 as id; """) - context = Context(paths=tmp_path, config=example_project_config, state_sync=state_sync) + context = Context( + paths=tmp_path, config=example_project_config, state_sync=state_sync + ) context.plan(auto_apply=True) # dump state @@ -576,8 +628,12 @@ def test_roundtrip_includes_environment_statements(tmp_path: Path) -> None: config = Config( gateways={ "main": GatewayConfig( - connection=DuckDBConnectionConfig(database=str(tmp_path / "warehouse.db")), - state_connection=DuckDBConnectionConfig(database=str(tmp_path / "state.db")), + connection=DuckDBConnectionConfig( + database=str(tmp_path / "warehouse.db") + ), + state_connection=DuckDBConnectionConfig( + database=str(tmp_path / "state.db") + ), ) }, default_gateway="main", @@ -596,8 +652,13 @@ def test_roundtrip_includes_environment_statements(tmp_path: Path) -> None: environments = json.loads(state_file.read_text(encoding="utf8"))["environments"] - assert environments["prod"]["statements"][0]["before_all"][0] == "select 1 as before_all" - assert environments["prod"]["statements"][0]["after_all"][0] == "select 2 as after_all" + assert ( + environments["prod"]["statements"][0]["before_all"][0] + == "select 1 as before_all" + ) + assert ( + environments["prod"]["statements"][0]["after_all"][0] == "select 2 as after_all" + ) assert not context.plan().has_changes diff --git a/tests/core/state_sync/test_state_sync.py b/tests/core/state_sync/test_state_sync.py index 348a883fd5..378463f54e 100644 --- a/tests/core/state_sync/test_state_sync.py +++ b/tests/core/state_sync/test_state_sync.py @@ -16,38 +16,17 @@ from sqlmesh.core.dialect import parse_one from sqlmesh.core.engine_adapter import create_engine_adapter from sqlmesh.core.environment import Environment, EnvironmentStatements -from sqlmesh.core.model import ( - FullKind, - IncrementalByTimeRangeKind, - Seed, - SeedKind, - SeedModel, - SqlModel, -) -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotChangeCategory, - SnapshotId, - SnapshotIntervals, - SnapshotNameVersion, - SnapshotTableCleanupTask, - missing_intervals, -) -from sqlmesh.core.state_sync import ( - CachingStateSync, - EngineAdapterStateSync, -) -from sqlmesh.core.state_sync.base import ( - SCHEMA_VERSION, - SQLGLOT_VERSION, - Versions, -) -from sqlmesh.core.state_sync.common import ( - ExpiredBatchRange, - LimitBoundary, - PromotionResult, - RowBoundary, -) +from sqlmesh.core.model import (FullKind, IncrementalByTimeRangeKind, Seed, + SeedKind, SeedModel, SqlModel) +from sqlmesh.core.snapshot import (Snapshot, SnapshotChangeCategory, + SnapshotId, SnapshotIntervals, + SnapshotNameVersion, + SnapshotTableCleanupTask, missing_intervals) +from sqlmesh.core.state_sync import CachingStateSync, EngineAdapterStateSync +from sqlmesh.core.state_sync.base import (SCHEMA_VERSION, SQLGLOT_VERSION, + Versions) +from sqlmesh.core.state_sync.common import (ExpiredBatchRange, LimitBoundary, + PromotionResult, RowBoundary) from sqlmesh.utils.date import now_timestamp, to_datetime, to_timestamp from sqlmesh.utils.errors import SQLMeshError, StateMigrationError @@ -159,7 +138,9 @@ def test_push_snapshots( state_sync.push_snapshots([snapshot_a, snapshot_b]) - assert state_sync.get_snapshots([snapshot_a.snapshot_id, snapshot_b.snapshot_id]) == { + assert state_sync.get_snapshots( + [snapshot_a.snapshot_id, snapshot_b.snapshot_id] + ) == { snapshot_a.snapshot_id: snapshot_a, snapshot_b.snapshot_id: snapshot_b, } @@ -169,7 +150,10 @@ def test_push_snapshots( state_sync.push_snapshots([snapshot_a]) assert str({snapshot_a.snapshot_id}) == mock_logger.call_args[0][1] state_sync.push_snapshots([snapshot_a, snapshot_b]) - assert str({snapshot_a.snapshot_id, snapshot_b.snapshot_id}) == mock_logger.call_args[0][1] + assert ( + str({snapshot_a.snapshot_id, snapshot_b.snapshot_id}) + == mock_logger.call_args[0][1] + ) # test serialization state_sync.push_snapshots( @@ -178,12 +162,10 @@ def test_push_snapshots( SqlModel( name="a", kind=FullKind(), - query=parse_one( - """ + query=parse_one(""" select 'x' + ' ' as y, "z" + '\' as z, - """ - ), + """), ), version="1", ) @@ -191,7 +173,9 @@ def test_push_snapshots( ) -def test_duplicates(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable) -> None: +def test_duplicates( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +) -> None: snapshot_a = make_snapshot( SqlModel( name="a", @@ -226,14 +210,18 @@ def test_duplicates(state_sync: EngineAdapterStateSync, make_snapshot: t.Callabl ) -def test_snapshots_exists(state_sync: EngineAdapterStateSync, snapshots: t.List[Snapshot]) -> None: +def test_snapshots_exists( + state_sync: EngineAdapterStateSync, snapshots: t.List[Snapshot] +) -> None: state_sync.push_snapshots(snapshots) snapshot_ids = {snapshot.snapshot_id for snapshot in snapshots} assert state_sync.snapshots_exist(snapshot_ids) == snapshot_ids @pytest.fixture -def get_snapshot_intervals(state_sync) -> t.Callable[[Snapshot], t.Optional[SnapshotIntervals]]: +def get_snapshot_intervals( + state_sync, +) -> t.Callable[[Snapshot], t.Optional[SnapshotIntervals]]: def _get_snapshot_intervals(snapshot: Snapshot) -> t.Optional[SnapshotIntervals]: intervals = state_sync.interval_state.get_snapshot_intervals([snapshot]) return intervals[0] if intervals else None @@ -275,7 +263,9 @@ def test_add_interval( snapshot.change_category = SnapshotChangeCategory.BREAKING snapshot.forward_only = True - state_sync.add_interval(snapshot, to_datetime("2020-01-16"), "2020-01-20", is_dev=True) + state_sync.add_interval( + snapshot, to_datetime("2020-01-16"), "2020-01-20", is_dev=True + ) intervals = get_snapshot_intervals(snapshot) assert intervals.intervals == [ (to_timestamp("2020-01-01"), to_timestamp("2020-01-04")), @@ -311,7 +301,9 @@ def test_add_interval_partial( ] -def test_remove_interval(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable) -> None: +def test_remove_interval( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +) -> None: snapshot_a = make_snapshot( SqlModel( name="a", @@ -341,7 +333,9 @@ def test_remove_interval(state_sync: EngineAdapterStateSync, make_snapshot: t.Ca remove_records_count = state_sync.engine_adapter.fetchone( "SELECT COUNT(*) FROM sqlmesh._intervals WHERE name = '\"a\"' AND version = 'a' AND is_removed" - )[0] # type: ignore + )[ + 0 + ] # type: ignore assert remove_records_count == num_of_removals * 2 # 2 * snapshots snapshots = state_sync.get_snapshots([snapshot_a, snapshot_b]) @@ -417,11 +411,15 @@ def test_refresh_snapshot_intervals( assert not snapshot.intervals state_sync.refresh_snapshot_intervals([snapshot]) - assert snapshot.intervals == [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))] + assert snapshot.intervals == [ + (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) + ] def test_get_snapshot_intervals( - state_sync: EngineAdapterStateSync, make_snapshot: t.Callable, get_snapshot_intervals + state_sync: EngineAdapterStateSync, + make_snapshot: t.Callable, + get_snapshot_intervals, ) -> None: state_sync.interval_state.SNAPSHOT_BATCH_SIZE = 1 @@ -460,8 +458,12 @@ def test_get_snapshot_intervals( a_intervals = get_snapshot_intervals(snapshot_a) c_intervals = get_snapshot_intervals(snapshot_c) - assert a_intervals.intervals == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] - assert c_intervals.intervals == [(to_timestamp("2020-01-03"), to_timestamp("2020-01-04"))] + assert a_intervals.intervals == [ + (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) + ] + assert c_intervals.intervals == [ + (to_timestamp("2020-01-03"), to_timestamp("2020-01-04")) + ] def test_compact_intervals( @@ -537,7 +539,9 @@ def test_compact_intervals_delete_batches( ) -def test_promote_snapshots(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_promote_snapshots( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): snapshot_a = make_snapshot( SqlModel( name="a", @@ -582,9 +586,13 @@ def test_promote_snapshots(state_sync: EngineAdapterStateSync, make_snapshot: t. state_sync.push_snapshots([snapshot_a, snapshot_b_old, snapshot_b, snapshot_c]) - promotion_result = promote_snapshots(state_sync, [snapshot_a, snapshot_b_old], "prod") + promotion_result = promote_snapshots( + state_sync, [snapshot_a, snapshot_b_old], "prod" + ) - assert set(promotion_result.added) == set([snapshot_a.table_info, snapshot_b_old.table_info]) + assert set(promotion_result.added) == set( + [snapshot_a.table_info, snapshot_b_old.table_info] + ) assert not promotion_result.removed assert not promotion_result.removed_environment_naming_info promotion_result = promote_snapshots( @@ -617,7 +625,9 @@ def test_promote_snapshots(state_sync: EngineAdapterStateSync, make_snapshot: t. > prev_snapshot_c_updated_ts ) assert ( - state_sync.get_snapshots([snapshot_b_old])[snapshot_b_old.snapshot_id].updated_ts + state_sync.get_snapshots([snapshot_b_old])[ + snapshot_b_old.snapshot_id + ].updated_ts > prev_snapshot_b_old_updated_ts ) @@ -806,10 +816,15 @@ def test_promote_snapshots_catalog_name_override_change( # B is not removed because it's catalog did not change and therefore removing would actually result # in dropping what we just added. # A is removed because it was explicitly removed from the promotion. - assert set(promotion_result.removed) == {snapshot_a.table_info, snapshot_c.table_info} + assert set(promotion_result.removed) == { + snapshot_a.table_info, + snapshot_c.table_info, + } # Make sure the removed suffix target correctly has the old catalog name set assert promotion_result.removed_environment_naming_info - assert promotion_result.removed_environment_naming_info.catalog_name_override is None + assert ( + promotion_result.removed_environment_naming_info.catalog_name_override is None + ) promotion_result = promote_snapshots( state_sync, @@ -837,7 +852,10 @@ def test_promote_snapshots_catalog_name_override_change( } # Make sure the removed suffix target correctly has the old catalog name set assert promotion_result.removed_environment_naming_info - assert promotion_result.removed_environment_naming_info.catalog_name_override == "catalog1" + assert ( + promotion_result.removed_environment_naming_info.catalog_name_override + == "catalog1" + ) def test_promote_snapshots_parent_plan_id_mismatch( @@ -933,7 +951,9 @@ def test_promote_environment_expired( assert promotion_result.added == [snapshot.table_info] -def test_promote_snapshots_no_gaps(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_promote_snapshots_no_gaps( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): model = SqlModel( name="a", query=parse_one("select 1, ds"), @@ -948,7 +968,9 @@ def test_promote_snapshots_no_gaps(state_sync: EngineAdapterStateSync, make_snap promote_snapshots(state_sync, [snapshot], "prod", no_gaps=True) new_snapshot_same_version = make_snapshot(model, version="a") - new_snapshot_same_version.change_category = SnapshotChangeCategory.INDIRECT_NON_BREAKING + new_snapshot_same_version.change_category = ( + SnapshotChangeCategory.INDIRECT_NON_BREAKING + ) new_snapshot_same_version.fingerprint = snapshot.fingerprint.copy( update={"data_hash": "new_snapshot_same_version"} ) @@ -967,7 +989,9 @@ def test_promote_snapshots_no_gaps(state_sync: EngineAdapterStateSync, make_snap SQLMeshError, match=r".*Detected missing intervals for model .*, interrupting your current plan. Please re-apply your plan to resolve this error.*", ): - promote_snapshots(state_sync, [new_snapshot_missing_interval], "prod", no_gaps=True) + promote_snapshots( + state_sync, [new_snapshot_missing_interval], "prod", no_gaps=True + ) new_snapshot_same_interval = make_snapshot(model, version="c") new_snapshot_same_interval.change_category = SnapshotChangeCategory.BREAKING @@ -1014,7 +1038,9 @@ def test_promote_snapshots_no_gaps_lookback( update={"data_hash": "new_snapshot_same_version"} ) state_sync.push_snapshots([new_snapshot_same_version]) - state_sync.add_interval(new_snapshot_same_version, "2023-01-01", "2023-01-08 15:00:00") + state_sync.add_interval( + new_snapshot_same_version, "2023-01-01", "2023-01-08 15:00:00" + ) promote_snapshots(state_sync, [new_snapshot_same_version], "prod", no_gaps=True) @@ -1085,7 +1111,9 @@ def test_start_date_gap(state_sync: EngineAdapterStateSync, make_snapshot: t.Cal promote_snapshots(state_sync, [snapshot], "prod", no_gaps=True) -def test_delete_expired_environments(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_delete_expired_environments( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): snapshot = make_snapshot( SqlModel( name="a", @@ -1118,7 +1146,9 @@ def test_delete_expired_environments(state_sync: EngineAdapterStateSync, make_sn state_sync.promote(env_a, environment_statements=environment_statements) - env_b = env_a.copy(update={"name": "test_environment_b", "expiration_ts": now_ts + 1000}) + env_b = env_a.copy( + update={"name": "test_environment_b", "expiration_ts": now_ts + 1000} + ) state_sync.promote(env_b) env_a = Environment(**json.loads(env_a.json())) @@ -1174,10 +1204,14 @@ def test_get_expired_environments_filtered_by_name( expired_all = state_sync.get_expired_environments(current_ts=now_ts) assert {e.name for e in expired_all} == {"test_env_a", "test_env_b"} - expired_filtered = state_sync.get_expired_environments(current_ts=now_ts, name="test_env_a") + expired_filtered = state_sync.get_expired_environments( + current_ts=now_ts, name="test_env_a" + ) assert [e.name for e in expired_filtered] == ["test_env_a"] - expired_non_expired = state_sync.get_expired_environments(current_ts=now_ts, name="test_env_c") + expired_non_expired = state_sync.get_expired_environments( + current_ts=now_ts, name="test_env_c" + ) assert expired_non_expired == [] @@ -1222,7 +1256,9 @@ def test_get_expired_environments_nonexistent_name( ): """get_expired_environments with a name that does not exist returns an empty list without error.""" now_ts = now_timestamp() - result = state_sync.get_expired_environments(current_ts=now_ts, name="nonexistent_env") + result = state_sync.get_expired_environments( + current_ts=now_ts, name="nonexistent_env" + ) assert result == [] @@ -1252,7 +1288,9 @@ def test_delete_expired_environments_non_expired_name_is_not_deleted( ) state_sync.promote(env) - deleted = state_sync.delete_expired_environments(current_ts=now_ts, name="non_expired_env") + deleted = state_sync.delete_expired_environments( + current_ts=now_ts, name="non_expired_env" + ) assert deleted == [] assert state_sync.get_environment("non_expired_env") is not None @@ -1311,7 +1349,9 @@ def test_invalidate_environment_sync_does_not_delete_sibling_expired_envs( assert state_sync.get_environment("sibling_env") is not None -def test_delete_expired_snapshots(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_delete_expired_snapshots( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): now_ts = now_timestamp() snapshot = make_snapshot( @@ -1344,14 +1384,18 @@ def test_delete_expired_snapshots(state_sync: EngineAdapterStateSync, make_snaps assert _get_cleanup_tasks(state_sync) == [ SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True), - SnapshotTableCleanupTask(snapshot=new_snapshot.table_info, dev_table_only=False), + SnapshotTableCleanupTask( + snapshot=new_snapshot.table_info, dev_table_only=False + ), ] state_sync.delete_expired_snapshots(batch_range=ExpiredBatchRange.all_batch_range()) assert not state_sync.get_snapshots(all_snapshots) -def test_get_expired_snapshot_batch(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_get_expired_snapshot_batch( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): now_ts = now_timestamp() snapshots = [] @@ -1875,12 +1919,16 @@ def test_delete_expired_snapshots_previous_finalized_snapshots( # previous_finalized_snapshots in a non-finalized environment assert not _get_cleanup_tasks(state_sync) state_sync.delete_expired_snapshots(batch_range=ExpiredBatchRange.all_batch_range()) - assert state_sync.snapshots_exist([old_snapshot.snapshot_id]) == {old_snapshot.snapshot_id} + assert state_sync.snapshots_exist([old_snapshot.snapshot_id]) == { + old_snapshot.snapshot_id + } # Once the environment is finalized, the expired snapshot should be removed successfully state_sync.finalize(env) assert _get_cleanup_tasks(state_sync) == [ - SnapshotTableCleanupTask(snapshot=old_snapshot.table_info, dev_table_only=False), + SnapshotTableCleanupTask( + snapshot=old_snapshot.table_info, dev_table_only=False + ), ] state_sync.delete_expired_snapshots(batch_range=ExpiredBatchRange.all_batch_range()) assert not state_sync.snapshots_exist([old_snapshot.snapshot_id]) @@ -2009,7 +2057,9 @@ def test_delete_expired_snapshots_ignore_ttl( # default TTL = 1 week, nothing to clean up yet if we take TTL into account assert not _get_cleanup_tasks(state_sync) state_sync.delete_expired_snapshots(batch_range=ExpiredBatchRange.all_batch_range()) - assert state_sync.snapshots_exist([snapshot_c.snapshot_id]) == {snapshot_c.snapshot_id} + assert state_sync.snapshots_exist([snapshot_c.snapshot_id]) == { + snapshot_c.snapshot_id + } # If we ignore TTL, only snapshot_c should get cleaned up because snapshot_a and snapshot_b are part of an environment assert snapshot_a.table_info != snapshot_b.table_info != snapshot_c.table_info @@ -2076,7 +2126,9 @@ def test_delete_expired_snapshots_cleanup_intervals( ] # Check new snapshot's intervals - stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[new_snapshot.snapshot_id] + stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[ + new_snapshot.snapshot_id + ] assert stored_new_snapshot.intervals == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-06")), ] @@ -2084,7 +2136,9 @@ def test_delete_expired_snapshots_cleanup_intervals( assert _get_cleanup_tasks(state_sync) == [ SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True), - SnapshotTableCleanupTask(snapshot=new_snapshot.table_info, dev_table_only=False), + SnapshotTableCleanupTask( + snapshot=new_snapshot.table_info, dev_table_only=False + ), ] state_sync.delete_expired_snapshots(batch_range=ExpiredBatchRange.all_batch_range()) @@ -2129,7 +2183,9 @@ def test_delete_expired_snapshots_cleanup_intervals_shared_version( ) # Check new snapshot's intervals - stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[new_snapshot.snapshot_id] + stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[ + new_snapshot.snapshot_id + ] assert stored_new_snapshot.intervals == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-06")), ] @@ -2156,7 +2212,9 @@ def test_delete_expired_snapshots_cleanup_intervals_shared_version( version=snapshot.version, dev_version=snapshot.dev_version, intervals=[(to_timestamp("2023-01-01"), to_timestamp("2023-01-04"))], - dev_intervals=[(to_timestamp("2023-01-01"), to_timestamp("2023-01-04"))], + dev_intervals=[ + (to_timestamp("2023-01-01"), to_timestamp("2023-01-04")) + ], ), SnapshotIntervals( name='"a"', @@ -2177,7 +2235,9 @@ def test_delete_expired_snapshots_cleanup_intervals_shared_version( assert not state_sync.get_snapshots([snapshot]) # Check new snapshot's intervals - stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[new_snapshot.snapshot_id] + stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[ + new_snapshot.snapshot_id + ] assert stored_new_snapshot.intervals == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-06")), ] @@ -2247,7 +2307,9 @@ def test_delete_expired_snapshots_cleanup_intervals_shared_dev_version( ) # Check new snapshot's intervals - stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[new_snapshot.snapshot_id] + stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[ + new_snapshot.snapshot_id + ] assert stored_new_snapshot.intervals == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-04")), ] @@ -2276,14 +2338,18 @@ def test_delete_expired_snapshots_cleanup_intervals_shared_dev_version( version=snapshot.version, dev_version=snapshot.dev_version, intervals=[(to_timestamp("2023-01-01"), to_timestamp("2023-01-04"))], - dev_intervals=[(to_timestamp("2023-01-04"), to_timestamp("2023-01-08"))], + dev_intervals=[ + (to_timestamp("2023-01-04"), to_timestamp("2023-01-08")) + ], ), SnapshotIntervals( name='"a"', identifier=new_snapshot.identifier, version=snapshot.version, dev_version=new_snapshot.dev_version, - dev_intervals=[(to_timestamp("2023-01-08"), to_timestamp("2023-01-11"))], + dev_intervals=[ + (to_timestamp("2023-01-08"), to_timestamp("2023-01-11")) + ], ), ], key=compare_snapshot_intervals, @@ -2295,7 +2361,9 @@ def test_delete_expired_snapshots_cleanup_intervals_shared_dev_version( assert not state_sync.get_snapshots([snapshot]) # Check new snapshot's intervals - stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[new_snapshot.snapshot_id] + stored_new_snapshot = state_sync.get_snapshots([new_snapshot])[ + new_snapshot.snapshot_id + ] assert stored_new_snapshot.intervals == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-04")), ] @@ -2321,14 +2389,18 @@ def test_delete_expired_snapshots_cleanup_intervals_shared_dev_version( identifier=None, version=snapshot.version, dev_version=snapshot.dev_version, - dev_intervals=[(to_timestamp("2023-01-04"), to_timestamp("2023-01-08"))], + dev_intervals=[ + (to_timestamp("2023-01-04"), to_timestamp("2023-01-08")) + ], ), SnapshotIntervals( name='"a"', identifier=new_snapshot.identifier, version=snapshot.version, dev_version=new_snapshot.dev_version, - dev_intervals=[(to_timestamp("2023-01-08"), to_timestamp("2023-01-11"))], + dev_intervals=[ + (to_timestamp("2023-01-08"), to_timestamp("2023-01-11")) + ], ), ], key=compare_snapshot_intervals, @@ -2421,7 +2493,9 @@ def test_compact_intervals_after_cleanup( assert ( sorted( - state_sync.interval_state.get_snapshot_intervals([snapshot_a, snapshot_b, snapshot_c]), + state_sync.interval_state.get_snapshot_intervals( + [snapshot_a, snapshot_b, snapshot_c] + ), key=lambda x: (x.identifier or "", x.dev_version or ""), ) == expected_intervals @@ -2432,7 +2506,9 @@ def test_compact_intervals_after_cleanup( assert state_sync.engine_adapter.fetchone("SELECT COUNT(*) FROM sqlmesh._intervals")[0] == 4 # type: ignore assert ( sorted( - state_sync.interval_state.get_snapshot_intervals([snapshot_a, snapshot_b, snapshot_c]), + state_sync.interval_state.get_snapshot_intervals( + [snapshot_a, snapshot_b, snapshot_c] + ), key=lambda x: (x.identifier or "", x.dev_version or ""), ) == expected_intervals @@ -2467,10 +2543,14 @@ def test_environment_start_as_timestamp( stored_env = state_sync.get_environment(env.name) assert stored_env - assert stored_env.start_at == to_datetime(now_ts).replace(tzinfo=None).isoformat(sep=" ") + assert stored_env.start_at == to_datetime(now_ts).replace(tzinfo=None).isoformat( + sep=" " + ) -def test_unpause_snapshots(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_unpause_snapshots( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): snapshot = make_snapshot( SqlModel( name="test_snapshot", @@ -2503,13 +2583,17 @@ def test_unpause_snapshots(state_sync: EngineAdapterStateSync, make_snapshot: t. actual_snapshots = state_sync.get_snapshots([snapshot, new_snapshot]) assert not actual_snapshots[snapshot.snapshot_id].unpaused_ts - assert actual_snapshots[new_snapshot.snapshot_id].unpaused_ts == to_timestamp(unpaused_dt) + assert actual_snapshots[new_snapshot.snapshot_id].unpaused_ts == to_timestamp( + unpaused_dt + ) assert actual_snapshots[snapshot.snapshot_id].unrestorable assert not actual_snapshots[new_snapshot.snapshot_id].unrestorable -def test_unrestorable_snapshot(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_unrestorable_snapshot( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): snapshot = make_snapshot( SqlModel( name="test_snapshot", @@ -2533,26 +2617,34 @@ def test_unrestorable_snapshot(state_sync: EngineAdapterStateSync, make_snapshot new_indirect_non_breaking_snapshot = make_snapshot( SqlModel(name="test_snapshot", query=parse_one("select 2, ds"), cron="@daily") ) - new_indirect_non_breaking_snapshot.categorize_as(SnapshotChangeCategory.INDIRECT_NON_BREAKING) + new_indirect_non_breaking_snapshot.categorize_as( + SnapshotChangeCategory.INDIRECT_NON_BREAKING + ) new_indirect_non_breaking_snapshot.version = "a" assert not new_indirect_non_breaking_snapshot.unpaused_ts state_sync.push_snapshots([new_indirect_non_breaking_snapshot]) state_sync.unpause_snapshots([new_indirect_non_breaking_snapshot], unpaused_dt) - actual_snapshots = state_sync.get_snapshots([snapshot, new_indirect_non_breaking_snapshot]) + actual_snapshots = state_sync.get_snapshots( + [snapshot, new_indirect_non_breaking_snapshot] + ) assert not actual_snapshots[snapshot.snapshot_id].unpaused_ts assert actual_snapshots[ new_indirect_non_breaking_snapshot.snapshot_id ].unpaused_ts == to_timestamp(unpaused_dt) assert not actual_snapshots[snapshot.snapshot_id].unrestorable - assert not actual_snapshots[new_indirect_non_breaking_snapshot.snapshot_id].unrestorable + assert not actual_snapshots[ + new_indirect_non_breaking_snapshot.snapshot_id + ].unrestorable new_forward_only_snapshot = make_snapshot( SqlModel(name="test_snapshot", query=parse_one("select 3, ds"), cron="@daily") ) - new_forward_only_snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) + new_forward_only_snapshot.categorize_as( + SnapshotChangeCategory.BREAKING, forward_only=True + ) new_forward_only_snapshot.version = "a" assert not new_forward_only_snapshot.unpaused_ts @@ -2563,10 +2655,12 @@ def test_unrestorable_snapshot(state_sync: EngineAdapterStateSync, make_snapshot [snapshot, new_indirect_non_breaking_snapshot, new_forward_only_snapshot] ) assert not actual_snapshots[snapshot.snapshot_id].unpaused_ts - assert not actual_snapshots[new_indirect_non_breaking_snapshot.snapshot_id].unpaused_ts - assert actual_snapshots[new_forward_only_snapshot.snapshot_id].unpaused_ts == to_timestamp( - unpaused_dt - ) + assert not actual_snapshots[ + new_indirect_non_breaking_snapshot.snapshot_id + ].unpaused_ts + assert actual_snapshots[ + new_forward_only_snapshot.snapshot_id + ].unpaused_ts == to_timestamp(unpaused_dt) assert actual_snapshots[snapshot.snapshot_id].unrestorable assert actual_snapshots[new_indirect_non_breaking_snapshot.snapshot_id].unrestorable @@ -2608,7 +2702,9 @@ def test_unrestorable_snapshot_target_not_forward_only( actual_snapshots = state_sync.get_snapshots([snapshot, updated_snapshot]) assert not actual_snapshots[snapshot.snapshot_id].unpaused_ts - assert actual_snapshots[updated_snapshot.snapshot_id].unpaused_ts == to_timestamp(unpaused_dt) + assert actual_snapshots[updated_snapshot.snapshot_id].unpaused_ts == to_timestamp( + unpaused_dt + ) assert actual_snapshots[snapshot.snapshot_id].unrestorable assert not actual_snapshots[updated_snapshot.snapshot_id].unrestorable @@ -2691,7 +2787,9 @@ def test_version_sqlmesh(state_sync: EngineAdapterStateSync) -> None: # sqlmesh version is ahead sqlmesh_version_minor_decrease = f"{major}.{int(minor) - 1}.{patch}" error = rf"SQLMesh \(local\) is using version '{re.escape(SQLMESH_VERSION)}' which is ahead of '{sqlmesh_version_minor_decrease}'" - state_sync.version_state.update_versions(sqlmesh_version=sqlmesh_version_minor_decrease) + state_sync.version_state.update_versions( + sqlmesh_version=sqlmesh_version_minor_decrease + ) with pytest.raises(SQLMeshError, match=error): state_sync.get_versions() state_sync.get_versions(validate=False) @@ -2734,7 +2832,9 @@ def test_empty_versions() -> None: assert empty_versions.sqlmesh_version == "0.0.0" -def test_migrate(state_sync: EngineAdapterStateSync, mocker: MockerFixture, tmp_path) -> None: +def test_migrate( + state_sync: EngineAdapterStateSync, mocker: MockerFixture, tmp_path +) -> None: from sqlmesh import __version__ as SQLMESH_VERSION migrate_rows_mock = mocker.patch( @@ -2766,7 +2866,9 @@ def test_migrate(state_sync: EngineAdapterStateSync, mocker: MockerFixture, tmp_ assert ( state_sync.engine_adapter.fetchone( "SELECT COUNT(*) FROM sqlmesh._snapshots WHERE ttl_ms IS NULL" - )[0] # type: ignore + )[ + 0 + ] # type: ignore == 0 ) @@ -2795,9 +2897,15 @@ def test_rollback(state_sync: EngineAdapterStateSync, mocker: MockerFixture) -> f"{state_sync.schema}._versions", f"{state_sync.schema}._versions_backup", ) in calls - assert not state_sync.engine_adapter.table_exists(f"{state_sync.schema}._snapshots_backup") - assert not state_sync.engine_adapter.table_exists(f"{state_sync.schema}._environments_backup") - assert not state_sync.engine_adapter.table_exists(f"{state_sync.schema}._versions_backup") + assert not state_sync.engine_adapter.table_exists( + f"{state_sync.schema}._snapshots_backup" + ) + assert not state_sync.engine_adapter.table_exists( + f"{state_sync.schema}._environments_backup" + ) + assert not state_sync.engine_adapter.table_exists( + f"{state_sync.schema}._versions_backup" + ) def test_first_migration_failure(duck_conn, mocker: MockerFixture, tmp_path) -> None: @@ -2806,21 +2914,31 @@ def test_first_migration_failure(duck_conn, mocker: MockerFixture, tmp_path) -> schema=c.SQLMESH, cache_dir=tmp_path / c.CACHE, ) - mocker.patch.object(state_sync.migrator, "_migrate_rows", side_effect=Exception("mocked error")) + mocker.patch.object( + state_sync.migrator, "_migrate_rows", side_effect=Exception("mocked error") + ) with pytest.raises( SQLMeshError, match="SQLMesh migration failed.", ): state_sync.migrate() - assert not state_sync.engine_adapter.table_exists(state_sync.snapshot_state.snapshots_table) + assert not state_sync.engine_adapter.table_exists( + state_sync.snapshot_state.snapshots_table + ) assert not state_sync.engine_adapter.table_exists( state_sync.environment_state.environments_table ) - assert not state_sync.engine_adapter.table_exists(state_sync.version_state.versions_table) - assert not state_sync.engine_adapter.table_exists(state_sync.interval_state.intervals_table) + assert not state_sync.engine_adapter.table_exists( + state_sync.version_state.versions_table + ) + assert not state_sync.engine_adapter.table_exists( + state_sync.interval_state.intervals_table + ) -def test_migrate_rows(state_sync: EngineAdapterStateSync, mocker: MockerFixture) -> None: +def test_migrate_rows( + state_sync: EngineAdapterStateSync, mocker: MockerFixture +) -> None: state_sync.engine_adapter.replace_query( "sqlmesh._versions", pd.read_json("tests/fixtures/migrations/versions.json"), @@ -2885,13 +3003,21 @@ def test_migrate_rows(state_sync: EngineAdapterStateSync, mocker: MockerFixture) }, ) - old_snapshots = state_sync.engine_adapter.fetchdf("select * from sqlmesh._snapshots") - old_environments = state_sync.engine_adapter.fetchdf("select * from sqlmesh._environments") + old_snapshots = state_sync.engine_adapter.fetchdf( + "select * from sqlmesh._snapshots" + ) + old_environments = state_sync.engine_adapter.fetchdf( + "select * from sqlmesh._environments" + ) state_sync.migrate(skip_backup=True) - new_snapshots = state_sync.engine_adapter.fetchdf("select * from sqlmesh._snapshots") - new_environments = state_sync.engine_adapter.fetchdf("select * from sqlmesh._environments") + new_snapshots = state_sync.engine_adapter.fetchdf( + "select * from sqlmesh._snapshots" + ) + new_environments = state_sync.engine_adapter.fetchdf( + "select * from sqlmesh._environments" + ) assert len(old_snapshots) == 24 assert len(new_snapshots) == 36 @@ -2917,7 +3043,9 @@ def test_migrate_rows(state_sync: EngineAdapterStateSync, mocker: MockerFixture) assert not missing_intervals(dev_snapshots, start=start, end=end) - assert not missing_intervals(dev_snapshots, start="2023-01-08", end="2023-01-10") == 8 + assert ( + not missing_intervals(dev_snapshots, start="2023-01-08", end="2023-01-10") == 8 + ) all_snapshot_ids = [ SnapshotId(name=name, identifier=identifier) @@ -2937,7 +3065,9 @@ def test_migrate_rows(state_sync: EngineAdapterStateSync, mocker: MockerFixture) ) -def test_backup_state(state_sync: EngineAdapterStateSync, mocker: MockerFixture) -> None: +def test_backup_state( + state_sync: EngineAdapterStateSync, mocker: MockerFixture +) -> None: state_sync.engine_adapter.replace_query( "sqlmesh._snapshots", pd.read_json("tests/fixtures/migrations/snapshots.json"), @@ -2969,7 +3099,9 @@ def test_restore_snapshots_table(state_sync: EngineAdapterStateSync) -> None: target_columns_to_types=snapshot_columns_to_types, ) - old_snapshots = state_sync.engine_adapter.fetchdf("select * from sqlmesh._snapshots") + old_snapshots = state_sync.engine_adapter.fetchdf( + "select * from sqlmesh._snapshots" + ) old_snapshots_count = state_sync.engine_adapter.fetchone( "select count(*) from sqlmesh._snapshots" ) @@ -2977,14 +3109,18 @@ def test_restore_snapshots_table(state_sync: EngineAdapterStateSync) -> None: state_sync.migrator._backup_state() state_sync.engine_adapter.delete_from("sqlmesh._snapshots", "TRUE") - snapshots_count = state_sync.engine_adapter.fetchone("select count(*) from sqlmesh._snapshots") + snapshots_count = state_sync.engine_adapter.fetchone( + "select count(*) from sqlmesh._snapshots" + ) assert snapshots_count == (0,) state_sync.migrator._restore_table( table_name="sqlmesh._snapshots", backup_table_name="sqlmesh._snapshots_backup", ) - new_snapshots = state_sync.engine_adapter.fetchdf("select * from sqlmesh._snapshots") + new_snapshots = state_sync.engine_adapter.fetchdf( + "select * from sqlmesh._snapshots" + ) pd.testing.assert_frame_equal( old_snapshots, new_snapshots, @@ -3012,7 +3148,9 @@ def test_seed_hydration( assert snapshot.model.seed.content == "header\n1\n2" state_sync.snapshot_state.clear_cache() - stored_snapshot = state_sync.get_snapshots([snapshot.snapshot_id])[snapshot.snapshot_id] + stored_snapshot = state_sync.get_snapshots([snapshot.snapshot_id])[ + snapshot.snapshot_id + ] assert isinstance(stored_snapshot.model, SeedModel) assert not stored_snapshot.model.is_hydrated assert stored_snapshot.model.seed.content == "" @@ -3035,7 +3173,9 @@ def test_nodes_exist(state_sync: EngineAdapterStateSync, make_snapshot: t.Callab assert state_sync.nodes_exist([snapshot.name]) == {snapshot.name} -def test_invalidate_environment(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): +def test_invalidate_environment( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): snapshot = make_snapshot( SqlModel( name="a", @@ -3074,14 +3214,18 @@ def test_invalidate_environment(state_sync: EngineAdapterStateSync, make_snapsho stored_env = state_sync.get_environment("test_environment") assert stored_env - assert stored_env.expiration_ts and stored_env.expiration_ts < original_expiration_ts + assert ( + stored_env.expiration_ts and stored_env.expiration_ts < original_expiration_ts + ) deleted_environments = state_sync.delete_expired_environments() assert len(deleted_environments) == 1 assert deleted_environments[0].name == "test_environment" assert state_sync.get_environment_statements(env.name) == [] - with pytest.raises(SQLMeshError, match="Cannot invalidate the production environment."): + with pytest.raises( + SQLMeshError, match="Cannot invalidate the production environment." + ): state_sync.invalidate_environment("prod") @@ -3154,11 +3298,15 @@ def test_cache(state_sync, make_snapshot, mocker): # prime the cache with a real snapshot cache.push_snapshots([snapshot]) - assert cache.get_snapshots([snapshot.snapshot_id]) == {snapshot.snapshot_id: snapshot} + assert cache.get_snapshots([snapshot.snapshot_id]) == { + snapshot.snapshot_id: snapshot + } # cache hit with patch.object(state_sync, "get_snapshots") as mock: - assert cache.get_snapshots([snapshot.snapshot_id]) == {snapshot.snapshot_id: snapshot} + assert cache.get_snapshots([snapshot.snapshot_id]) == { + snapshot.snapshot_id: snapshot + } mock.assert_not_called() # clear the cache by adding intervals @@ -3168,10 +3316,14 @@ def test_cache(state_sync, make_snapshot, mocker): mock.assert_called() # clear the cache by removing intervals - cache.remove_intervals([(snapshot, snapshot.inclusive_exclusive("2020-01-01", "2020-01-01"))]) + cache.remove_intervals( + [(snapshot, snapshot.inclusive_exclusive("2020-01-01", "2020-01-01"))] + ) # prime the cache - assert cache.get_snapshots([snapshot.snapshot_id]) == {snapshot.snapshot_id: snapshot} + assert cache.get_snapshots([snapshot.snapshot_id]) == { + snapshot.snapshot_id: snapshot + } # cache hit half way now_timestamp.return_value = to_timestamp("2023-01-01 00:00:05") @@ -3218,7 +3370,9 @@ def test_max_interval_end_per_model( environment_name = "test_max_interval_end_for_environment" assert state_sync.max_interval_end_per_model(environment_name) == {} - assert state_sync.max_interval_end_per_model(environment_name, {snapshot_a.name}) == {} + assert ( + state_sync.max_interval_end_per_model(environment_name, {snapshot_a.name}) == {} + ) state_sync.promote( Environment( @@ -3231,13 +3385,13 @@ def test_max_interval_end_per_model( ) ) - assert state_sync.max_interval_end_per_model(environment_name, {snapshot_a.name}) == { - snapshot_a.name: to_timestamp("2023-01-04") - } + assert state_sync.max_interval_end_per_model( + environment_name, {snapshot_a.name} + ) == {snapshot_a.name: to_timestamp("2023-01-04")} - assert state_sync.max_interval_end_per_model(environment_name, {snapshot_b.name}) == { - snapshot_b.name: to_timestamp("2023-01-03") - } + assert state_sync.max_interval_end_per_model( + environment_name, {snapshot_b.name} + ) == {snapshot_b.name: to_timestamp("2023-01-03")} assert state_sync.max_interval_end_per_model( environment_name, {snapshot_a.name, snapshot_b.name} @@ -3290,7 +3444,9 @@ def test_max_interval_end_per_model_with_pending_restatements( ) snapshot = state_sync.get_snapshots([snapshot.snapshot_id])[snapshot.snapshot_id] - assert snapshot.intervals == [(to_timestamp("2023-01-01"), to_timestamp("2023-01-04"))] + assert snapshot.intervals == [ + (to_timestamp("2023-01-01"), to_timestamp("2023-01-04")) + ] assert snapshot.pending_restatement_intervals == [ (to_timestamp("2023-01-04"), to_timestamp("2023-01-05")) ] @@ -3345,7 +3501,9 @@ def test_max_interval_end_per_model_ensure_finalized_snapshots( environment_name = "test_max_interval_end_for_environment" assert state_sync.max_interval_end_per_model(environment_name) == {} - assert state_sync.max_interval_end_per_model(environment_name, {snapshot_a.name}) == {} + assert ( + state_sync.max_interval_end_per_model(environment_name, {snapshot_a.name}) == {} + ) state_sync.promote( Environment( @@ -3370,7 +3528,9 @@ def test_max_interval_end_per_model_ensure_finalized_snapshots( ) == {snapshot_b.name: to_timestamp("2023-01-03")} assert state_sync.max_interval_end_per_model( - environment_name, {snapshot_a.name, snapshot_b.name}, ensure_finalized_snapshots=True + environment_name, + {snapshot_a.name, snapshot_b.name}, + ensure_finalized_snapshots=True, ) == {snapshot_b.name: to_timestamp("2023-01-03")} assert state_sync.max_interval_end_per_model( @@ -3411,7 +3571,9 @@ def test_snapshot_batching(state_sync, mocker, make_snapshot): ) ) calls = mock.delete_from.call_args_list - identifiers = sorted([snapshot_a.identifier, snapshot_b.identifier, snapshot_c.identifier]) + identifiers = sorted( + [snapshot_a.identifier, snapshot_b.identifier, snapshot_c.identifier] + ) assert mock.delete_from.call_args_list == [ call( exp.to_table("sqlmesh._snapshots"), @@ -3518,11 +3680,16 @@ def test_snapshot_cache( state_sync.snapshot_state._snapshot_cache = cache_mock snapshot = make_snapshot(SqlModel(name="a", query=parse_one("select 1"))) - cache_mock.get_or_load.return_value = ({snapshot.snapshot_id: snapshot}, {snapshot.snapshot_id}) + cache_mock.get_or_load.return_value = ( + {snapshot.snapshot_id: snapshot}, + {snapshot.snapshot_id}, + ) state_sync.snapshot_state.push_snapshots([snapshot]) - assert state_sync.get_snapshots([snapshot.snapshot_id]) == {snapshot.snapshot_id: snapshot} + assert state_sync.get_snapshots([snapshot.snapshot_id]) == { + snapshot.snapshot_id: snapshot + } cache_mock.get_or_load.assert_called_once_with({snapshot.snapshot_id}, mocker.ANY) # Update the snapshot in the state and make sure this update is reflected on the cached instance. @@ -3531,7 +3698,9 @@ def test_snapshot_cache( state_sync.snapshot_state._update_snapshots( [snapshot.snapshot_id], unpaused_ts=1, unrestorable=True ) - new_snapshot = state_sync.get_snapshots([snapshot.snapshot_id])[snapshot.snapshot_id] + new_snapshot = state_sync.get_snapshots([snapshot.snapshot_id])[ + snapshot.snapshot_id + ] assert new_snapshot.unpaused_ts == 1 assert new_snapshot.unrestorable @@ -3540,10 +3709,18 @@ def test_snapshot_cache( assert state_sync.get_snapshots([snapshot.snapshot_id]) == {} -def test_update_auto_restatements(state_sync: EngineAdapterStateSync, make_snapshot: t.Callable): - snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 1")), version="1") - snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select 2")), version="2") - snapshot_c = make_snapshot(SqlModel(name="c", query=parse_one("select 3")), version="3") +def test_update_auto_restatements( + state_sync: EngineAdapterStateSync, make_snapshot: t.Callable +): + snapshot_a = make_snapshot( + SqlModel(name="a", query=parse_one("select 1")), version="1" + ) + snapshot_b = make_snapshot( + SqlModel(name="b", query=parse_one("select 2")), version="2" + ) + snapshot_c = make_snapshot( + SqlModel(name="c", query=parse_one("select 3")), version="3" + ) state_sync.snapshot_state.push_snapshots([snapshot_a, snapshot_b, snapshot_c]) @@ -3751,14 +3928,18 @@ def test_compact_intervals_pending_restatement_shared_version( pending_restatement_intervals=[], ), ] - expected_intervals = sorted(expected_intervals, key=lambda x: (x.name, x.identifier or "")) + expected_intervals = sorted( + expected_intervals, key=lambda x: (x.name, x.identifier or "") + ) with time_machine.travel("2020-01-05 01:00:00 UTC"): # Add a new interval for the new snapshot state_sync.add_interval(snapshot_b, "2020-01-03", "2020-01-03") assert ( sorted( - state_sync.interval_state.get_snapshot_intervals([snapshot_a, snapshot_b]), + state_sync.interval_state.get_snapshot_intervals( + [snapshot_a, snapshot_b] + ), key=lambda x: (x.name, x.identifier or ""), ) == expected_intervals @@ -3836,14 +4017,18 @@ def test_compact_intervals_pending_restatement_shared_version( pending_restatement_intervals=[], ), ] - expected_intervals = sorted(expected_intervals, key=lambda x: (x.name, x.identifier or "")) + expected_intervals = sorted( + expected_intervals, key=lambda x: (x.name, x.identifier or "") + ) with time_machine.travel("2020-01-05 02:00:00 UTC"): # Add a new interval for the previous snapshot state_sync.add_interval(snapshot_a, "2020-01-04", "2020-01-04") assert ( sorted( - state_sync.interval_state.get_snapshot_intervals([snapshot_a, snapshot_b]), + state_sync.interval_state.get_snapshot_intervals( + [snapshot_a, snapshot_b] + ), key=lambda x: (x.name, x.identifier or ""), ) == expected_intervals @@ -3896,13 +4081,17 @@ def test_compact_intervals_pending_restatement_shared_version( pending_restatement_intervals=[], ), ] - expected_intervals = sorted(expected_intervals, key=lambda x: (x.name, x.identifier or "")) + expected_intervals = sorted( + expected_intervals, key=lambda x: (x.name, x.identifier or "") + ) with time_machine.travel("2020-01-05 03:00:00 UTC"): state_sync.add_interval(snapshot_b, "2020-01-05", "2020-01-05") assert ( sorted( - state_sync.interval_state.get_snapshot_intervals([snapshot_a, snapshot_b]), + state_sync.interval_state.get_snapshot_intervals( + [snapshot_a, snapshot_b] + ), key=lambda x: (x.name, x.identifier or ""), ) == expected_intervals @@ -3952,7 +4141,9 @@ def test_get_environments_summary( state_sync.promote(env_a) env_b_ttl = now_ts + 1000 - env_b = env_a.copy(update={"name": "test_environment_b", "expiration_ts": env_b_ttl}) + env_b = env_a.copy( + update={"name": "test_environment_b", "expiration_ts": env_b_ttl} + ) state_sync.promote(env_b) prod = Environment( @@ -4079,7 +4270,9 @@ def test_update_environment_statements(state_sync: EngineAdapterStateSync): environment.name, environment.plan_id, environment_statements ) - environment_statements_dev = state_sync.get_environment_statements(environment="dev") + environment_statements_dev = state_sync.get_environment_statements( + environment="dev" + ) assert environment_statements_dev[0].before_all == [ "CREATE OR REPLACE TABLE table_1 AS SELECT 'a'" ] @@ -4103,7 +4296,9 @@ def test_update_environment_statements(state_sync: EngineAdapterStateSync): environment.name, environment.plan_id, environment_statements ) - environment_statements_dev = state_sync.get_environment_statements(environment="dev") + environment_statements_dev = state_sync.get_environment_statements( + environment="dev" + ) assert environment_statements_dev[0].before_all == [ "CREATE OR REPLACE TABLE table_1 AS SELECT 'a'" ] @@ -4139,12 +4334,15 @@ def test_get_snapshots_by_names( state_sync.push_snapshots([snap_a_v1, snap_a_v2, snap_b]) - assert {s.snapshot_id for s in state_sync.get_snapshots_by_names(snapshot_names=['"a"'])} == { + assert { + s.snapshot_id for s in state_sync.get_snapshots_by_names(snapshot_names=['"a"']) + } == { snap_a_v1.snapshot_id, snap_a_v2.snapshot_id, } assert { - s.snapshot_id for s in state_sync.get_snapshots_by_names(snapshot_names=['"a"', '"b"']) + s.snapshot_id + for s in state_sync.get_snapshots_by_names(snapshot_names=['"a"', '"b"']) } == { snap_a_v1.snapshot_id, snap_a_v2.snapshot_id, @@ -4181,11 +4379,15 @@ def test_get_snapshots_by_names_include_expired( assert { s.snapshot_id - for s in state_sync.get_snapshots_by_names(snapshot_names=['"a"'], current_ts=now_ts) + for s in state_sync.get_snapshots_by_names( + snapshot_names=['"a"'], current_ts=now_ts + ) } == {normal_a.snapshot_id} assert { s.snapshot_id - for s in state_sync.get_snapshots_by_names(snapshot_names=['"a"'], exclude_expired=False) + for s in state_sync.get_snapshots_by_names( + snapshot_names=['"a"'], exclude_expired=False + ) } == { normal_a.snapshot_id, expired_a.snapshot_id, @@ -4206,7 +4408,13 @@ def test_state_version_is_too_old( state_sync.engine_adapter.replace_query( "sqlmesh._versions", pd.DataFrame( - [{"schema_version": 59, "sqlmesh_version": "0.133.0", "sqlglot_version": "25.31.4"}] + [ + { + "schema_version": 59, + "sqlmesh_version": "0.133.0", + "sqlglot_version": "25.31.4", + } + ] ), target_columns_to_types={ "schema_version": exp.DataType.build("int"), diff --git a/tests/core/test_audit.py b/tests/core/test_audit.py index b5226563b5..1c4805d031 100644 --- a/tests/core/test_audit.py +++ b/tests/core/test_audit.py @@ -1,27 +1,18 @@ import json + import pytest from sqlglot import exp, parse_one from sqlmesh.core import constants as c +from sqlmesh.core.audit import (ModelAudit, StandaloneAudit, builtin, + load_audit, load_multiple_audits) from sqlmesh.core.config.model import ModelDefaultsConfig from sqlmesh.core.context import Context +from sqlmesh.core.dialect import jinja_query, parse +from sqlmesh.core.model import (FullKind, IncrementalByTimeRangeKind, Model, + SeedModel, create_sql_model, + load_sql_based_model) from sqlmesh.core.node import DbtNodeInfo -from sqlmesh.core.audit import ( - ModelAudit, - StandaloneAudit, - builtin, - load_audit, - load_multiple_audits, -) -from sqlmesh.core.dialect import parse, jinja_query -from sqlmesh.core.model import ( - FullKind, - IncrementalByTimeRangeKind, - Model, - SeedModel, - create_sql_model, - load_sql_based_model, -) from sqlmesh.utils.errors import AuditConfigError from sqlmesh.utils.jinja import JinjaMacroRegistry, MacroExtractor from sqlmesh.utils.metaprogramming import Executable @@ -47,8 +38,7 @@ def model_default_catalog() -> Model: def test_load(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" -- Audit comment Audit ( name my_audit, @@ -62,8 +52,7 @@ def test_load(assert_exp_eq): db.table WHERE col IS NULL - """ - ) + """) audit = load_audit(expressions, path="/path/to/audit", dialect="duckdb") assert isinstance(audit, ModelAudit) @@ -86,8 +75,7 @@ def test_load(assert_exp_eq): def test_load_standalone(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" Audit ( name my_audit, dialect spark, @@ -103,8 +91,7 @@ def test_load_standalone(assert_exp_eq): db.table WHERE col IS NULL - """ - ) + """) audit = load_audit(expressions, path="/path/to/audit", dialect="duckdb") assert isinstance(audit, StandaloneAudit) @@ -129,8 +116,7 @@ def test_load_standalone(assert_exp_eq): def test_load_standalone_default_catalog(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" Audit ( name my_audit, dialect spark, @@ -146,11 +132,13 @@ def test_load_standalone_default_catalog(assert_exp_eq): db.table WHERE col IS NULL - """ - ) + """) audit = load_audit( - expressions, path="/path/to/audit", dialect="duckdb", default_catalog="test_catalog" + expressions, + path="/path/to/audit", + dialect="duckdb", + default_catalog="test_catalog", ) assert isinstance(audit, StandaloneAudit) assert audit.dialect == "spark" @@ -186,8 +174,7 @@ def test_load_standalone_default_catalog(assert_exp_eq): def test_load_standalone_with_macros(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, owner owner_name, @@ -206,14 +193,17 @@ def test_load_standalone_with_macros(assert_exp_eq): db.table t1 WHERE col IS NULL - """ - ) + """) audit = load_audit( expressions, macros={ - "test_macro": Executable(payload="def test_macro(evaluator, v):\n return v"), - "extra_macro": Executable(payload="def extra_macro(evaluator, v):\n return v + 1"), + "test_macro": Executable( + payload="def test_macro(evaluator, v):\n return v" + ), + "extra_macro": Executable( + payload="def extra_macro(evaluator, v):\n return v + 1" + ), }, ) @@ -222,8 +212,7 @@ def test_load_standalone_with_macros(assert_exp_eq): def test_load_standalone_with_jinja_macros(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, owner owner_name, @@ -240,8 +229,7 @@ def test_load_standalone_with_jinja_macros(assert_exp_eq): WHERE col IS NULL JINJA_QUERY_END; - """ - ) + """) macros = """ {% macro test_macro(v) %}{{ v }}{% endmacro %} @@ -261,8 +249,7 @@ def test_load_standalone_with_jinja_macros(assert_exp_eq): def test_load_multiple(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" Audit ( name first_audit, dialect spark, @@ -281,8 +268,7 @@ def test_load_multiple(assert_exp_eq): SELECT * FROM db.table WHERE col2 IS NULL; - """ - ) + """) first_audit, second_audit = load_multiple_audits(expressions, path="/path/to/audit") assert first_audit.dialect == "spark" @@ -311,8 +297,7 @@ def test_load_multiple(assert_exp_eq): def test_load_with_dictionary_defaults(): - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, dialect spark, @@ -323,8 +308,7 @@ def test_load_with_dictionary_defaults(): ); SELECT 1 - """ - ) + """) audit = load_audit(expressions, dialect="spark") assert audit.defaults.keys() == {"field1", "field2"} @@ -334,8 +318,7 @@ def test_load_with_dictionary_defaults(): def test_load_with_single_defaults(): # testing it also works with a single default with no trailing comma - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, defaults ( @@ -344,8 +327,7 @@ def test_load_with_single_defaults(): ); SELECT 1 - """ - ) + """) audit = load_audit(expressions, dialect="duckdb") assert audit.defaults.keys() == {"field1"} @@ -354,26 +336,22 @@ def test_load_with_single_defaults(): def test_no_audit_statement(): - expressions = parse( - """ + expressions = parse(""" SELECT 1 - """ - ) + """) with pytest.raises(AuditConfigError) as ex: load_audit(expressions, path="/path/to/audit", dialect="duckdb") assert "Incomplete audit definition" in str(ex.value) def test_unordered_audit_statements(): - expressions = parse( - """ + expressions = parse(""" SELECT 1; AUDIT ( name my_audit, ); - """ - ) + """) with pytest.raises(AuditConfigError) as ex: load_audit(expressions, path="/path/to/audit", dialect="duckdb") @@ -381,15 +359,13 @@ def test_unordered_audit_statements(): def test_no_query(): - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, ); @DEF(x, 1) - """ - ) + """) with pytest.raises(AuditConfigError) as ex: load_audit(expressions, path="/path/to/audit", dialect="duckdb") @@ -438,8 +414,7 @@ def test_resolve_template(model_default_catalog: Model): def test_load_with_defaults(model, assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" Audit ( name my_audit, defaults ( @@ -458,8 +433,7 @@ def test_load_with_defaults(model, assert_exp_eq): AND @IF(@field4 = 'overridden', @field4 IN ('some string', 'other string'), 1=1) AND @field1 = @field2 AND @field3 != @field4 - """ - ) + """) audit = load_audit(expressions, path="/path/to/audit", dialect="duckdb") assert audit.defaults == { "field1": exp.to_column("some_column"), @@ -515,7 +489,9 @@ def test_not_null_audit_default_catalog(model_default_catalog: Model): def test_unique_values_audit(model: Model): rendered_query_a = model.render_audit_query( - builtin.unique_values_audit, columns=[exp.to_column("a")], condition=parse_one("b IS NULL") + builtin.unique_values_audit, + columns=[exp.to_column("a")], + condition=parse_one("b IS NULL"), ) assert ( rendered_query_a.sql() @@ -593,7 +569,10 @@ def test_accepted_range_audit(model: Model): == 'SELECT * FROM (SELECT * FROM "db"."test_model" AS "test_model" WHERE "ds" BETWEEN \'1970-01-01\' AND \'1970-01-01\') AS "_0" WHERE "a" < 0 AND TRUE' ) rendered_query = model.render_audit_query( - builtin.accepted_range_audit, column=exp.to_column("a"), max_v=100, inclusive=exp.false() + builtin.accepted_range_audit, + column=exp.to_column("a"), + max_v=100, + inclusive=exp.false(), ) assert ( rendered_query.sql() @@ -699,14 +678,16 @@ def test_pattern_audits(model: Model): def test_standalone_audit(model: Model, assert_exp_eq): audit = StandaloneAudit( - name="test_audit", query=parse_one(f"SELECT * FROM {model.name} WHERE col IS NULL") + name="test_audit", + query=parse_one(f"SELECT * FROM {model.name} WHERE col IS NULL"), ) assert audit.depends_on == {model.fqn} rendered_query = audit.render_audit_query() assert_exp_eq( - rendered_query, """SELECT * FROM "db"."test_model" AS "test_model" WHERE "col" IS NULL""" + rendered_query, + """SELECT * FROM "db"."test_model" AS "test_model" WHERE "col" IS NULL""", ) with pytest.raises(AuditConfigError) as ex: @@ -716,8 +697,7 @@ def test_standalone_audit(model: Model, assert_exp_eq): def test_render_definition(): - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, dialect spark, @@ -736,12 +716,15 @@ def test_render_definition(): db.table t1 WHERE col IS NULL - """ - ) + """) audit = load_audit( expressions, - macros={"test_macro": Executable(payload="def test_macro(evaluator, v):\n return v")}, + macros={ + "test_macro": Executable( + payload="def test_macro(evaluator, v):\n return v" + ) + }, ) from sqlmesh.core.dialect import format_model_expressions @@ -752,7 +735,9 @@ def test_render_definition(): ) == format_model_expressions(expressions) # Should include the macro implementation. - assert "def test_macro(evaluator, v):" in format_model_expressions(audit.render_definition()) + assert "def test_macro(evaluator, v):" in format_model_expressions( + audit.render_definition() + ) def test_render_definition_dbt_node_info(): @@ -760,11 +745,11 @@ def test_render_definition_dbt_node_info(): unique_id="test.project.my_audit", name="my_audit", fqn="project.my_audit" ) - audit = StandaloneAudit(name="my_audit", dbt_node_info=node_info, query=jinja_query("select 1")) + audit = StandaloneAudit( + name="my_audit", dbt_node_info=node_info, query=jinja_query("select 1") + ) - assert ( - audit.render_definition()[0].sql(pretty=True) - == """AUDIT ( + assert audit.render_definition()[0].sql(pretty=True) == """AUDIT ( name my_audit, dbt_node_info ( fqn := 'project.my_audit', @@ -773,12 +758,10 @@ def test_render_definition_dbt_node_info(): ), standalone TRUE )""" - ) def test_text_diff(): - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, dialect spark, @@ -793,12 +776,15 @@ def test_text_diff(): db.table t1 WHERE col IS NULL - """ - ) + """) audit = load_audit( expressions, - macros={"test_macro": Executable(payload="def test_macro(evaluator, v):\n return v")}, + macros={ + "test_macro": Executable( + payload="def test_macro(evaluator, v):\n return v" + ) + }, ) modified_audit = audit.copy() @@ -826,7 +812,10 @@ def test_non_blocking_builtin(): assert BUILT_IN_AUDITS["not_null_non_blocking"].blocking is False assert BUILT_IN_AUDITS["not_null_non_blocking"].name == "not_null_non_blocking" - assert BUILT_IN_AUDITS["not_null"].query == BUILT_IN_AUDITS["not_null_non_blocking"].query + assert ( + BUILT_IN_AUDITS["not_null"].query + == BUILT_IN_AUDITS["not_null_non_blocking"].query + ) def test_string_length_between_audit(model: Model): @@ -844,7 +833,9 @@ def test_string_length_between_audit(model: Model): def test_not_constant_audit(model: Model): rendered_query = model.render_audit_query( - builtin.not_constant_audit, column=exp.column("x"), condition=exp.condition("x > 1") + builtin.not_constant_audit, + column=exp.column("x"), + condition=exp.condition("x > 1"), ) assert ( rendered_query.sql() @@ -865,8 +856,7 @@ def test_condition_with_macro_var(model: Model): def test_variables(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" Audit ( name my_audit, dialect bigquery, @@ -879,8 +869,7 @@ def test_variables(assert_exp_eq): db.table WHERE col = @VAR('test_var') - """ - ) + """) audit = load_audit( expressions, @@ -888,7 +877,9 @@ def test_variables(assert_exp_eq): dialect="bigquery", variables={"test_var": "test_val", "test_var_unused": "unused_val"}, ) - assert audit.python_env[c.SQLMESH_VARS] == Executable.value({"test_var": "test_val"}) + assert audit.python_env[c.SQLMESH_VARS] == Executable.value( + {"test_var": "test_val"} + ) assert ( audit.render_audit_query().sql(dialect="bigquery") == "SELECT * FROM `db`.`table` AS `table` WHERE `col` = 'test_val'" @@ -896,8 +887,7 @@ def test_variables(assert_exp_eq): def test_load_inline_audits(assert_exp_eq): - expressions = parse( - """ + expressions = parse(""" MODEL ( name db.table, dialect spark, @@ -919,8 +909,7 @@ def test_load_inline_audits(assert_exp_eq): FROM @this_model WHERE id < 0; - """ - ) + """) model = load_sql_based_model(expressions) assert len(model.audits) == 2 @@ -937,7 +926,9 @@ def test_model_inline_audits(sushi_context: Context): assert isinstance(model, SeedModel) assert len(model.audit_definitions) == 3 assert isinstance(model.audit_definitions["assert_valid_name"], ModelAudit) - model.render_audit_query(model.audit_definitions["assert_positive_id"]).sql() == expected_query + model.render_audit_query( + model.audit_definitions["assert_positive_id"] + ).sql() == expected_query def test_audit_query_normalization(): @@ -959,11 +950,13 @@ def test_audit_query_normalization(): def test_rendered_diff(): audit1 = StandaloneAudit( - name="test_audit", query=parse_one("SELECT * FROM 'test' WHERE @AND(TRUE, NULL) > 2") + name="test_audit", + query=parse_one("SELECT * FROM 'test' WHERE @AND(TRUE, NULL) > 2"), ) audit2 = StandaloneAudit( - name="test_audit", query=parse_one("SELECT * FROM 'test' WHERE @OR(FALSE, NULL) > 2") + name="test_audit", + query=parse_one("SELECT * FROM 'test' WHERE @OR(FALSE, NULL) > 2"), ) assert """@@ -6,4 +6,4 @@ @@ -976,8 +969,7 @@ def test_rendered_diff(): def test_multiple_audits_with_same_name(): - expressions = parse( - """ + expressions = parse(""" MODEL ( name db.table, dialect spark, @@ -995,8 +987,7 @@ def test_multiple_audits_with_same_name(): ); SELECT * FROM @this_model WHERE @column >= @threshold; - """ - ) + """) model = load_sql_based_model(expressions) assert len(model.audits) == 3 assert len(model.audits_with_args) == 3 @@ -1022,7 +1013,8 @@ def test_default_audits_included_when_no_model_audits(): """) model_defaults = ModelDefaultsConfig( - dialect="duckdb", audits=["not_null(columns := ['id'])", "unique_values(columns := ['id'])"] + dialect="duckdb", + audits=["not_null(columns := ['id'])", "unique_values(columns := ['id'])"], ) model = load_sql_based_model(expressions, defaults=model_defaults.dict()) @@ -1050,8 +1042,7 @@ def test_default_audits_included_when_no_model_audits(): def test_model_defaults_audits_with_same_name(): - expressions = parse( - """ + expressions = parse(""" MODEL ( name db.table, dialect spark, @@ -1069,8 +1060,7 @@ def test_model_defaults_audits_with_same_name(): ); SELECT * FROM @this_model WHERE @column >= @threshold; - """ - ) + """) model_defaults = ModelDefaultsConfig( dialect="duckdb", @@ -1125,8 +1115,7 @@ def test_model_defaults_audits_with_same_name(): def test_audit_formatting_flag_serde(): - expressions = parse( - """ + expressions = parse(""" AUDIT ( name my_audit, dialect bigquery, @@ -1134,8 +1123,7 @@ def test_audit_formatting_flag_serde(): ); SELECT * FROM db.table WHERE col = @VAR('test_var') - """ - ) + """) audit = load_audit( expressions, diff --git a/tests/core/test_config.py b/tests/core/test_config.py index 0da5b6e22f..442818b504 100644 --- a/tests/core/test_config.py +++ b/tests/core/test_config.py @@ -1,41 +1,36 @@ import os import pathlib import re +import typing as t from pathlib import Path from unittest import mock -import typing as t import pytest from pytest_mock import MockerFixture from sqlglot import exp -from sqlmesh.core.config import ( - Config, - DuckDBConnectionConfig, - GatewayConfig, - ModelDefaultsConfig, - BigQueryConnectionConfig, - MotherDuckConnectionConfig, - BuiltInSchedulerConfig, - EnvironmentSuffixTarget, - TableNamingConvention, -) -from sqlmesh.core.config.connection import DuckDBAttachOptions, RedshiftConnectionConfig -from sqlmesh.core.config.loader import ( - load_config_from_env, - load_config_from_paths, - load_config_from_python_module, - load_configs, -) +from sqlmesh.core.config import (BigQueryConnectionConfig, + BuiltInSchedulerConfig, Config, + DuckDBConnectionConfig, + EnvironmentSuffixTarget, GatewayConfig, + ModelDefaultsConfig, + MotherDuckConnectionConfig, + TableNamingConvention) +from sqlmesh.core.config.connection import (DuckDBAttachOptions, + RedshiftConnectionConfig) +from sqlmesh.core.config.loader import (load_config_from_env, + load_config_from_paths, + load_config_from_python_module, + load_configs) from sqlmesh.core.context import Context from sqlmesh.core.engine_adapter.athena import AthenaEngineAdapter from sqlmesh.core.engine_adapter.duckdb import DuckDBEngineAdapter from sqlmesh.core.engine_adapter.redshift import RedshiftEngineAdapter from sqlmesh.core.notification_target import ConsoleNotificationTarget from sqlmesh.core.user import User -from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils import yaml from sqlmesh.dbt.loader import DbtLoader +from sqlmesh.utils import yaml +from sqlmesh.utils.errors import ConfigError from tests.utils.test_filesystem import create_temp_file @@ -43,8 +38,7 @@ def yaml_config_path(tmp_path_factory) -> Path: config_path = tmp_path_factory.mktemp("yaml_config") / "config.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: another_gateway: connection: @@ -53,8 +47,7 @@ def yaml_config_path(tmp_path_factory) -> Path: model_defaults: dialect: '' - """ - ) + """) return config_path @@ -82,9 +75,9 @@ def test_update_with_gateways(): Config(gateways=gateway0_config) ) == Config(gateways={"gateway1": gateway1_config, "": gateway0_config}) - assert Config(gateways=gateway0_config).update_with(Config(gateways=gateway1_config)) == Config( - gateways=gateway1_config - ) + assert Config(gateways=gateway0_config).update_with( + Config(gateways=gateway1_config) + ) == Config(gateways=gateway1_config) assert Config(gateways={"gateway0": gateway0_config}).update_with( Config(gateways={"gateway1": gateway1_config}) @@ -118,7 +111,9 @@ def test_update_with_notification_targets(): def test_update_with_model_defaults(): - config_a = Config(model_defaults=ModelDefaultsConfig(start="2022-01-01", dialect="duckdb")) + config_a = Config( + model_defaults=ModelDefaultsConfig(start="2022-01-01", dialect="duckdb") + ) config_b = Config(model_defaults=ModelDefaultsConfig(dialect="spark")) assert config_a.update_with(config_b) == Config( @@ -141,7 +136,9 @@ def test_default_gateway(): assert config.get_gateway() == gateway_a - assert config.copy(update={"default_gateway": "gateway2"}).get_gateway() == gateway_c + assert ( + config.copy(update={"default_gateway": "gateway2"}).get_gateway() == gateway_c + ) assert ( Config( @@ -168,7 +165,9 @@ def test_load_config_from_paths(yaml_config_path: Path, python_config_path: Path assert config == Config( gateways={ # type: ignore - "another_gateway": GatewayConfig(connection=DuckDBConnectionConfig(database="test_db")), + "another_gateway": GatewayConfig( + connection=DuckDBConnectionConfig(database="test_db") + ), "": GatewayConfig(connection=DuckDBConnectionConfig()), }, model_defaults=ModelDefaultsConfig(dialect=""), @@ -184,12 +183,16 @@ def test_load_config_multiple_config_files_in_folder(tmp_path): with open(config_b_path, "w", encoding="utf-8") as fd: fd.write("project: project_b") - with pytest.raises(ConfigError, match=r"^Multiple configuration files found in folder.*"): + with pytest.raises( + ConfigError, match=r"^Multiple configuration files found in folder.*" + ): load_config_from_paths(Config, project_paths=[config_a_path, config_b_path]) def test_load_config_no_config(): - with pytest.raises(ConfigError, match=r"^SQLMesh project config could not be found.*"): + with pytest.raises( + ConfigError, match=r"^SQLMesh project config could not be found.*" + ): load_config_from_paths(Config, load_from_env=False) @@ -217,12 +220,14 @@ def test_load_config_no_dialect(tmp_path): ) with pytest.raises( - ConfigError, match=r"^Default model SQL dialect is a required configuration parameter.*" + ConfigError, + match=r"^Default model SQL dialect is a required configuration parameter.*", ): load_config_from_paths(Config, project_paths=[tmp_path / "config.yaml"]) with pytest.raises( - ConfigError, match=r"^Default model SQL dialect is a required configuration parameter.*" + ConfigError, + match=r"^Default model SQL dialect is a required configuration parameter.*", ): load_config_from_paths(Config, project_paths=[tmp_path / "config.py"]) @@ -254,7 +259,9 @@ def test_load_config_unsupported_extension(tmp_path): config_path = tmp_path / "config.txt" config_path.touch() - with pytest.raises(ConfigError, match=r"^Unsupported config file extension 'txt'.*"): + with pytest.raises( + ConfigError, match=r"^Unsupported config file extension 'txt'.*" + ): load_config_from_paths(Config, project_paths=[config_path]) @@ -300,12 +307,16 @@ def test_load_config_from_env(): }, ): assert Config.parse_obj(load_config_from_env()) == Config( - gateways=GatewayConfig(connection=DuckDBConnectionConfig(database="test_db")), + gateways=GatewayConfig( + connection=DuckDBConnectionConfig(database="test_db") + ), ) def test_load_config_from_env_fails(): - with mock.patch.dict(os.environ, {"SQLMESH__GATEWAYS__ABCDEF__CONNECTION__PASSWORD": "..."}): + with mock.patch.dict( + os.environ, {"SQLMESH__GATEWAYS__ABCDEF__CONNECTION__PASSWORD": "..."} + ): with pytest.raises( ConfigError, match="Missing connection type.\n\nVerify your config.yaml and environment variables.", @@ -340,8 +351,7 @@ def test_load_config_from_env_invalid_variable_name(): def test_load_yaml_config_env_var_gateway_override(tmp_path_factory): config_path = tmp_path_factory.mktemp("yaml_config") / "config.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: testing: connection: @@ -349,8 +359,7 @@ def test_load_yaml_config_env_var_gateway_override(tmp_path_factory): database: blah model_defaults: dialect: bigquery - """ - ) + """) with mock.patch.dict( os.environ, { @@ -437,7 +446,10 @@ def test_load_config_from_python_module_invalid_config_object(tmp_path): [ ( "'^dev$': dev_catalog\n '^other$': other_catalog", - {re.compile("^dev$"): "dev_catalog", re.compile("^other$"): "other_catalog"}, + { + re.compile("^dev$"): "dev_catalog", + re.compile("^other$"): "other_catalog", + }, "duckdb", "", ), @@ -461,11 +473,12 @@ def test_load_config_from_python_module_invalid_config_object(tmp_path): ), ], ) -def test_environment_catalog_mapping(tmp_path_factory, mapping, expected, dialect, raise_error): +def test_environment_catalog_mapping( + tmp_path_factory, mapping, expected, dialect, raise_error +): config_path = tmp_path_factory.mktemp("yaml_config") / "config.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - f""" + fd.write(f""" gateways: local: connection: @@ -476,8 +489,7 @@ def test_environment_catalog_mapping(tmp_path_factory, mapping, expected, dialec environment_catalog_mapping: {mapping} - """ - ) + """) if raise_error: with pytest.raises(ConfigError, match=raise_error): load_config_from_paths( @@ -486,17 +498,22 @@ def test_environment_catalog_mapping(tmp_path_factory, mapping, expected, dialec ) else: assert ( - load_config_from_paths(Config, project_paths=[config_path]).environment_catalog_mapping + load_config_from_paths( + Config, project_paths=[config_path] + ).environment_catalog_mapping == expected ) -def test_physical_schema_mapping_mutually_exclusive_with_physical_schema_override() -> None: +def test_physical_schema_mapping_mutually_exclusive_with_physical_schema_override() -> ( + None +): Config(physical_schema_override={"foo": "bar"}) # type: ignore Config(physical_schema_mapping={"^foo$": "bar"}) with pytest.raises( - ConfigError, match=r"Only one.*physical_schema_override.*physical_schema_mapping" + ConfigError, + match=r"Only one.*physical_schema_override.*physical_schema_mapping", ): Config(physical_schema_override={"foo": "bar"}, physical_schema_mapping={"^foo$": "bar"}) # type: ignore @@ -512,7 +529,9 @@ class DerivedConfig(Config): assert config == DerivedConfig( gateways={ # type: ignore - "another_gateway": GatewayConfig(connection=DuckDBConnectionConfig(database="test_db")), + "another_gateway": GatewayConfig( + connection=DuckDBConnectionConfig(database="test_db") + ), "": GatewayConfig(connection=DuckDBConnectionConfig()), }, model_defaults=ModelDefaultsConfig(dialect=""), @@ -566,7 +585,8 @@ def test_variables(): "UPPERCASE_VAR": 2, } config = Config( - variables=variables, gateways={"local": GatewayConfig(variables=gateway_variables)} + variables=variables, + gateways={"local": GatewayConfig(variables=gateway_variables)}, ) assert config.variables == variables assert config.get_gateway("local").variables == {"uppercase_var": 2} @@ -581,8 +601,7 @@ def test_variables(): def test_load_duckdb_attach_config(tmp_path): config_path = tmp_path / "config_duckdb_attach.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: another_gateway: connection: @@ -598,24 +617,30 @@ def test_load_duckdb_attach_config(tmp_path): read_only: true model_defaults: dialect: '' - """ - ) + """) config = load_config_from_paths( Config, project_paths=[config_path], ) - assert config.gateways["another_gateway"].connection.catalogs.get("memory") == ":memory:" + assert ( + config.gateways["another_gateway"].connection.catalogs.get("memory") + == ":memory:" + ) - attach_config_1 = config.gateways["another_gateway"].connection.catalogs.get("sqlite") + attach_config_1 = config.gateways["another_gateway"].connection.catalogs.get( + "sqlite" + ) assert isinstance(attach_config_1, DuckDBAttachOptions) assert attach_config_1.type == "sqlite" assert attach_config_1.path == "test.db" assert attach_config_1.read_only is False - attach_config_2 = config.gateways["another_gateway"].connection.catalogs.get("postgres") + attach_config_2 = config.gateways["another_gateway"].connection.catalogs.get( + "postgres" + ) assert isinstance(attach_config_2, DuckDBAttachOptions) assert attach_config_2.type == "postgres" @@ -626,15 +651,13 @@ def test_load_duckdb_attach_config(tmp_path): def test_load_model_defaults_audits(tmp_path): config_path = tmp_path / "config_model_defaults_audits.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" model_defaults: dialect: '' audits: - assert_positive_order_ids - does_not_exceed_threshold(column := id, threshold := 1000) - """ - ) + """) config = load_config_from_paths( Config, @@ -653,8 +676,7 @@ def test_load_model_defaults_audits(tmp_path): def test_load_model_defaults_statements(tmp_path): config_path = tmp_path / "config_model_defaults_statements.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" model_defaults: dialect: duckdb pre_statements: @@ -666,8 +688,7 @@ def test_load_model_defaults_statements(tmp_path): - SET memory_limit = '5GB' on_virtual_update: - UPDATE stats_table SET last_update = CURRENT_TIMESTAMP - """ - ) + """) config = load_config_from_paths( Config, @@ -677,32 +698,42 @@ def test_load_model_defaults_statements(tmp_path): assert config.model_defaults.pre_statements is not None assert len(config.model_defaults.pre_statements) == 2 assert isinstance(exp.maybe_parse(config.model_defaults.pre_statements[0]), exp.Set) - assert isinstance(exp.maybe_parse(config.model_defaults.pre_statements[1]), exp.Create) + assert isinstance( + exp.maybe_parse(config.model_defaults.pre_statements[1]), exp.Create + ) assert config.model_defaults.post_statements is not None assert len(config.model_defaults.post_statements) == 3 - assert isinstance(exp.maybe_parse(config.model_defaults.post_statements[0]), exp.Drop) - assert isinstance(exp.maybe_parse(config.model_defaults.post_statements[1]), exp.Analyze) - assert isinstance(exp.maybe_parse(config.model_defaults.post_statements[2]), exp.Set) + assert isinstance( + exp.maybe_parse(config.model_defaults.post_statements[0]), exp.Drop + ) + assert isinstance( + exp.maybe_parse(config.model_defaults.post_statements[1]), exp.Analyze + ) + assert isinstance( + exp.maybe_parse(config.model_defaults.post_statements[2]), exp.Set + ) assert config.model_defaults.on_virtual_update is not None assert len(config.model_defaults.on_virtual_update) == 1 - assert isinstance(exp.maybe_parse(config.model_defaults.on_virtual_update[0]), exp.Update) + assert isinstance( + exp.maybe_parse(config.model_defaults.on_virtual_update[0]), exp.Update + ) def test_load_model_defaults_validation_statements(tmp_path): config_path = tmp_path / "config_model_defaults_statements_wrong.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" model_defaults: dialect: duckdb pre_statements: - 313 - """ - ) + """) - with pytest.raises(ConfigError, match=r"Invalid field 'model_defaults\.pre_statements\.0"): + with pytest.raises( + ConfigError, match=r"Invalid field 'model_defaults\.pre_statements\.0" + ): config = load_config_from_paths( Config, project_paths=[config_path], @@ -712,8 +743,7 @@ def test_load_model_defaults_validation_statements(tmp_path): def test_scheduler_config(tmp_path_factory): config_path = tmp_path_factory.mktemp("yaml_config") / "config.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: builtin_gateway: scheduler: @@ -724,22 +754,22 @@ def test_scheduler_config(tmp_path_factory): model_defaults: dialect: bigquery - """ - ) + """) config = load_config_from_paths( Config, project_paths=[config_path], ) assert isinstance(config.default_scheduler, BuiltInSchedulerConfig) - assert isinstance(config.get_gateway("builtin_gateway").scheduler, BuiltInSchedulerConfig) + assert isinstance( + config.get_gateway("builtin_gateway").scheduler, BuiltInSchedulerConfig + ) def test_multi_gateway_config(tmp_path, mocker: MockerFixture): config_path = tmp_path / "config_athena_redshift.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: redshift: connection: @@ -770,8 +800,7 @@ def test_multi_gateway_config(tmp_path, mocker: MockerFixture): model_defaults: dialect: redshift - """ - ) + """) config = load_config_from_paths( Config, @@ -794,8 +823,7 @@ def test_multi_gateway_config(tmp_path, mocker: MockerFixture): def test_multi_gateway_single_threaded_config(tmp_path): config_path = tmp_path / "config_duck_athena.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: duckdb: connection: @@ -811,8 +839,7 @@ def test_multi_gateway_single_threaded_config(tmp_path): default_gateway: duckdb model_defaults: dialect: duckdb - """ - ) + """) config = load_config_from_paths( Config, @@ -832,8 +859,7 @@ def test_multi_gateway_single_threaded_config(tmp_path): def test_trino_schema_location_mapping_syntax(tmp_path): config_path = tmp_path / "config_trino.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: trino: connection: @@ -849,8 +875,7 @@ def test_trino_schema_location_mapping_syntax(tmp_path): model_defaults: dialect: trino - """ - ) + """) config = load_config_from_paths( Config, @@ -868,8 +893,7 @@ def test_trino_schema_location_mapping_syntax(tmp_path): def test_trino_source_option(tmp_path): config_path = tmp_path / "config_trino_source.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: trino: connection: @@ -883,8 +907,7 @@ def test_trino_source_option(tmp_path): model_defaults: dialect: trino - """ - ) + """) config = load_config_from_paths( Config, @@ -901,8 +924,7 @@ def test_trino_source_option(tmp_path): def test_gcp_postgres_ip_and_scopes(tmp_path): config_path = tmp_path / "config_gcp_postgres.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: gcp_postgres: connection: @@ -921,8 +943,7 @@ def test_gcp_postgres_ip_and_scopes(tmp_path): model_defaults: dialect: postgres - """ - ) + """) config = load_config_from_paths( Config, @@ -942,9 +963,15 @@ def test_gcp_postgres_ip_and_scopes(tmp_path): def test_gateway_model_defaults(tmp_path): global_defaults = ModelDefaultsConfig( - dialect="snowflake", owner="foo", optimize_query=True, enabled=True, cron="@daily" + dialect="snowflake", + owner="foo", + optimize_query=True, + enabled=True, + cron="@daily", + ) + gateway_defaults = ModelDefaultsConfig( + dialect="duckdb", owner="baz", optimize_query=False ) - gateway_defaults = ModelDefaultsConfig(dialect="duckdb", owner="baz", optimize_query=False) config = Config( gateways={ @@ -993,14 +1020,12 @@ def test_model_defaults_cron_tz(tmp_path): config_path = tmp_path / "config_model_defaults_cron_tz.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" model_defaults: dialect: duckdb cron: '@daily' cron_tz: 'America/Los_Angeles' - """ - ) + """) config = load_config_from_paths( Config, @@ -1047,8 +1072,7 @@ def test_gateway_model_defaults_cron_tz(tmp_path): def test_redshift_merge_flag(tmp_path, mocker: MockerFixture): config_path = tmp_path / "config_redshift_merge.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: redshift: connection: @@ -1070,8 +1094,7 @@ def test_redshift_merge_flag(tmp_path, mocker: MockerFixture): model_defaults: dialect: redshift - """ - ) + """) config = load_config_from_paths( Config, @@ -1092,8 +1115,7 @@ def test_redshift_merge_flag(tmp_path, mocker: MockerFixture): def test_environment_statements_config(tmp_path): config_path = tmp_path / "config_before_after_all.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" gateways: postgres: connection: @@ -1114,8 +1136,7 @@ def test_environment_statements_config(tmp_path): model_defaults: dialect: postgres - """ - ) + """) config = load_config_from_paths( Config, @@ -1267,12 +1288,10 @@ def test_load_python_config_dot_env_vars(tmp_path_factory): # SQLMESH__ variables override config fields directly if they follow the naming structure dot_path = main_dir / ".env" with open(dot_path, "w", encoding="utf-8") as fd: - fd.write( - """SQLMESH__GATEWAYS__DUCKDB_GATEWAY__STATE_CONNECTION__TYPE="bigquery" + fd.write("""SQLMESH__GATEWAYS__DUCKDB_GATEWAY__STATE_CONNECTION__TYPE="bigquery" SQLMESH__GATEWAYS__DUCKDB_GATEWAY__STATE_CONNECTION__CHECK_IMPORT="false" SQLMESH__DEFAULT_GATEWAY="duckdb_gateway" - """ - ) + """) # Use mock.patch.dict to isolate environment variables between the tests with mock.patch.dict(os.environ, {}, clear=True): @@ -1298,8 +1317,7 @@ def test_load_yaml_config_dot_env_vars(tmp_path_factory): main_dir = tmp_path_factory.mktemp("yaml_config") config_path = main_dir / "config.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """gateways: + fd.write("""gateways: duckdb_gateway: connection: type: duckdb @@ -1314,21 +1332,18 @@ def test_load_yaml_config_dot_env_vars(tmp_path_factory): secret: {{ env_var('S3_SECRET') }} model_defaults: dialect: "" -""" - ) +""") # This test checks both using SQLMESH__ prefixed environment variables with underscores # and setting a regular environment variable for use with env_var(). dot_path = main_dir / ".env" with open(dot_path, "w", encoding="utf-8") as fd: - fd.write( - """S3_BUCKET="s3://metrics_bucket/sales.db" + fd.write("""S3_BUCKET="s3://metrics_bucket/sales.db" S3_KEY="S3_KEY_ID" S3_SECRET="XXX_S3_SECRET_XXX" SQLMESH__DEFAULT_GATEWAY="duckdb_gateway" SQLMESH__MODEL_DEFAULTS__DIALECT="athena" -""" - ) +""") # Use mock.patch.dict to isolate environment variables between the tests with mock.patch.dict(os.environ, {}, clear=True): @@ -1347,7 +1362,13 @@ def test_load_yaml_config_dot_env_vars(tmp_path_factory): "cloud_sales": "s3://metrics_bucket/sales.db", }, extensions=[{"name": "httpfs"}], - secrets=[{"type": "s3", "key_id": "S3_KEY_ID", "secret": "XXX_S3_SECRET_XXX"}], + secrets=[ + { + "type": "s3", + "key_id": "S3_KEY_ID", + "secret": "XXX_S3_SECRET_XXX", + } + ], ), ), }, @@ -1360,16 +1381,14 @@ def test_load_config_dotenv_directory_not_loaded(tmp_path_factory): main_dir = tmp_path_factory.mktemp("config_with_env_dir") config_path = main_dir / "config.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """gateways: + fd.write("""gateways: test_gateway: connection: type: duckdb database: test.db model_defaults: dialect: duckdb -""" - ) +""") # Create a .env directory instead of a file to simulate a Python virtual environment env_dir = main_dir / ".env" @@ -1380,16 +1399,14 @@ def test_load_config_dotenv_directory_not_loaded(tmp_path_factory): other_dir = tmp_path_factory.mktemp("config_with_env_file") other_config_path = other_dir / "config.yaml" with open(other_config_path, "w", encoding="utf-8") as fd: - fd.write( - """gateways: + fd.write("""gateways: test_gateway: connection: type: duckdb database: test.db model_defaults: dialect: duckdb -""" - ) +""") env_file = other_dir / ".env" with open(env_file, "w", encoding="utf-8") as fd: @@ -1420,25 +1437,21 @@ def test_load_yaml_config_custom_dotenv_path(tmp_path_factory): main_dir = tmp_path_factory.mktemp("yaml_config_2") config_path = main_dir / "config.yaml" with open(config_path, "w", encoding="utf-8") as fd: - fd.write( - """gateways: + fd.write("""gateways: test_gateway: connection: type: duckdb database: {{ env_var('DB_NAME') }} -""" - ) +""") # Create a custom dot env file in a different location custom_env_dir = tmp_path_factory.mktemp("custom_env") custom_env_path = custom_env_dir / ".my_env" with open(custom_env_path, "w", encoding="utf-8") as fd: - fd.write( - """DB_NAME="custom_database.db" + fd.write("""DB_NAME="custom_database.db" SQLMESH__DEFAULT_GATEWAY="test_gateway" SQLMESH__MODEL_DEFAULTS__DIALECT="postgres" -""" - ) +""") # Test that without custom dotenv path, env vars are not loaded with mock.patch.dict(os.environ, {}, clear=True): @@ -1483,9 +1496,13 @@ def test_load_yaml_config_custom_dotenv_path(tmp_path_factory): ], ) def test_physical_table_naming_convention( - convention_str: t.Optional[str], expected: t.Optional[TableNamingConvention], tmp_path: Path + convention_str: t.Optional[str], + expected: t.Optional[TableNamingConvention], + tmp_path: Path, ): - config_part = f"physical_table_naming_convention: {convention_str}" if convention_str else "" + config_part = ( + f"physical_table_naming_convention: {convention_str}" if convention_str else "" + ) (tmp_path / "config.yaml").write_text(f""" gateways: test_gateway: @@ -1527,13 +1544,17 @@ def test_load_configs_without_main_connection(tmp_path: Path): with config_file.open("w") as f: yaml.dump( { - "gateways": {"": {"state_connection": {"type": "duckdb", "database": "state.db"}}}, + "gateways": { + "": {"state_connection": {"type": "duckdb", "database": "state.db"}} + }, "model_defaults": {"dialect": "duckdb", "start": "2020-01-01"}, }, f, ) - configs = list(load_configs(config=None, config_type=Config, paths=[tmp_path]).values()) + configs = list( + load_configs(config=None, config_type=Config, paths=[tmp_path]).values() + ) assert len(configs) == 1 config = configs[0] @@ -1572,7 +1593,9 @@ def test_load_configs_in_dbt_project_without_config_py(tmp_path: Path): start: '2020-01-01' """) - configs = list(load_configs(config=None, config_type=Config, paths=[tmp_path]).values()) + configs = list( + load_configs(config=None, config_type=Config, paths=[tmp_path]).values() + ) assert len(configs) == 1 config = configs[0] @@ -1607,7 +1630,9 @@ def test_canonicalize_sorts_sets() -> None: def test_canonicalize_recurses_into_containers() -> None: from sqlmesh.core.config.root import _canonicalize - assert _canonicalize({"rules": {"z", "a"}, "nested": [{3, 1}, ("x", {"q", "b"})]}) == { + assert _canonicalize( + {"rules": {"z", "a"}, "nested": [{3, 1}, ("x", {"q", "b"})]} + ) == { "rules": ["a", "z"], "nested": [[1, 3], ("x", ["b", "q"])], } diff --git a/tests/core/test_connection_config.py b/tests/core/test_connection_config.py index c506d401d9..213d6829a0 100644 --- a/tests/core/test_connection_config.py +++ b/tests/core/test_connection_config.py @@ -7,28 +7,26 @@ from _pytest.fixtures import FixtureRequest from sqlglot import exp -from sqlmesh.core.config.connection import ( - INIT_DISPLAY_INFO_TO_TYPE, - SUPPORTS_MSSQL_PYTHON_DRIVER, - AthenaConnectionConfig, - BigQueryConnectionConfig, - ClickhouseConnectionConfig, - ConnectionConfig, - DatabricksConnectionConfig, - DuckDBAttachOptions, - DuckDBConnectionConfig, - FabricConnectionConfig, - GCPPostgresConnectionConfig, - MotherDuckConnectionConfig, - MSSQLConnectionConfig, - MySQLConnectionConfig, - PostgresConnectionConfig, - SnowflakeConnectionConfig, - StarRocksConnectionConfig, - TrinoAuthenticationMethod, - _connection_config_validator, - _get_engine_import_validator, -) +from sqlmesh.core.config.connection import (INIT_DISPLAY_INFO_TO_TYPE, + SUPPORTS_MSSQL_PYTHON_DRIVER, + AthenaConnectionConfig, + BigQueryConnectionConfig, + ClickhouseConnectionConfig, + ConnectionConfig, + DatabricksConnectionConfig, + DuckDBAttachOptions, + DuckDBConnectionConfig, + FabricConnectionConfig, + GCPPostgresConnectionConfig, + MotherDuckConnectionConfig, + MSSQLConnectionConfig, + MySQLConnectionConfig, + PostgresConnectionConfig, + SnowflakeConnectionConfig, + StarRocksConnectionConfig, + TrinoAuthenticationMethod, + _connection_config_validator, + _get_engine_import_validator) from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.pydantic import PydanticModel @@ -92,7 +90,9 @@ def snowflake_oauth_access_token() -> str: return "eyJhbGciOiJSUzI1NiIsImtpZCI6ImFmZmM2MjkwN2E0NDYxODJhZGMxZmE0ZTgxZmRiYTYzMTBkY2U2M2YifQ.eyJhenAiOiIyNzIxOTYwNjkxNzMtZm81ZWI0MXQzbmR1cTZ1ZXRkc2pkdWdzZXV0ZnBtc3QuYXBwcy5nb29nbGV1c2VyY29udGVudC5jb20iLCJhdWQiOiIyNzIxOTYwNjkxNzMtZm81ZWI0MXQzbmR1cTZ1ZXRkc2pkdWdzZXV0ZnBtc3QuYXBwcy5nb29nbGV1c2VyY29udGVudC5jb20iLCJzdWIiOiIxMTc4NDc5MTI4NzU5MTM5MDU0OTMiLCJlbWFpbCI6ImFhcm9uLnBhcmVja2lAZ21haWwuY29tIiwiZW1haWxfdmVyaWZpZWQiOnRydWUsImF0X2hhc2giOiJpRVljNDBUR0luUkhoVEJidWRncEpRIiwiZXhwIjoxNTI0NTk5MDU2LCJpc3MiOiJodHRwczovL2FjY291bnRzLmdvb2dsZS5jb20iLCJpYXQiOjE1MjQ1OTU0NTZ9.ho2czp_1JWsglJ9jN8gCgWfxDi2gY4X5-QcT56RUGkgh5BJaaWdlrRhhN_eNuJyN3HRPhvVA_KJVy1tMltTVd2OQ6VkxgBNfBsThG_zLPZriw7a1lANblarwxLZID4fXDYG-O8U-gw4xb-NIsOzx6xsxRBdfKKniavuEg56Sd3eKYyqrMA0DWnIagqLiKE6kpZkaGImIpLcIxJPF0-yeJTMt_p1NoJF7uguHHLYr6752hqppnBpMjFL2YMDVeg3jl1y5DeSKNPh6cZ8H2p4Xb2UIrJguGbQHVIJvtm_AspRjrmaTUQKrzXDRCfDROSUU-h7XKIWRrEd2-W9UkV5oCg" -def test_snowflake(make_config, snowflake_key_passphrase_bytes, snowflake_oauth_access_token): +def test_snowflake( + make_config, snowflake_key_passphrase_bytes, snowflake_oauth_access_token +): # Authenticator and user/password is fine config = make_config( type="snowflake", @@ -103,26 +103,34 @@ def test_snowflake(make_config, snowflake_key_passphrase_bytes, snowflake_oauth_ ) assert isinstance(config, SnowflakeConnectionConfig) # Auth with no user/password is fine - config = make_config(type="snowflake", account="test", authenticator="externalbrowser") + config = make_config( + type="snowflake", account="test", authenticator="externalbrowser" + ) assert isinstance(config, SnowflakeConnectionConfig) # No auth and no user raises with pytest.raises( - ConfigError, match="User and password must be provided if using default authentication" + ConfigError, + match="User and password must be provided if using default authentication", ): make_config(type="snowflake", account="test", password="test") # No auth and no password raises with pytest.raises( - ConfigError, match="User and password must be provided if using default authentication" + ConfigError, + match="User and password must be provided if using default authentication", ): make_config(type="snowflake", account="test", user="test") # No auth and no user/password raises with pytest.raises( - ConfigError, match="User and password must be provided if using default authentication" + ConfigError, + match="User and password must be provided if using default authentication", ): make_config(type="snowflake", account="test") # Private key and username with no authenticator is fine config = make_config( - type="snowflake", account="test", private_key=snowflake_key_passphrase_bytes, user="test" + type="snowflake", + account="test", + private_key=snowflake_key_passphrase_bytes, + user="test", ) assert isinstance(config, SnowflakeConnectionConfig) # Private key with jwt auth is fine @@ -136,7 +144,8 @@ def test_snowflake(make_config, snowflake_key_passphrase_bytes, snowflake_oauth_ assert isinstance(config, SnowflakeConnectionConfig) # Private key without username raises with pytest.raises( - ConfigError, match=r"User must be provided when using SNOWFLAKE_JWT authentication" + ConfigError, + match=r"User must be provided when using SNOWFLAKE_JWT authentication", ): make_config( type="snowflake", @@ -146,7 +155,8 @@ def test_snowflake(make_config, snowflake_key_passphrase_bytes, snowflake_oauth_ ) # Private key with password raises with pytest.raises( - ConfigError, match=r"Password cannot be provided when using SNOWFLAKE_JWT authentication" + ConfigError, + match=r"Password cannot be provided when using SNOWFLAKE_JWT authentication", ): make_config( type="snowflake", @@ -189,7 +199,9 @@ def test_snowflake(make_config, snowflake_key_passphrase_bytes, snowflake_oauth_ ) assert isinstance(config, SnowflakeConnectionConfig) assert config.get_catalog() == "test_catalog" - with pytest.raises(ConfigError, match=r"Token must be provided if using oauth authentication"): + with pytest.raises( + ConfigError, match=r"Token must be provided if using oauth authentication" + ): make_config( type="snowflake", account="test", @@ -279,18 +291,23 @@ def test_snowflake_private_key_pass( def test_validator(): assert _connection_config_validator(None, None) is None - snowflake_config = SnowflakeConnectionConfig(account="test", authenticator="externalbrowser") + snowflake_config = SnowflakeConnectionConfig( + account="test", authenticator="externalbrowser" + ) assert _connection_config_validator(None, snowflake_config) == snowflake_config assert ( _connection_config_validator( - None, dict(type="snowflake", account="test", authenticator="externalbrowser") + None, + dict(type="snowflake", account="test", authenticator="externalbrowser"), ) == snowflake_config ) with pytest.raises(ConfigError, match="Missing connection type."): - _connection_config_validator(None, dict(account="test", authenticator="externalbrowser")) + _connection_config_validator( + None, dict(account="test", authenticator="externalbrowser") + ) with pytest.raises(ConfigError, match="Unknown connection type 'invalid'."): _connection_config_validator( @@ -325,7 +342,8 @@ def test_trino(make_config): assert config.catalog == "catalog" with pytest.raises( - ConfigError, match="Username and Password must be provided if using basic authentication" + ConfigError, + match="Username and Password must be provided if using basic authentication", ): make_config(method="basic", **required_kwargs) @@ -334,7 +352,8 @@ def test_trino(make_config): assert config.method == TrinoAuthenticationMethod.LDAP with pytest.raises( - ConfigError, match="Username and Password must be provided if using ldap authentication" + ConfigError, + match="Username and Password must be provided if using ldap authentication", ): make_config(method="ldap", **required_kwargs) @@ -389,7 +408,9 @@ def test_trino(make_config): assert config.port == 443 # Validate http is only for basic and no auth - config = make_config(method="basic", password="password", http_scheme="http", **required_kwargs) + config = make_config( + method="basic", password="password", http_scheme="http", **required_kwargs + ) assert config.method == TrinoAuthenticationMethod.BASIC assert config.http_scheme == "http" @@ -469,9 +490,9 @@ def test_trino_timestamp_mapping(make_config): ) assert config.timestamp_mapping is not None - assert config.timestamp_mapping[exp.DataType.build("TIMESTAMP")] == exp.DataType.build( - "TIMESTAMP(6)" - ) + assert config.timestamp_mapping[ + exp.DataType.build("TIMESTAMP") + ] == exp.DataType.build("TIMESTAMP(6)") # Test with invalid source type with pytest.raises(ConfigError) as exc_info: @@ -847,14 +868,18 @@ def test_duckdb_attach_ducklake_catalog(make_config): ) ducklake_catalog_with_prefix = config_with_prefix.catalogs.get("ducklake") generated_sql_with_prefix = ducklake_catalog_with_prefix.to_sql("ducklake") - assert "ATTACH IF NOT EXISTS 'ducklake:catalog.ducklake'" in generated_sql_with_prefix + assert ( + "ATTACH IF NOT EXISTS 'ducklake:catalog.ducklake'" in generated_sql_with_prefix + ) # Ensure we don't have double prefixes assert "'ducklake:catalog.ducklake" in generated_sql_with_prefix def test_duckdb_attach_options(): options = DuckDBAttachOptions( - type="postgres", path="dbname=postgres user=postgres host=127.0.0.1", read_only=True + type="postgres", + path="dbname=postgres user=postgres host=127.0.0.1", + read_only=True, ) assert ( @@ -956,7 +981,9 @@ def test_motherduck_attach_catalog(make_config): def test_motherduck_attach_options(): options = DuckDBAttachOptions( - type="postgres", path="dbname=postgres user=postgres host=127.0.0.1", read_only=True + type="postgres", + path="dbname=postgres user=postgres host=127.0.0.1", + read_only=True, ) assert ( @@ -1041,7 +1068,9 @@ def test_motherduck_token_mask(make_config): # motherduck format assert config_1._mask_sensitive_data(config_1.database) == "whodunnit" assert ( - config_1._mask_sensitive_data(f"md:{config_1.database}?motherduck_token={config_1.token}") + config_1._mask_sensitive_data( + f"md:{config_1.database}?motherduck_token={config_1.token}" + ) == "md:whodunnit?motherduck_token=********" ) assert ( @@ -1051,7 +1080,9 @@ def test_motherduck_token_mask(make_config): == "md:whodunnit?attach_mode=single&motherduck_token=********" ) assert ( - config_2._mask_sensitive_data(f"md:{config_2.database}?motherduck_token={config_2.token}") + config_2._mask_sensitive_data( + f"md:{config_2.database}?motherduck_token={config_2.token}" + ) == "md:whodunnit?motherduck_token=********" ) assert ( @@ -1067,7 +1098,9 @@ def test_motherduck_token_mask(make_config): == "md:whodunnit?motherduck_token=********" ) assert ( - config_1._mask_sensitive_data("md:whodunnit?motherduck_token=longtoken123456789") + config_1._mask_sensitive_data( + "md:whodunnit?motherduck_token=longtoken123456789" + ) == "md:whodunnit?motherduck_token=********" ) assert ( @@ -1096,9 +1129,14 @@ def test_motherduck_token_mask(make_config): ) == "postgres:dbname=testdb password=******** user=admin" ) - assert config_1._mask_sensitive_data("postgres:password=short") == "postgres:password=********" assert ( - config_1._mask_sensitive_data("postgres:host=localhost password=p@ssw0rd! dbname=db") + config_1._mask_sensitive_data("postgres:password=short") + == "postgres:password=********" + ) + assert ( + config_1._mask_sensitive_data( + "postgres:host=localhost password=p@ssw0rd! dbname=db" + ) == "postgres:host=localhost password=******** dbname=db" ) @@ -1108,13 +1146,17 @@ def test_motherduck_token_mask(make_config): ) assert ( - config_1._mask_sensitive_data("md:db?motherduck_token=token123 postgres:password=secret") + config_1._mask_sensitive_data( + "md:db?motherduck_token=token123 postgres:password=secret" + ) == "md:db?motherduck_token=******** postgres:password=********" ) # MySQL format assert ( - config_1._mask_sensitive_data("host=localhost user=root password=mysql123 database=mydb") + config_1._mask_sensitive_data( + "host=localhost user=root password=mysql123 database=mydb" + ) == "host=localhost user=root password=******** database=mydb" ) @@ -1157,7 +1199,9 @@ def test_bigquery(make_config): ) with pytest.raises(ConfigError, match="you must also specify the `project` field"): - make_config(type="bigquery", execution_project="execution_project", check_import=False) + make_config( + type="bigquery", execution_project="execution_project", check_import=False + ) with pytest.raises(ConfigError, match="you must also specify the `project` field"): make_config(type="bigquery", quota_project="quota_project", check_import=False) @@ -1257,7 +1301,9 @@ def test_clickhouse(make_config): pool = config._connection_factory.keywords["pool_mgr"] assert pool.connection_pool_kw["server_hostname"] == "server_host_name" - assert pool.connection_pool_kw["assert_hostname"] == "server_host_name" # because verify=True + assert ( + pool.connection_pool_kw["assert_hostname"] == "server_host_name" + ) # because verify=True assert pool.connection_pool_kw["ca_certs"] == "ca_cert" assert pool.connection_pool_kw["cert_file"] == "client_cert" assert pool.connection_pool_kw["key_file"] == "client_cert_key" @@ -1329,7 +1375,9 @@ def test_athena_s3_staging_dir_or_workgroup(make_config): def test_athena_s3_locations_valid(make_config): with pytest.raises(ConfigError, match=r".*must be a s3:// URI.*"): make_config( - type="athena", work_group="primary", s3_warehouse_location="hdfs://legacy/location" + type="athena", + work_group="primary", + s3_warehouse_location="hdfs://legacy/location", ) with pytest.raises(ConfigError, match=r".*must be a s3:// URI.*"): @@ -1346,7 +1394,10 @@ def test_athena_s3_locations_valid(make_config): assert config.s3_warehouse_location == "s3://bucket/prod/warehouse/" config = make_config( - type="athena", work_group="primary", s3_staging_dir=None, s3_warehouse_location=None + type="athena", + work_group="primary", + s3_staging_dir=None, + s3_warehouse_location=None, ) assert isinstance(config, AthenaConnectionConfig) @@ -1403,9 +1454,13 @@ def test_databricks(make_config): assert oauth_u2m_config.oauth_client_secret is None # auth_type must match the AuthType enum if specified - with pytest.raises(ConfigError, match=r".*nonexist does not match a valid option.*"): + with pytest.raises( + ConfigError, match=r".*nonexist does not match a valid option.*" + ): make_config( - type="databricks", server_hostname="dbc-test.cloud.databricks.com", auth_type="nonexist" + type="databricks", + server_hostname="dbc-test.cloud.databricks.com", + auth_type="nonexist", ) # if client_secret is specified, client_id must also be specified @@ -1457,7 +1512,9 @@ def test_databricks__u2m_oauth__shared_connection_pool(make_config): @patch.object(DatabricksConnectionConfig, "_connection_factory_with_kwargs") -def test_databricks__m2m_oauth__connection_pool(mock_connection_factory_with_kwargs, make_config): +def test_databricks__m2m_oauth__connection_pool( + mock_connection_factory_with_kwargs, make_config +): from sqlmesh.utils.connection_pool import ThreadLocalConnectionPool config = make_config( @@ -1505,7 +1562,9 @@ def test_engine_import_validator(): ): class TestConfigA(PydanticModel): - _engine_import_validator = _get_engine_import_validator("missing", "bigquery") + _engine_import_validator = _get_engine_import_validator( + "missing", "bigquery" + ) TestConfigA() @@ -1540,7 +1599,8 @@ def test_engine_display_order(): This test ensures that those integers begin with 1, are unique, and are sequential. """ display_numbers = [ - info[0] for info in sorted(INIT_DISPLAY_INFO_TO_TYPE.values(), key=lambda x: x[0]) + info[0] + for info in sorted(INIT_DISPLAY_INFO_TO_TYPE.values(), key=lambda x: x[0]) ] assert display_numbers == list(range(1, len(display_numbers) + 1)) @@ -1556,7 +1616,9 @@ def test_mssql_engine_import_validator(): # Test MSSQL Python driver suggests mssql-python extra when import fails if SUPPORTS_MSSQL_PYTHON_DRIVER: - with pytest.raises(ConfigError, match=r"pip install \"sqlmesh\[mssql-python\]\""): + with pytest.raises( + ConfigError, match=r"pip install \"sqlmesh\[mssql-python\]\"" + ): with patch("importlib.import_module") as mock_import: mock_import.side_effect = ImportError("No module named 'mssql_python'") MSSQLConnectionConfig(host="localhost", driver="mssql-python") @@ -1588,7 +1650,9 @@ def test_mssql_connection_config_parameter_validation(make_config): assert config.driver == "pymssql" # Test explicit pyodbc driver - config = make_config(type="mssql", host="localhost", driver="pyodbc", check_import=False) + config = make_config( + type="mssql", host="localhost", driver="pyodbc", check_import=False + ) assert isinstance(config, MSSQLConnectionConfig) assert config.driver == "pyodbc" @@ -1601,7 +1665,9 @@ def test_mssql_connection_config_parameter_validation(make_config): assert config.driver == "mssql-python" # Test explicit pymssql driver - config = make_config(type="mssql", host="localhost", driver="pymssql", check_import=False) + config = make_config( + type="mssql", host="localhost", driver="pymssql", check_import=False + ) assert isinstance(config, MSSQLConnectionConfig) assert config.driver == "pymssql" @@ -1620,7 +1686,9 @@ def test_mssql_connection_config_parameter_validation(make_config): assert config.driver_name == "ODBC Driver 18 for SQL Server" assert config.trust_server_certificate is True assert config.encrypt is False - assert config.odbc_properties == {"Authentication": "ActiveDirectoryServicePrincipal"} + assert config.odbc_properties == { + "Authentication": "ActiveDirectoryServicePrincipal" + } # Test mssql-python specific parameters if SUPPORTS_MSSQL_PYTHON_DRIVER: @@ -1636,7 +1704,9 @@ def test_mssql_connection_config_parameter_validation(make_config): assert isinstance(config, MSSQLConnectionConfig) assert config.trust_server_certificate is True assert config.encrypt is False - assert config.odbc_properties == {"Authentication": "ActiveDirectoryServicePrincipal"} + assert config.odbc_properties == { + "Authentication": "ActiveDirectoryServicePrincipal" + } # Test pymssql specific parameters config = make_config( @@ -1655,7 +1725,9 @@ def test_mssql_connection_config_parameter_validation(make_config): def test_mssql_connection_kwargs_keys(): """Test _connection_kwargs_keys returns correct keys for each driver variant.""" # Test pymssql driver keys - config = MSSQLConnectionConfig(host="localhost", driver="pymssql", check_import=False) + config = MSSQLConnectionConfig( + host="localhost", driver="pymssql", check_import=False + ) pymssql_keys = config._connection_kwargs_keys expected_pymssql_keys = { "password", @@ -1674,7 +1746,9 @@ def test_mssql_connection_kwargs_keys(): assert pymssql_keys == expected_pymssql_keys # Test pyodbc driver keys - config = MSSQLConnectionConfig(host="localhost", driver="pyodbc", check_import=False) + config = MSSQLConnectionConfig( + host="localhost", driver="pyodbc", check_import=False + ) pyodbc_keys = config._connection_kwargs_keys expected_pyodbc_keys = { "password", @@ -1700,7 +1774,9 @@ def test_mssql_connection_kwargs_keys(): # Test mssql-python driver keys if SUPPORTS_MSSQL_PYTHON_DRIVER: - config = MSSQLConnectionConfig(host="localhost", driver="mssql-python", check_import=False) + config = MSSQLConnectionConfig( + host="localhost", driver="mssql-python", check_import=False + ) mssql_python_keys = config._connection_kwargs_keys expected_mssql_python_keys = { "password", @@ -1838,7 +1914,9 @@ def test_mssql_pyodbc_connection_string_minimal(): assert mock_pyodbc_connect.call_args[1]["autocommit"] is True -@pytest.mark.xfail(not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported") +@pytest.mark.xfail( + not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported" +) def test_mssql_mssql_python_connection_string_generation(): """Test mssql_python.connect gets invoked with the correct ODBC connection string.""" with patch("mssql_python.connect") as mock_mssql_python_connect: @@ -1887,7 +1965,9 @@ def test_mssql_mssql_python_connection_string_generation(): assert call_args[1]["autocommit"] is False -@pytest.mark.xfail(not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported") +@pytest.mark.xfail( + not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported" +) def test_mssql_mssql_python_connection_string_with_odbc_properties(): """Test mssql-python connection string includes custom ODBC properties.""" with patch("mssql_python.connect") as mock_mssql_python_connect: @@ -1936,7 +2016,9 @@ def test_mssql_mssql_python_connection_string_with_odbc_properties(): assert conn_str.count("ConnectRetryInterval") == 1 -@pytest.mark.xfail(not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported") +@pytest.mark.xfail( + not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported" +) def test_mssql_mssql_python_connection_string_minimal(): """Test mssql-python connection string with minimal configuration.""" with patch("mssql_python.connect") as mock_mssql_python_connect: @@ -2111,12 +2193,16 @@ def mock_add_output_converter(sql_type, converter_func): result = converter_func(binary_data) - expected_dt = datetime(2023, 1, 1, 12, 0, 0, 0, timezone(timedelta(hours=-8, minutes=0))) + expected_dt = datetime( + 2023, 1, 1, 12, 0, 0, 0, timezone(timedelta(hours=-8, minutes=0)) + ) assert result == expected_dt assert result.tzinfo == timezone(timedelta(hours=-8)) -@pytest.mark.xfail(not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported") +@pytest.mark.xfail( + not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported" +) def test_mssql_mssql_python_connection_datetimeoffset_handling(): """Test that the MSSQL mssql-python connection properly handles DATETIMEOFFSET conversion.""" import struct @@ -2189,7 +2275,9 @@ def mock_add_output_converter(sql_type, converter_func): assert result.tzinfo == timezone(timedelta(hours=5, minutes=30)) -@pytest.mark.xfail(not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported") +@pytest.mark.xfail( + not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported" +) def test_mssql_mssql_python_connection_negative_timezone_offset(): """Test DATETIMEOFFSET handling with negative timezone offset at connection level.""" import struct @@ -2239,7 +2327,9 @@ def mock_add_output_converter(sql_type, converter_func): result = converter_func(binary_data) - expected_dt = datetime(2023, 1, 1, 12, 0, 0, 0, timezone(timedelta(hours=-8, minutes=0))) + expected_dt = datetime( + 2023, 1, 1, 12, 0, 0, 0, timezone(timedelta(hours=-8, minutes=0)) + ) assert result == expected_dt assert result.tzinfo == timezone(timedelta(hours=-8)) @@ -2282,14 +2372,22 @@ def test_fabric_pyodbc_connection_config_parameter_validation(make_config): assert config.driver_name == "ODBC Driver 18 for SQL Server" assert config.trust_server_certificate is True assert config.encrypt is False - assert config.odbc_properties == {"Authentication": "ActiveDirectoryServicePrincipal"} + assert config.odbc_properties == { + "Authentication": "ActiveDirectoryServicePrincipal" + } # Test that specifying a different driver for Fabric raises an error - with pytest.raises(ConfigError, match=r"Input should be 'pyodbc' or 'mssql-python'"): - make_config(type="fabric", host="localhost", driver="pymssql", check_import=False) + with pytest.raises( + ConfigError, match=r"Input should be 'pyodbc' or 'mssql-python'" + ): + make_config( + type="fabric", host="localhost", driver="pymssql", check_import=False + ) -@pytest.mark.xfail(not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported") +@pytest.mark.xfail( + not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported" +) def test_fabric_mssql_python_connection_config_parameter_validation(make_config): """Test Fabric mssql-python connection config parameter validation.""" # Test that FabricConnectionConfig correctly handles mssql-python-specific parameters. @@ -2308,11 +2406,17 @@ def test_fabric_mssql_python_connection_config_parameter_validation(make_config) assert config.driver == "mssql-python" # Driver is fixed to mssql-python assert config.trust_server_certificate is True assert config.encrypt is False - assert config.odbc_properties == {"Authentication": "ActiveDirectoryServicePrincipal"} + assert config.odbc_properties == { + "Authentication": "ActiveDirectoryServicePrincipal" + } # Test that specifying a different driver for Fabric raises an error - with pytest.raises(ConfigError, match=r"Input should be 'pyodbc' or 'mssql-python'"): - make_config(type="fabric", host="localhost", driver="pymssql", check_import=False) + with pytest.raises( + ConfigError, match=r"Input should be 'pyodbc' or 'mssql-python'" + ): + make_config( + type="fabric", host="localhost", driver="pymssql", check_import=False + ) def test_fabric_pyodbc_connection_string_generation(): @@ -2362,7 +2466,9 @@ def test_fabric_pyodbc_connection_string_generation(): assert call_args[1]["autocommit"] is True -@pytest.mark.xfail(not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported") +@pytest.mark.xfail( + not SUPPORTS_MSSQL_PYTHON_DRIVER, reason="mssql-python driver not supported" +) def test_fabric_mssql_python_connection_string_generation(): """Test that the Fabric mssql-python connection gets invoked with the correct connection string.""" with patch("mssql_python.connect") as mock_mssql_python_connect: diff --git a/tests/core/test_context.py b/tests/core/test_context.py index e41382b078..59928dfa92 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -1,71 +1,54 @@ import logging import pathlib -import typing as t import re -from datetime import date, timedelta, datetime +import typing as t +from datetime import date, datetime, timedelta +from pathlib import Path from tempfile import TemporaryDirectory from unittest.mock import PropertyMock, call, patch -import time_machine -import pytest import pandas as pd # noqa: TID253 -from pathlib import Path +import pytest +import time_machine from pytest_mock.plugin import MockerFixture -from sqlglot import ParseError, exp, parse_one, Dialect +from sqlglot import Dialect, ParseError, exp, parse_one from sqlglot.errors import SchemaError import sqlmesh.core.constants from sqlmesh.cli.project_init import init_example_project -from sqlmesh.core.console import TerminalConsole -from sqlmesh.core import dialect as d, constants as c -from sqlmesh.core.config import ( - load_configs, - AutoCategorizationMode, - CategorizerConfig, - Config, - DuckDBConnectionConfig, - EnvironmentSuffixTarget, - GatewayConfig, - LinterConfig, - ModelDefaultsConfig, - PlanConfig, - SnowflakeConnectionConfig, -) +from sqlmesh.core import constants as c +from sqlmesh.core import dialect as d +from sqlmesh.core.config import (AutoCategorizationMode, CategorizerConfig, + Config, DuckDBConnectionConfig, + EnvironmentSuffixTarget, GatewayConfig, + LinterConfig, ModelDefaultsConfig, PlanConfig, + SnowflakeConnectionConfig, load_configs) +from sqlmesh.core.console import TerminalConsole, create_console, get_console from sqlmesh.core.context import Context -from sqlmesh.core.console import create_console, get_console from sqlmesh.core.dialect import parse, schema_ from sqlmesh.core.engine_adapter.duckdb import DuckDBEngineAdapter -from sqlmesh.core.environment import Environment, EnvironmentNamingInfo, EnvironmentStatements -from sqlmesh.core.plan.definition import Plan +from sqlmesh.core.environment import (Environment, EnvironmentNamingInfo, + EnvironmentStatements) from sqlmesh.core.macros import MacroEvaluator, RuntimeStage -from sqlmesh.core.model import load_sql_based_model, model, SqlModel, Model -from sqlmesh.core.model.common import ParsableSql +from sqlmesh.core.model import Model, SqlModel, load_sql_based_model, model from sqlmesh.core.model.cache import OptimizedQueryCache -from sqlmesh.core.snapshot import SnapshotChangeCategory -from sqlmesh.core.renderer import render_statements +from sqlmesh.core.model.common import ParsableSql from sqlmesh.core.model.kind import ModelKindName +from sqlmesh.core.plan.definition import Plan +from sqlmesh.core.renderer import render_statements +from sqlmesh.core.snapshot import SnapshotChangeCategory from sqlmesh.core.state_sync.cache import CachingStateSync from sqlmesh.core.state_sync.db import EngineAdapterStateSync -from sqlmesh.utils.connection_pool import SingletonConnectionPool, ThreadLocalSharedConnectionPool -from sqlmesh.utils.date import ( - make_inclusive_end, - now, - to_date, - to_datetime, - to_timestamp, - yesterday_ds, -) -from sqlmesh.utils.errors import ( - ConfigError, - SQLMeshError, - LinterError, - PlanError, - NoChangesPlanError, -) +from sqlmesh.utils.connection_pool import (SingletonConnectionPool, + ThreadLocalSharedConnectionPool) +from sqlmesh.utils.date import (make_inclusive_end, now, to_date, to_datetime, + to_timestamp, yesterday_ds) +from sqlmesh.utils.errors import (ConfigError, LinterError, NoChangesPlanError, + PlanError, SQLMeshError) from sqlmesh.utils.metaprogramming import Executable from sqlmesh.utils.windows import IS_WINDOWS, fix_windows_path -from tests.utils.test_helpers import use_terminal_console from tests.utils.test_filesystem import create_temp_file +from tests.utils.test_helpers import use_terminal_console def test_global_config(copy_to_temp_path: t.Callable): @@ -93,20 +76,29 @@ def test_missing_named_config(copy_to_temp_path: t.Callable): def test_config_parameter(copy_to_temp_path: t.Callable): - config = Config(model_defaults=ModelDefaultsConfig(dialect="presto"), project="test_project") + config = Config( + model_defaults=ModelDefaultsConfig(dialect="presto"), project="test_project" + ) context = Context(paths=copy_to_temp_path("examples/sushi"), config=config) assert context.config.dialect == "presto" assert context.config.project == "test_project" def test_generate_table_name_in_dialect(mocker: MockerFixture): - context = Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="bigquery"))) + context = Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="bigquery")) + ) mocker.patch( "sqlmesh.core.context.GenericContext._model_tables", - PropertyMock(return_value={'"project-id"."dataset"."table"': '"project-id".dataset.table'}), + PropertyMock( + return_value={ + '"project-id"."dataset"."table"': '"project-id".dataset.table' + } + ), ) assert ( - context.resolve_table('"project-id"."dataset"."table"') == "`project-id`.`dataset`.`table`" + context.resolve_table('"project-id"."dataset"."table"') + == "`project-id`.`dataset`.`table`" ) @@ -123,7 +115,9 @@ def test_custom_macros(sushi_context): def test_dag(sushi_context): - assert set(sushi_context.dag.upstream('"memory"."sushi"."customer_revenue_by_day"')) == { + assert set( + sushi_context.dag.upstream('"memory"."sushi"."customer_revenue_by_day"') + ) == { '"memory"."sushi"."items"', '"memory"."sushi"."orders"', '"memory"."sushi"."order_items"', @@ -209,15 +203,21 @@ def test_render_sql_model(sushi_context, assert_exp_eq, copy_to_temp_path: t.Cal @pytest.mark.slow -def test_render_non_deployable_parent(sushi_context, assert_exp_eq, copy_to_temp_path: t.Callable): +def test_render_non_deployable_parent( + sushi_context, assert_exp_eq, copy_to_temp_path: t.Callable +): model = sushi_context.get_model("sushi.waiter_revenue_by_day") forward_only_kind = model.kind.copy(update={"forward_only": True}) - model = model.copy(update={"kind": forward_only_kind, "stamp": "trigger forward-only change"}) + model = model.copy( + update={"kind": forward_only_kind, "stamp": "trigger forward-only change"} + ) sushi_context.upsert_model(model) sushi_context.plan("dev", no_prompts=True, auto_apply=True) expected_table_name = parse_one( - sushi_context.get_snapshot("sushi.waiter_revenue_by_day").table_name(is_deployable=False), + sushi_context.get_snapshot("sushi.waiter_revenue_by_day").table_name( + is_deployable=False + ), into=exp.Table, ).this.this @@ -281,7 +281,9 @@ def test_diff(sushi_context: Context, mocker: MockerFixture): yesterday = yesterday_ds() success = sushi_context.run(start=yesterday, end=yesterday) - sushi_context.upsert_model("sushi.customers", query=parse_one("select 1 as customer_id")) + sushi_context.upsert_model( + "sushi.customers", query=parse_one("select 1 as customer_id") + ) sushi_context.diff("test") assert mock_console.show_environment_difference_summary.called assert mock_console.show_model_difference_summary.called @@ -292,33 +294,43 @@ def test_diff(sushi_context: Context, mocker: MockerFixture): def test_evaluate_limit(): context = Context(config=Config()) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL(name with_limit, kind FULL); - SELECT t.v as v FROM (VALUES (1), (2), (3), (4), (5)) AS t(v) LIMIT 1 + 2""" - ) - ) - ) + SELECT t.v as v FROM (VALUES (1), (2), (3), (4), (5)) AS t(v) LIMIT 1 + 2"""))) - assert context.evaluate("with_limit", "2020-01-01", "2020-01-02", "2020-01-02").size == 3 - assert context.evaluate("with_limit", "2020-01-01", "2020-01-02", "2020-01-02", 4).size == 3 - assert context.evaluate("with_limit", "2020-01-01", "2020-01-02", "2020-01-02", 2).size == 2 + assert ( + context.evaluate("with_limit", "2020-01-01", "2020-01-02", "2020-01-02").size + == 3 + ) + assert ( + context.evaluate("with_limit", "2020-01-01", "2020-01-02", "2020-01-02", 4).size + == 3 + ) + assert ( + context.evaluate("with_limit", "2020-01-01", "2020-01-02", "2020-01-02", 2).size + == 2 + ) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL(name without_limit, kind FULL); - SELECT t.v as v FROM (VALUES (1), (2), (3), (4), (5)) AS t(v)""" - ) - ) - ) + SELECT t.v as v FROM (VALUES (1), (2), (3), (4), (5)) AS t(v)"""))) - assert context.evaluate("without_limit", "2020-01-01", "2020-01-02", "2020-01-02").size == 5 - assert context.evaluate("without_limit", "2020-01-01", "2020-01-02", "2020-01-02", 4).size == 4 - assert context.evaluate("without_limit", "2020-01-01", "2020-01-02", "2020-01-02", 2).size == 2 + assert ( + context.evaluate("without_limit", "2020-01-01", "2020-01-02", "2020-01-02").size + == 5 + ) + assert ( + context.evaluate( + "without_limit", "2020-01-01", "2020-01-02", "2020-01-02", 4 + ).size + == 4 + ) + assert ( + context.evaluate( + "without_limit", "2020-01-01", "2020-01-02", "2020-01-02", 2 + ).size + == 2 + ) def test_gateway_specific_adapters(copy_to_temp_path, mocker): @@ -349,24 +361,25 @@ def test_multiple_gateways(tmp_path: Path): context = Context(config=config) gateway_model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name staging.stg_model, start '2024-01-01',kind FULL, gateway staging); - SELECT t.v as v FROM (VALUES (1), (2), (3), (4), (5)) AS t(v)""" - ), + SELECT t.v as v FROM (VALUES (1), (2), (3), (4), (5)) AS t(v)"""), default_catalog="db", ) assert gateway_model.gateway == "staging" context.upsert_model(gateway_model) - assert context.evaluate("staging.stg_model", "2020-01-01", "2020-01-02", "2020-01-02").size == 5 + assert ( + context.evaluate( + "staging.stg_model", "2020-01-01", "2020-01-02", "2020-01-02" + ).size + == 5 + ) default_model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name main.final_model, start '2024-01-01',kind FULL); - SELECT v FROM staging.stg_model""" - ), + SELECT v FROM staging.stg_model"""), default_catalog="db", ) @@ -384,20 +397,28 @@ def test_multiple_gateways(tmp_path: Path): physical_schemas = [snapshot.physical_schema for snapshot in sorted_snapshots] assert physical_schemas == ["sqlmesh__main", "sqlmesh__staging"] - view_schemas = [snapshot.qualified_view_name.schema_name for snapshot in sorted_snapshots] + view_schemas = [ + snapshot.qualified_view_name.schema_name for snapshot in sorted_snapshots + ] assert view_schemas == ["main", "staging"] assert ( str(context.fetchdf("select * from staging.stg_model")) == " v\n0 1\n1 2\n2 3\n3 4\n4 5" ) - assert str(context.fetchdf("select * from final_model")) == " v\n0 1\n1 2\n2 3\n3 4\n4 5" + assert ( + str(context.fetchdf("select * from final_model")) + == " v\n0 1\n1 2\n2 3\n3 4\n4 5" + ) assert ( context.snapshots['"db"."main"."final_model"'].parents[0].name == '"db"."staging"."stg_model"' ) - assert context.dag._sorted == ['"db"."staging"."stg_model"', '"db"."main"."final_model"'] + assert context.dag._sorted == [ + '"db"."staging"."stg_model"', + '"db"."main"."final_model"', + ] def test_multi_gateway_catalog_aware_and_unsupported(tmp_path: Path, mocker): @@ -460,22 +481,26 @@ def test_multi_gateway_catalog_aware_and_unsupported(tmp_path: Path, mocker): # Loading models for both gateways must not raise a SchemaError. duckdb_model = load_sql_based_model( - parse("MODEL(name main.duckdb_tbl, kind FULL, gateway duckdb_gw);\nSELECT 1 AS col"), + parse( + "MODEL(name main.duckdb_tbl, kind FULL, gateway duckdb_gw);\nSELECT 1 AS col" + ), default_catalog="db", ) ch_model = load_sql_based_model( - parse("MODEL(name mydb.ch_tbl, kind FULL, gateway clickhouse_gw);\nSELECT 1 AS col"), + parse( + "MODEL(name mydb.ch_tbl, kind FULL, gateway clickhouse_gw);\nSELECT 1 AS col" + ), default_catalog="__clickhouse_gw__", ) # Both models must have 3-level FQNs so MappingSchema nesting is uniform. # count(".") == 2 means 3 parts (catalog.db.table), i.e. a 3-level FQN. - assert duckdb_model.fqn.count(".") == 2, ( - f"Expected 3-level FQN for duckdb model, got: {duckdb_model.fqn}" - ) - assert ch_model.fqn.count(".") == 2, ( - f"Expected 3-level FQN for ch model, got: {ch_model.fqn}" - ) # 3 parts = 2 dots + assert ( + duckdb_model.fqn.count(".") == 2 + ), f"Expected 3-level FQN for duckdb model, got: {duckdb_model.fqn}" + assert ( + ch_model.fqn.count(".") == 2 + ), f"Expected 3-level FQN for ch model, got: {ch_model.fqn}" # 3 parts = 2 dots # Both models loaded into the same MappingSchema must not raise a nesting SchemaError. from sqlglot.schema import MappingSchema @@ -590,7 +615,9 @@ def test_cleanup_environments_initializes_virtual_catalog_without_probing_unrela virtual_catalog=virtual_catalog, ) unavailable_adapter = make_mocked_engine_adapter(DuckDBEngineAdapter) - unavailable_adapter.cursor.execute.side_effect = RuntimeError("unrelated gateway unavailable") + unavailable_adapter.cursor.execute.side_effect = RuntimeError( + "unrelated gateway unavailable" + ) context = Context(config=Config(), load=False) context._engine_adapter = duck_adapter @@ -636,12 +663,16 @@ def test_cleanup_environments_initializes_virtual_catalog_without_probing_unrela @pytest.mark.fast def test_cleanup_environments_does_not_leak_historical_virtual_catalog( - mocker: MockerFixture, make_mocked_engine_adapter: t.Callable, make_snapshot: t.Callable + mocker: MockerFixture, + make_mocked_engine_adapter: t.Callable, + make_snapshot: t.Callable, ): """Historical cleanup catalogs must not mutate root or evaluator adapters.""" from sqlmesh.core.engine_adapter.clickhouse import ClickhouseEngineAdapter - duck_adapter = make_mocked_engine_adapter(DuckDBEngineAdapter, default_catalog="main") + duck_adapter = make_mocked_engine_adapter( + DuckDBEngineAdapter, default_catalog="main" + ) clickhouse_adapter = make_mocked_engine_adapter( ClickhouseEngineAdapter, virtual_catalog="new_catalog", @@ -683,19 +714,30 @@ def test_cleanup_environments_does_not_leak_historical_virtual_catalog( context._state_sync = state_sync evaluator_before_cleanup = context.snapshot_evaluator - assert evaluator_before_cleanup.adapters["clickhouse_gw"]._default_catalog == "new_catalog" + assert ( + evaluator_before_cleanup.adapters["clickhouse_gw"]._default_catalog + == "new_catalog" + ) assert context._cleanup_environments(name="dev") == [] assert clickhouse_adapter._default_catalog == "new_catalog" - assert evaluator_before_cleanup.adapters["clickhouse_gw"]._default_catalog == "new_catalog" + assert ( + evaluator_before_cleanup.adapters["clickhouse_gw"]._default_catalog + == "new_catalog" + ) context._snapshot_evaluator = None - assert context.snapshot_evaluator.adapters["clickhouse_gw"]._default_catalog == "new_catalog" + assert ( + context.snapshot_evaluator.adapters["clickhouse_gw"]._default_catalog + == "new_catalog" + ) @pytest.mark.fast def test_cleanup_environments_supports_legacy_virtual_catalog_adapter( - mocker: MockerFixture, make_mocked_engine_adapter: t.Callable, make_snapshot: t.Callable + mocker: MockerFixture, + make_mocked_engine_adapter: t.Callable, + make_snapshot: t.Callable, ): """Cleanup must support opt-in adapters whose injection override accepts only a gateway.""" from sqlmesh.core.engine_adapter.clickhouse import ClickhouseEngineAdapter @@ -704,7 +746,9 @@ class LegacyVirtualCatalogAdapter(ClickhouseEngineAdapter): def inject_virtual_catalog(self, gateway: str) -> None: self._default_catalog = f"__{gateway}__" - duck_adapter = make_mocked_engine_adapter(DuckDBEngineAdapter, default_catalog="main") + duck_adapter = make_mocked_engine_adapter( + DuckDBEngineAdapter, default_catalog="main" + ) legacy_adapter = make_mocked_engine_adapter(LegacyVirtualCatalogAdapter) context = Context(config=Config(), load=False) @@ -749,7 +793,9 @@ def inject_virtual_catalog(self, gateway: str) -> None: @pytest.mark.fast def test_cleanup_environments_initializes_scoped_third_party_adapter( - mocker: MockerFixture, make_mocked_engine_adapter: t.Callable, make_snapshot: t.Callable + mocker: MockerFixture, + make_mocked_engine_adapter: t.Callable, + make_snapshot: t.Callable, ): """Cleanup clones must run the virtual-catalog hook before restoring persisted state.""" from sqlmesh.core.engine_adapter.clickhouse import ClickhouseEngineAdapter @@ -775,7 +821,9 @@ def inject_virtual_catalog(self, gateway: str) -> None: self.virtual_catalog_enabled = True self._default_catalog = f"__{gateway}__" - duck_adapter = make_mocked_engine_adapter(DuckDBEngineAdapter, default_catalog="main") + duck_adapter = make_mocked_engine_adapter( + DuckDBEngineAdapter, default_catalog="main" + ) stateful_adapter = make_mocked_engine_adapter(StatefulVirtualCatalogAdapter) context = Context(config=Config(), load=False) @@ -821,12 +869,16 @@ def inject_virtual_catalog(self, gateway: str) -> None: @pytest.mark.fast def test_cleanup_environments_rejects_ambiguous_persisted_virtual_catalogs( - mocker: MockerFixture, make_mocked_engine_adapter: t.Callable, make_snapshot: t.Callable + mocker: MockerFixture, + make_mocked_engine_adapter: t.Callable, + make_snapshot: t.Callable, ): """Cleanup must not partially drop views when one gateway has multiple persisted catalogs.""" from sqlmesh.core.engine_adapter.clickhouse import ClickhouseEngineAdapter - duck_adapter = make_mocked_engine_adapter(DuckDBEngineAdapter, default_catalog="main") + duck_adapter = make_mocked_engine_adapter( + DuckDBEngineAdapter, default_catalog="main" + ) clickhouse_adapter = make_mocked_engine_adapter( ClickhouseEngineAdapter, virtual_catalog="old_catalog", @@ -879,7 +931,9 @@ def test_cleanup_environments_rejects_ambiguous_persisted_virtual_catalogs( @pytest.mark.fast -def test_multi_gateway_virtual_catalog_create_schema_strips_prefix(tmp_path: Path, mocker): +def test_multi_gateway_virtual_catalog_create_schema_strips_prefix( + tmp_path: Path, mocker +): """Integration test: create_schema with a 3-level virtual-catalog FQN must strip the synthetic catalog prefix before sending DDL to ClickHouse. @@ -940,22 +994,26 @@ def test_multi_gateway_virtual_catalog_create_schema_strips_prefix(tmp_path: Pat # --- Phase 2: FQN uniformity --- ch_model = load_sql_based_model( - parse("MODEL(name mydb.ch_tbl, kind FULL, gateway clickhouse_gw);\nSELECT 1 AS col"), + parse( + "MODEL(name mydb.ch_tbl, kind FULL, gateway clickhouse_gw);\nSELECT 1 AS col" + ), default_catalog="__clickhouse_gw__", ) duckdb_model = load_sql_based_model( - parse("MODEL(name main.duckdb_tbl, kind FULL, gateway duckdb_gw);\nSELECT 1 AS col"), + parse( + "MODEL(name main.duckdb_tbl, kind FULL, gateway duckdb_gw);\nSELECT 1 AS col" + ), default_catalog=catalog_per_gw["duckdb_gw"], ) # Both models must have 3-level FQNs (catalog.db.table → 2 dots) so MappingSchema nesting # is uniform and does not raise a SchemaError. - assert ch_model.fqn.count(".") == 2, ( - f"Expected 3-level FQN for ClickHouse model, got: {ch_model.fqn}" - ) - assert duckdb_model.fqn.count(".") == 2, ( - f"Expected 3-level FQN for DuckDB model, got: {duckdb_model.fqn}" - ) + assert ( + ch_model.fqn.count(".") == 2 + ), f"Expected 3-level FQN for ClickHouse model, got: {ch_model.fqn}" + assert ( + duckdb_model.fqn.count(".") == 2 + ), f"Expected 3-level FQN for DuckDB model, got: {duckdb_model.fqn}" from sqlglot.schema import MappingSchema @@ -981,7 +1039,9 @@ def _capture_create_schema( else str(schema_name) ) - mocker.patch.object(ch_adapter, "_create_schema", side_effect=_capture_create_schema) + mocker.patch.object( + ch_adapter, "_create_schema", side_effect=_capture_create_schema + ) # Call create_schema with the 3-level virtual-catalog-prefixed schema name. ch_adapter.create_schema("__clickhouse_gw__.mydb") @@ -989,9 +1049,9 @@ def _capture_create_schema( assert len(create_schema_calls) == 1, "Expected exactly one _create_schema call" passed_schema = create_schema_calls[0] # The virtual catalog prefix must NOT appear in the SQL sent to the wire. - assert "__clickhouse_gw__" not in passed_schema, ( - f"Virtual catalog prefix should be stripped before reaching _create_schema, got: {passed_schema!r}" - ) + assert ( + "__clickhouse_gw__" not in passed_schema + ), f"Virtual catalog prefix should be stripped before reaching _create_schema, got: {passed_schema!r}" @pytest.mark.fast @@ -1016,7 +1076,10 @@ def test_warn_if_virtual_catalog_rematerialization_emits_warning(mocker): # Override engine_adapters so the context sees our prepared adapter. mocker.patch.object( - type(ctx), "engine_adapters", new_callable=PropertyMock, return_value={"ch_gw": ch_adapter} + type(ctx), + "engine_adapters", + new_callable=PropertyMock, + return_value={"ch_gw": ch_adapter}, ) # Build a mock snapshot with a 3-level name that has the virtual catalog prefix. @@ -1046,7 +1109,9 @@ def test_warn_if_virtual_catalog_rematerialization_emits_warning(mocker): @pytest.mark.fast -def test_warn_if_virtual_catalog_rematerialization_no_warning_when_genuinely_new(mocker): +def test_warn_if_virtual_catalog_rematerialization_no_warning_when_genuinely_new( + mocker, +): """_warn_if_virtual_catalog_rematerialization must NOT warn when there is no matching old 2-level name — i.e. the model is a brand-new model, not a renamed existing one.""" from unittest.mock import MagicMock @@ -1062,7 +1127,10 @@ def test_warn_if_virtual_catalog_rematerialization_no_warning_when_genuinely_new ch_adapter._default_catalog = "__ch_gw__" mocker.patch.object( - type(ctx), "engine_adapters", new_callable=PropertyMock, return_value={"ch_gw": ch_adapter} + type(ctx), + "engine_adapters", + new_callable=PropertyMock, + return_value={"ch_gw": ch_adapter}, ) new_snapshot = MagicMock() @@ -1086,7 +1154,9 @@ def test_warn_if_virtual_catalog_rematerialization_no_warning_when_genuinely_new @pytest.mark.fast -def test_warn_if_virtual_catalog_rematerialization_no_warning_without_virtual_catalog(mocker): +def test_warn_if_virtual_catalog_rematerialization_no_warning_without_virtual_catalog( + mocker, +): """_warn_if_virtual_catalog_rematerialization must NOT warn when the ClickHouse adapter has no virtual catalog injected (i.e. _default_catalog is None).""" from unittest.mock import MagicMock @@ -1103,7 +1173,10 @@ def test_warn_if_virtual_catalog_rematerialization_no_warning_without_virtual_ca assert ch_adapter._default_catalog is None mocker.patch.object( - type(ctx), "engine_adapters", new_callable=PropertyMock, return_value={"ch_gw": ch_adapter} + type(ctx), + "engine_adapters", + new_callable=PropertyMock, + return_value={"ch_gw": ch_adapter}, ) new_snapshot = MagicMock() @@ -1127,10 +1200,7 @@ def test_warn_if_virtual_catalog_rematerialization_no_warning_without_virtual_ca def test_plan_execution_time(): context = Context(config=Config()) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name db.x, start '2024-01-01', @@ -1138,10 +1208,7 @@ def test_plan_execution_time(): ); SELECT @execution_date AS execution_date - """ - ) - ) - ) + """))) context.plan( "dev", @@ -1157,10 +1224,7 @@ def test_plan_execution_time(): def test_plan_execution_time_start_end(): context = Context(config=Config()) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name db.x, start '2020-01-01', @@ -1178,20 +1242,14 @@ def test_plan_execution_time_start_end(): ('5', '2024-01-01') ) data(id, ds) WHERE ds BETWEEN @start_ds AND @end_ds - """ - ) - ) - ) + """))) # prod plan - no fixed execution time so it defaults to now() and reads all the data prod_plan = context.plan(auto_apply=True) assert len(prod_plan.new_snapshots) == 1 - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name db.x, start '2020-01-01', @@ -1209,10 +1267,7 @@ def test_plan_execution_time_start_end(): ('5', '2024-01-01') ) data(id, ds) WHERE ds BETWEEN @start_ds AND @end_ds - """ - ) - ) - ) + """))) # dev plan with an execution time in the past and no explicit start/end specified # the plan end should be bounded to it and not exceed it even though in prod the last interval (used as a default end) @@ -1239,21 +1294,22 @@ def test_plan_execution_time_start_end(): ) # end should not be greater than execution_time # same as above but with a relative start and a relative end - dev_plan = context.plan("dev", start="2 days ago", execution_time="2020-01-05", end="1 day ago") + dev_plan = context.plan( + "dev", start="2 days ago", execution_time="2020-01-05", end="1 day ago" + ) assert to_datetime(dev_plan.start) == to_datetime( "2020-01-03" ) # start relative to execution_time assert to_datetime(dev_plan.execution_time) == to_datetime("2020-01-05") - assert to_datetime(dev_plan.end) == to_datetime("2020-01-04") # end relative to execution_time + assert to_datetime(dev_plan.end) == to_datetime( + "2020-01-04" + ) # end relative to execution_time def test_override_builtin_audit_blocking_mode(): context = Context(config=Config()) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name db.x, kind FULL, @@ -1264,17 +1320,15 @@ def test_override_builtin_audit_blocking_mode(): ); SELECT NULL AS c - """ - ) - ) - ) + """))) with patch.object(context.console, "log_warning") as mock_logger: plan = context.plan(auto_apply=True, no_prompts=True) new_snapshot = next(iter(plan.context_diff.new_snapshots.values())) assert ( - mock_logger.call_args_list[0][0][0] == "\ndb.x: 'not_null' audit error: 1 row failed." + mock_logger.call_args_list[0][0][0] + == "\ndb.x: 'not_null' audit error: 1 row failed." ) # Even though there are two builtin audits referenced in the above definition, we only @@ -1286,10 +1340,7 @@ def test_override_builtin_audit_blocking_mode(): assert list(args) == ["columns", "blocking"] context = Context(config=Config()) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name db.x, kind FULL, @@ -1299,10 +1350,7 @@ def test_override_builtin_audit_blocking_mode(): ); SELECT NULL AS c - """ - ) - ) - ) + """))) with pytest.raises(SQLMeshError): context.plan(auto_apply=True, no_prompts=True) @@ -1337,23 +1385,19 @@ def test_env_and_default_schema_normalization(mocker: MockerFixture): from sqlglot.dialects import DuckDB from sqlglot.dialects.dialect import NormalizationStrategy - mocker.patch.object(DuckDB, "NORMALIZATION_STRATEGY", NormalizationStrategy.UPPERCASE) + mocker.patch.object( + DuckDB, "NORMALIZATION_STRATEGY", NormalizationStrategy.UPPERCASE + ) context = Context(config=Config()) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name x, kind FULL ); SELECT 1 AS c - """ - ) - ) - ) + """))) context.plan("dev", auto_apply=True, no_prompts=True) assert list(context.fetchdf('select c from "DEFAULT__DEV"."X"')["c"])[0] == 1 @@ -1398,7 +1442,10 @@ def test_jinja_macro_undefined_variable_error(tmp_path: pathlib.Path): error_message = str(exc_info.value) assert "Failed to load model" in error_message assert "Could not render jinja for" in error_message - assert "Undefined macro/variable: 'target' in macro: 'generate_select'" in error_message + assert ( + "Undefined macro/variable: 'target' in macro: 'generate_select'" + in error_message + ) def test_clear_caches(tmp_path: pathlib.Path): @@ -1478,7 +1525,9 @@ def test_cache_path_configurations(tmp_path: pathlib.Path): # Test absolute path abs_cache = tmp_path / "abs_cache" - config_file.write_text(f"model_defaults:\n dialect: duckdb\ncache_dir: {abs_cache}") + config_file.write_text( + f"model_defaults:\n dialect: duckdb\ncache_dir: {abs_cache}" + ) context = Context(paths=str(project_dir)) assert context.cache_dir == abs_cache @@ -1607,7 +1656,11 @@ def test(): """, ) config = Config( - ignore_patterns=["models/ignore/**/*.sql", "macro_ignore.py", ".ipynb_checkpoints/*"] + ignore_patterns=[ + "models/ignore/**/*.sql", + "macro_ignore.py", + ".ipynb_checkpoints/*", + ] ) context = Context(paths=tmp_path, config=config) @@ -1680,7 +1733,9 @@ def test_project_config_person_config_overrides(tmp_path: pathlib.Path): def test_physical_schema_override(copy_to_temp_path: t.Callable) -> None: def get_schemas(context: Context): return { - snapshot.physical_schema for snapshot in context.snapshots.values() if snapshot.is_model + snapshot.physical_schema + for snapshot in context.snapshots.values() + if snapshot.is_model } def get_view_schemas(context: Context): @@ -1704,7 +1759,9 @@ def get_sushi_fingerprints(context: Context): assert get_view_schemas(no_mapping_context) == {"sushi", "raw"} no_mapping_fingerprints = get_sushi_fingerprints(no_mapping_context) context = Context(paths=project_path, config="map_config") - assert context.config.physical_schema_mapping == {re.compile("^sushi$"): "company_internal"} + assert context.config.physical_schema_mapping == { + re.compile("^sushi$"): "company_internal" + } assert get_schemas(context) == {"company_internal", "sqlmesh__raw"} assert get_view_schemas(context) == {"sushi", "raw"} sushi_fingerprints = get_sushi_fingerprints(context) @@ -1750,10 +1807,13 @@ def test_physical_schema_mapping(tmp_path: pathlib.Path) -> None: ctx.load() - physical_schemas = [snapshot.physical_schema for snapshot in sorted(ctx.snapshots.values())] + physical_schemas = [ + snapshot.physical_schema for snapshot in sorted(ctx.snapshots.values()) + ] view_schemas = [ - snapshot.qualified_view_name.schema_name for snapshot in sorted(ctx.snapshots.values()) + snapshot.qualified_view_name.schema_name + for snapshot in sorted(ctx.snapshots.values()) ] assert len(physical_schemas) == len(view_schemas) == 3 @@ -1788,7 +1848,9 @@ def test_janitor(sushi_context, mocker: MockerFixture) -> None: ), ] - state_sync_mock.get_expired_environments.return_value = [env.summary for env in environments] + state_sync_mock.get_expired_environments.return_value = [ + env.summary for env in environments + ] state_sync_mock.get_environment = lambda name: next( env for env in environments if env.name == name ) @@ -1881,7 +1943,9 @@ def get_expired(current_ts: int, name: t.Optional[str] = None) -> t.List: @pytest.mark.slow -def test_janitor_environment_not_expired_warning(sushi_context, mocker: MockerFixture) -> None: +def test_janitor_environment_not_expired_warning( + sushi_context, mocker: MockerFixture +) -> None: """Janitor with --environment emits a warning when the named environment is not expired.""" state_sync_mock = mocker.patch.object( type(sushi_context), "state_sync", new_callable=mocker.PropertyMock @@ -1916,7 +1980,9 @@ def test_invalidate_environment_sync_calls_cleanup_with_name( @pytest.mark.slow -def test_invalidate_environment_no_sync_skips_cleanup(sushi_context, mocker: MockerFixture) -> None: +def test_invalidate_environment_no_sync_skips_cleanup( + sushi_context, mocker: MockerFixture +) -> None: """invalidate_environment(..., sync=False) should not trigger _cleanup_environments at all.""" state_sync_mock = mocker.patch.object( type(sushi_context), "state_sync", new_callable=mocker.PropertyMock @@ -1928,7 +1994,9 @@ def test_invalidate_environment_no_sync_skips_cleanup(sushi_context, mocker: Moc state_sync_mock.delete_expired_environments.assert_not_called() -def test_invalidate_environment_nonexistent_raises(sushi_context, mocker: MockerFixture) -> None: +def test_invalidate_environment_nonexistent_raises( + sushi_context, mocker: MockerFixture +) -> None: """Invalidating an environment that does not exist should error instead of reporting success, so a mistyped name is caught rather than silently accepted.""" state_sync_mock = mocker.patch.object( @@ -1971,7 +2039,11 @@ def test_plan_default_end(sushi_context_pre_scheduling: Context): prod_plan_builder.apply() dev_plan = sushi_context_pre_scheduling.plan( - "test_env", no_prompts=True, include_unmodified=True, skip_backfill=True, auto_apply=True + "test_env", + no_prompts=True, + include_unmodified=True, + skip_backfill=True, + auto_apply=True, ) assert dev_plan.end is not None assert to_date(make_inclusive_end(dev_plan.end)) == plan_end @@ -1997,8 +2069,7 @@ def test_plan_start_ahead_of_end(copy_to_temp_path): context.close() with time_machine.travel("2024-01-03 00:00:00 UTC"): context = Context(paths=path, gateway="duckdb_persistent") - expression = d.parse( - """ + expression = d.parse(""" MODEL( name sushi.hourly, KIND FULL, @@ -2006,9 +2077,10 @@ def test_plan_start_ahead_of_end(copy_to_temp_path): start '2024-01-02 12:00:00', ); - SELECT 1""" + SELECT 1""") + model = load_sql_based_model( + expression, default_catalog=context.default_catalog ) - model = load_sql_based_model(expression, default_catalog=context.default_catalog) context.upsert_model(model) context.plan("prod", no_prompts=True, auto_apply=True) # Since the new start is ahead of the latest end loaded for prod, the table is deployed as empty @@ -2018,7 +2090,9 @@ def test_plan_start_ahead_of_end(copy_to_temp_path): i == to_timestamp("2024-01-02") for i in context.state_sync.max_interval_end_per_model("prod").values() ) - assert context.engine_adapter.fetchone("SELECT COUNT(*) FROM sushi.hourly")[0] == 0 + assert ( + context.engine_adapter.fetchone("SELECT COUNT(*) FROM sushi.hourly")[0] == 0 + ) context.close() @@ -2080,14 +2154,11 @@ def _write_daily_and_weekly_model_project(tmp_path: Path) -> None: report's project topology looked like. """ (tmp_path / "models").mkdir() - (tmp_path / "config.yaml").write_text( - """ + (tmp_path / "config.yaml").write_text(""" model_defaults: dialect: duckdb -""" - ) - (tmp_path / "models" / "daily_model.sql").write_text( - """ +""") + (tmp_path / "models" / "daily_model.sql").write_text(""" MODEL ( name daily_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -2098,10 +2169,8 @@ def _write_daily_and_weekly_model_project(tmp_path: Path) -> None: ); select @start_ds as start_ds, @end_ds as end_ds, @start_dt as start_dt, @end_dt as end_dt; -""" - ) - (tmp_path / "models" / "weekly_model.sql").write_text( - """ +""") + (tmp_path / "models" / "weekly_model.sql").write_text(""" MODEL ( name weekly_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -2112,15 +2181,20 @@ def _write_daily_and_weekly_model_project(tmp_path: Path) -> None: ); select @start_ds as start_ds, @end_ds as end_ds, @start_dt as start_dt, @end_dt as end_dt; -""" - ) +""") -def _missing_intervals_by_name(plan: Plan) -> t.Dict[str, t.Tuple[t.Tuple[int, int], ...]]: - return {si.snapshot_id.name: tuple(si.merged_intervals) for si in plan.missing_intervals} +def _missing_intervals_by_name( + plan: Plan, +) -> t.Dict[str, t.Tuple[t.Tuple[int, int], ...]]: + return { + si.snapshot_id.name: tuple(si.merged_intervals) for si in plan.missing_intervals + } -def test_plan_execution_time_ahead_of_prod_frontier_matches_run_for_all_models(tmp_path: Path): +def test_plan_execution_time_ahead_of_prod_frontier_matches_run_for_all_models( + tmp_path: Path, +): """Locks in that raising `max_interval_end_per_model` for an explicitly provided `execution_time` sweeps in *every* model with a recorded prod frontier, not just modified/selected ones. This is intentional, not an oversight: it's what makes a plain, @@ -2156,7 +2230,9 @@ def test_plan_execution_time_ahead_of_prod_frontier_matches_run_for_all_models(t # An equivalent `plan --run` at the same execution_time computes missing intervals with no # caps at all. If it matches exactly, that confirms the plain-plan raise reproduces the # `--run` result rather than under- or over-shooting it. - run_plan = context.plan_builder("prod", execution_time=execution_time, run=True).build() + run_plan = context.plan_builder( + "prod", execution_time=execution_time, run=True + ).build() assert run_plan.requires_backfill assert _missing_intervals_by_name(run_plan) == missing @@ -2210,8 +2286,7 @@ def test_plan_seed_model_excluded_from_default_end(copy_to_temp_path: t.Callable # a model that depends on this seed but has no interval in prod yet so only the seed would contribute to max_interval_end_per_model context.upsert_model( load_sql_based_model( - parse( - """ + parse(""" MODEL( name sushi.waiter_summary, kind INCREMENTAL_BY_TIME_RANGE ( @@ -2229,8 +2304,7 @@ def test_plan_seed_model_excluded_from_default_end(copy_to_temp_path: t.Callable sushi.waiter_names WHERE @start_ds BETWEEN @start_ds AND @end_ds - """ - ), + """), default_catalog=context.default_catalog, ) ) @@ -2243,7 +2317,10 @@ def test_plan_seed_model_excluded_from_default_end(copy_to_temp_path: t.Callable # the plan start date 2025-01-01 is after the seeds end date but shouldnt cause the plan to fail plan = context.plan( - "dev", start="2025-01-01", no_prompts=True, select_models=["*waiter_summary"] + "dev", + start="2025-01-01", + no_prompts=True, + select_models=["*waiter_summary"], ) # the end should fall back to execution_time rather than seeds end @@ -2263,37 +2340,27 @@ def test_schema_error_no_default(sushi_context_pre_scheduling) -> None: context = sushi_context_pre_scheduling with pytest.raises(SchemaError): - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL(name c); SELECT x FROM a - """ - ) - ) - ) + """))) @pytest.mark.slow def test_unrestorable_snapshot(sushi_context: Context) -> None: model_v1 = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_unrestorable); SELECT 1 AS one; - """ - ), + """), default_catalog=sushi_context.default_catalog, dialect=sushi_context.default_dialect, ) model_v2 = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_unrestorable); SELECT 2 AS two; - """ - ), + """), default_catalog=sushi_context.default_catalog, dialect=sushi_context.default_dialect, ) @@ -2396,7 +2463,9 @@ def _get_external_model_names(gateway=None): context = Context(paths=path, config="isolated_systems_config", gateway=gateway) external_model_names = [ - m.name for m in context.models.values() if m.kind.name == ModelKindName.EXTERNAL + m.name + for m in context.models.values() + if m.kind.name == ModelKindName.EXTERNAL ] assert len(external_model_names) > 0 @@ -2453,15 +2522,13 @@ def test_get_model_mixed_dialects(copy_to_temp_path): path = copy_to_temp_path("examples/sushi") context = Context(paths=path) - expression = d.parse( - """ + expression = d.parse(""" MODEL( name sushi.snowflake_dialect, dialect snowflake, ); - SELECT 1""" - ) + SELECT 1""") model = load_sql_based_model(expression, default_catalog=context.default_catalog) context.upsert_model(model) @@ -2470,7 +2537,9 @@ def test_get_model_mixed_dialects(copy_to_temp_path): def test_override_dialect_normalization_strategy(): config = Config( - model_defaults=ModelDefaultsConfig(dialect="duckdb,normalization_strategy=lowercase") + model_defaults=ModelDefaultsConfig( + dialect="duckdb,normalization_strategy=lowercase" + ) ) # This has the side-effect of mutating DuckDB globally to override its normalization strategy @@ -2576,14 +2645,18 @@ def test_duckdb_state_connection_automatic_multithreaded_mode(tmp_path): state_sync = context.state_sync.state_sync assert isinstance(state_sync, EngineAdapterStateSync) assert isinstance(state_sync.engine_adapter, DuckDBEngineAdapter) - assert isinstance(state_sync.engine_adapter._connection_pool, SingletonConnectionPool) + assert isinstance( + state_sync.engine_adapter._connection_pool, SingletonConnectionPool + ) context = Context(paths=[tmp_path], config=multi_threaded_config) assert isinstance(context.state_sync, CachingStateSync) state_sync = context.state_sync.state_sync assert isinstance(state_sync, EngineAdapterStateSync) assert isinstance(state_sync.engine_adapter, DuckDBEngineAdapter) - assert isinstance(state_sync.engine_adapter._connection_pool, ThreadLocalSharedConnectionPool) + assert isinstance( + state_sync.engine_adapter._connection_pool, ThreadLocalSharedConnectionPool + ) def test_requirements(copy_to_temp_path: t.Callable): @@ -2614,7 +2687,11 @@ def test_requirements(copy_to_temp_path: t.Callable): context._requirements = {"numpy": "2.1.2", "pandas": "2.2.1"} context._excluded_requirements = {"ipywidgets", "ruamel.yaml", "ruamel.yaml.clib"} - diff = context.plan_builder("dev", skip_tests=True, skip_backfill=True).build().context_diff + diff = ( + context.plan_builder("dev", skip_tests=True, skip_backfill=True) + .build() + .context_diff + ) assert set(diff.previous_requirements) == requirements reqs = {"numpy", "pandas"} if IS_WINDOWS: @@ -2639,10 +2716,7 @@ def test_deactivate_automatic_requirement_inference(copy_to_temp_path: t.Callabl def test_rendered_diff(): ctx = Context(config=Config()) - ctx.upsert_model( - load_sql_based_model( - parse( - """ + ctx.upsert_model(load_sql_based_model(parse(""" MODEL ( name test, ); @@ -2657,18 +2731,12 @@ def test_rendered_diff(): DROP VIEW @this_model ON_VIRTUAL_UPDATE_END; - """ - ) - ) - ) + """))) ctx.plan("dev", auto_apply=True, no_prompts=True) # Alter the model's query and pre/post/virtual statements to cause the diff - ctx.upsert_model( - load_sql_based_model( - parse( - """ + ctx.upsert_model(load_sql_based_model(parse(""" MODEL ( name test, ); @@ -2682,10 +2750,7 @@ def test_rendered_diff(): ON_VIRTUAL_UPDATE_BEGIN; DROP VIEW IF EXISTS @this_model ON_VIRTUAL_UPDATE_END; - """ - ) - ) - ) + """))) plan = ctx.plan("dev", auto_apply=True, no_prompts=True, diff_rendered=True) @@ -2715,7 +2780,9 @@ def test_rendered_diff(): ) -def test_plan_enable_preview_default(sushi_context: Context, sushi_dbt_context: Context): +def test_plan_enable_preview_default( + sushi_context: Context, sushi_dbt_context: Context +): assert sushi_context._plan_preview_enabled assert not sushi_dbt_context._plan_preview_enabled @@ -2747,7 +2814,9 @@ def test_raw_code_handling(sushi_test_dbt_context: Context): def test_dbt_models_are_not_validated(sushi_test_dbt_context: Context): model = sushi_test_dbt_context.models['"memory"."sushi"."non_validated_model"'] - assert model.render_query_or_raise().sql(comments=False) == 'SELECT 1 AS "c", 2 AS "c"' + assert ( + model.render_query_or_raise().sql(comments=False) == 'SELECT 1 AS "c", 2 AS "c"' + ) assert sushi_test_dbt_context.fetchdf( 'SELECT * FROM "memory"."sushi"."non_validated_model"' ).to_dict() == {"c": {0: 1}, "c_1": {0: 2}} @@ -2782,7 +2851,9 @@ def test_catalog_name_needs_to_be_quoted(): ) context = Context(config=config) parsed_model = parse("MODEL(name db.x, kind FULL); SELECT 1 AS c") - context.upsert_model(load_sql_based_model(parsed_model, default_catalog='"foo--bar"')) + context.upsert_model( + load_sql_based_model(parsed_model, default_catalog='"foo--bar"') + ) context.plan(auto_apply=True, no_prompts=True) assert context.fetchdf('select * from "foo--bar".db.x').to_dict() == {"c": {0: 1}} @@ -2808,7 +2879,9 @@ def test_plan_runs_audits_on_dev_previews(sushi_context: Context, capsys, caplog """ sushi_context.upsert_model( - load_sql_based_model(parse(test_model), default_catalog=sushi_context.default_catalog) + load_sql_based_model( + parse(test_model), default_catalog=sushi_context.default_catalog + ) ) plan = sushi_context.plan(auto_apply=True) @@ -2833,7 +2906,9 @@ def test_plan_runs_audits_on_dev_previews(sushi_context: Context, capsys, caplog """ sushi_context.upsert_model( - load_sql_based_model(parse(test_model), default_catalog=sushi_context.default_catalog) + load_sql_based_model( + parse(test_model), default_catalog=sushi_context.default_catalog + ) ) capsys.readouterr() # clear output buffer @@ -3106,7 +3181,8 @@ def access_adapter(evaluator): "evaluation_time", } assert ( - stats_table["physical_table"][0] == f"sqlmesh__db.db__test_stats_model__{snapshot.version}" + stats_table["physical_table"][0] + == f"sqlmesh__db.db__test_stats_model__{snapshot.version}" ) assert context.fetchdf("select * from memory.after_table").to_dict()["5"][0] == 5 @@ -3126,7 +3202,9 @@ def test_environment_statements_dialect(tmp_path: Path): before_all = [ "EXPORT DATA OPTIONS (URI='gs://path*.csv.gz', FORMAT='CSV') AS SELECT * FROM all_rows" ] - after_all = ["@IF(@this_env = 'prod', CREATE TABLE IF NOT EXISTS after_t AS SELECT 1)"] + after_all = [ + "@IF(@this_env = 'prod', CREATE TABLE IF NOT EXISTS after_t AS SELECT 1)" + ] config = Config( model_defaults=ModelDefaultsConfig(dialect="bigquery"), before_all=before_all, @@ -3159,15 +3237,21 @@ def assert_cached_violations_exist(cache: OptimizedQueryCache, model: Model): paths=tmp_path, ) - config_err = "Linter detected errors in the code. Please fix them before proceeding." + config_err = ( + "Linter detected errors in the code. Please fix them before proceeding." + ) # Case: Ensure load DOES NOT work if linter is enabled for query in ("SELECT * FROM tbl", "SELECT t.* FROM tbl"): with pytest.raises(LinterError, match=config_err): - ctx.upsert_model(load_sql_based_model(d.parse(f"MODEL (name test); {query}"))) + ctx.upsert_model( + load_sql_based_model(d.parse(f"MODEL (name test); {query}")) + ) ctx.plan(environment="dev", auto_apply=True, no_prompts=True) - error_model = load_sql_based_model(d.parse("MODEL (name test); SELECT * FROM (SELECT 1)")) + error_model = load_sql_based_model( + d.parse("MODEL (name test); SELECT * FROM (SELECT 1)") + ) with pytest.raises(LinterError, match=config_err): ctx.upsert_model(error_model) ctx.plan_builder("dev") @@ -3397,7 +3481,10 @@ def create_log_view(evaluator, view_name): model.on_virtual_update[0].sql(dialect=dialect) == "CREATE OR REPLACE TABLE log_schema AS SELECT @resolve_template('@{schema_name}') AS my_schema" ) - assert model.on_virtual_update[1].sql(dialect=dialect) == "@create_log_view(@this_model)" + assert ( + model.on_virtual_update[1].sql(dialect=dialect) + == "@create_log_view(@this_model)" + ) snapshot = context.get_snapshot("db.test_view_macro_this_model") assert snapshot and snapshot.version @@ -3412,7 +3499,9 @@ def create_log_view(evaluator, view_name): ) # Validate that from the macro evaluator this_model we get the environment-specific fqn - assert log_view["evaluator_this_model"][0] == '"db__dev"."test_view_macro_this_model"' + assert ( + log_view["evaluator_this_model"][0] == '"db__dev"."test_view_macro_this_model"' + ) # Validate the schema is retrieved using resolve_template for the environment-specific schema assert log_schema["my_schema"][0] == "db__dev" @@ -3420,13 +3509,11 @@ def create_log_view(evaluator, view_name): def test_plan_audit_intervals(tmp_path: pathlib.Path, caplog): ctx = Context( - paths=tmp_path, config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + paths=tmp_path, + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), ) - ctx.upsert_model( - load_sql_based_model( - parse( - """ + ctx.upsert_model(load_sql_based_model(parse(""" MODEL ( name sqlmesh_audit.date_example, kind INCREMENTAL_BY_TIME_RANGE( @@ -3442,15 +3529,9 @@ def test_plan_audit_intervals(tmp_path: pathlib.Path, caplog): DATE('2025-02-01') as date_id, ) SELECT date_id FROM sample_table WHERE date_id BETWEEN @start_ds AND @end_ds - """ - ) - ) - ) + """))) - ctx.upsert_model( - load_sql_based_model( - parse( - """ + ctx.upsert_model(load_sql_based_model(parse(""" MODEL ( name sqlmesh_audit.timestamp_example, kind INCREMENTAL_BY_TIME_RANGE( @@ -3466,18 +3547,21 @@ def test_plan_audit_intervals(tmp_path: pathlib.Path, caplog): TIMESTAMP('2025-02-01') as timestamp_id, ) SELECT timestamp_id FROM sample_table WHERE timestamp_id BETWEEN @start_ts AND @end_ts - """ - ) - ) - ) + """))) plan = ctx.plan( - environment="dev", auto_apply=True, no_prompts=True, start="2025-02-01", end="2025-02-01" + environment="dev", + auto_apply=True, + no_prompts=True, + start="2025-02-01", + end="2025-02-01", ) assert plan.missing_intervals date_snapshot = next(s for s in plan.new_snapshots if "date_example" in s.name) - timestamp_snapshot = next(s for s in plan.new_snapshots if "timestamp_example" in s.name) + timestamp_snapshot = next( + s for s in plan.new_snapshots if "timestamp_example" in s.name + ) # Case 1: The timestamp audit should be in the inclusive range ['2025-02-01 00:00:00', '2025-02-01 23:59:59.999999'] assert ( @@ -3497,10 +3581,14 @@ def test_check_intervals(sushi_context, mocker): SQLMeshError, match="Environment 'dev' was not found", ): - sushi_context.check_intervals(environment="dev", no_signals=False, select_models=[]) + sushi_context.check_intervals( + environment="dev", no_signals=False, select_models=[] + ) spy = mocker.spy(sqlmesh.core.snapshot.definition, "check_ready_intervals") - intervals = sushi_context.check_intervals(environment=None, no_signals=False, select_models=[]) + intervals = sushi_context.check_intervals( + environment=None, no_signals=False, select_models=[] + ) min_intervals = 19 assert spy.call_count == 2 @@ -3510,7 +3598,9 @@ def test_check_intervals(sushi_context, mocker): assert not i.intervals spy.reset_mock() - intervals = sushi_context.check_intervals(environment=None, no_signals=True, select_models=[]) + intervals = sushi_context.check_intervals( + environment=None, no_signals=True, select_models=[] + ) assert spy.call_count == 0 assert len(intervals) >= min_intervals @@ -3520,7 +3610,10 @@ def test_check_intervals(sushi_context, mocker): assert len(intervals) == 1 intervals = sushi_context.check_intervals( - environment=None, no_signals=False, select_models=["*waiter_as_customer*"], end="next week" + environment=None, + no_signals=False, + select_models=["*waiter_as_customer*"], + end="next week", ) assert tuple(intervals.values())[0].intervals @@ -3528,8 +3621,7 @@ def test_check_intervals(sushi_context, mocker): def test_audit(): context = Context(config=Config()) - parsed_model = parse( - """ + parsed_model = parse(""" MODEL ( name dummy, audits ( @@ -3538,15 +3630,15 @@ def test_audit(): ); SELECT NULL AS c - """ - ) + """) context.upsert_model(load_sql_based_model(parsed_model)) context.plan(no_prompts=True, auto_apply=True) - assert context.audit(models=["dummy"], start="2020-01-01", end="2020-01-01") is False + assert ( + context.audit(models=["dummy"], start="2020-01-01", end="2020-01-01") is False + ) - parsed_model = parse( - """ + parsed_model = parse(""" MODEL ( name dummy, audits ( @@ -3555,15 +3647,16 @@ def test_audit(): ); SELECT 1 AS c - """ - ) + """) context.upsert_model(load_sql_based_model(parsed_model)) context.plan(no_prompts=True, auto_apply=True) assert context.audit(models=["dummy"], start="2020-01-01", end="2020-01-01") is True -def test_prompt_if_uncategorized_snapshot(mocker: MockerFixture, tmp_path: Path) -> None: +def test_prompt_if_uncategorized_snapshot( + mocker: MockerFixture, tmp_path: Path +) -> None: init_example_project(tmp_path, engine_type="duckdb") config = Config( @@ -3582,10 +3675,14 @@ def test_prompt_if_uncategorized_snapshot(mocker: MockerFixture, tmp_path: Path) incremental_model = context.get_model("sqlmesh_example.incremental_model") incremental_model_query = incremental_model.render_query() - new_incremental_model_query = t.cast(exp.Select, incremental_model_query).select("1 AS z") + new_incremental_model_query = t.cast(exp.Select, incremental_model_query).select( + "1 AS z" + ) context.upsert_model( "sqlmesh_example.incremental_model", - query_=ParsableSql(sql=new_incremental_model_query.sql(dialect=incremental_model.dialect)), + query_=ParsableSql( + sql=new_incremental_model_query.sql(dialect=incremental_model.dialect) + ), ) mock_console = mocker.Mock() @@ -3603,14 +3700,20 @@ def test_prompt_if_uncategorized_snapshot(mocker: MockerFixture, tmp_path: Path) assert context.config.plan.no_prompts == True -def test_plan_explain_skips_tests(sushi_context: Context, mocker: MockerFixture) -> None: +def test_plan_explain_skips_tests( + sushi_context: Context, mocker: MockerFixture +) -> None: sushi_context.console = TerminalConsole() spy = mocker.spy(sushi_context, "_run_plan_tests") - sushi_context.plan(environment="dev", explain=True, no_prompts=True, include_unmodified=True) + sushi_context.plan( + environment="dev", explain=True, no_prompts=True, include_unmodified=True + ) spy.assert_called_once_with(skip_tests=True) -def test_dev_environment_virtual_update_with_environment_statements(tmp_path: Path) -> None: +def test_dev_environment_virtual_update_with_environment_statements( + tmp_path: Path, +) -> None: models_dir = tmp_path / "models" models_dir.mkdir() model_sql = """ @@ -3637,14 +3740,18 @@ def test_dev_environment_virtual_update_with_environment_statements(tmp_path: Pa context.plan("prod", auto_apply=True, no_prompts=True) # Try to create dev environment without changes (should fail) - with pytest.raises(NoChangesPlanError, match="Creating a new environment requires a change"): + with pytest.raises( + NoChangesPlanError, match="Creating a new environment requires a change" + ): context.plan("dev", auto_apply=True, no_prompts=True) # Now create a new context with only new environment statements config_with_statements = Config( model_defaults=ModelDefaultsConfig(dialect="duckdb"), gateways={"duckdb": GatewayConfig(connection=DuckDBConnectionConfig())}, - before_all=["CREATE TABLE IF NOT EXISTS audit_log (id INT, action VARCHAR(100))"], + before_all=[ + "CREATE TABLE IF NOT EXISTS audit_log (id INT, action VARCHAR(100))" + ], after_all=["INSERT INTO audit_log VALUES (1, 'environment_created')"], ) @@ -3657,7 +3764,9 @@ def test_dev_environment_virtual_update_with_environment_statements(tmp_path: Pa assert env.name == "dev" # Verify the environment statements were stored - stored_statements = context_with_statements.state_reader.get_environment_statements("dev") + stored_statements = context_with_statements.state_reader.get_environment_statements( + "dev" + ) assert len(stored_statements) == 1 assert stored_statements[0].before_all == [ "CREATE TABLE IF NOT EXISTS audit_log (id INT, action VARCHAR(100))" @@ -3708,7 +3817,8 @@ def test_plan_min_intervals(tmp_path: Path): init_example_project(tmp_path, engine_type="duckdb", dialect="duckdb") context = Context( - paths=tmp_path, config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + paths=tmp_path, + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), ) current_time = to_datetime("2020-02-01 00:00:01") @@ -3783,12 +3893,17 @@ def test_plan_min_intervals(tmp_path: Path): assert to_datetime(plan.end) == to_datetime("2020-02-01 00:00:01") assert to_datetime(plan.execution_time) == to_datetime("2020-02-01 00:00:01") - def _get_missing_intervals(plan: Plan, name: str) -> t.List[t.Tuple[datetime, datetime]]: + def _get_missing_intervals( + plan: Plan, name: str + ) -> t.List[t.Tuple[datetime, datetime]]: snapshot_id = context.get_snapshot(name, raise_if_missing=True).snapshot_id snapshot_intervals = next( si for si in plan.missing_intervals if si.snapshot_id == snapshot_id ) - return [(to_datetime(s), to_datetime(e)) for s, e in snapshot_intervals.merged_intervals] + return [ + (to_datetime(s), to_datetime(e)) + for s, e in snapshot_intervals.merged_intervals + ] # check initial intervals - should be full time range between start and execution time assert len(plan.missing_intervals) == 4 @@ -3848,7 +3963,9 @@ def _get_missing_intervals(plan: Plan, name: str) -> t.List[t.Tuple[datetime, da # show that the data was created (which shows that when the Plan became an EvaluatablePlan and eventually evaluated, the start date overrides didnt get dropped) assert context.engine_adapter.fetchall( "select start_dt, end_dt from sqlmesh_example__pr_env.daily_model" - ) == [(to_datetime("2020-01-31 00:00:00"), to_datetime("2020-01-31 23:59:59.999999"))] + ) == [ + (to_datetime("2020-01-31 00:00:00"), to_datetime("2020-01-31 23:59:59.999999")) + ] assert context.engine_adapter.fetchall( "select start_dt, end_dt from sqlmesh_example__pr_env.weekly_model" ) == [ @@ -3887,7 +4004,8 @@ def test_plan_min_intervals_adjusted_for_downstream(tmp_path: Path): init_example_project(tmp_path, engine_type="duckdb", dialect="duckdb") context = Context( - paths=tmp_path, config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + paths=tmp_path, + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), ) current_time = to_datetime("2020-02-01 00:00:01") @@ -3977,7 +4095,10 @@ def _get_missing_intervals(name: str) -> t.List[t.Tuple[datetime, datetime]]: snapshot_intervals = next( si for si in plan.missing_intervals if si.snapshot_id == snapshot_id ) - return [(to_datetime(s), to_datetime(e)) for s, e in snapshot_intervals.merged_intervals] + return [ + (to_datetime(s), to_datetime(e)) + for s, e in snapshot_intervals.merged_intervals + ] # We only operate on completed intervals, so given the current_time this is the range of the last completed week _get_missing_intervals("sqlmesh_example.weekly_model") == [ @@ -4011,23 +4132,33 @@ def _get_missing_intervals(name: str) -> t.List[t.Tuple[datetime, datetime]]: assert context.engine_adapter.fetchall( "select min(start_dt), max(end_dt) from sqlmesh_example__pr_env.weekly_model" - ) == [(to_datetime("2020-01-19 00:00:00"), to_datetime("2020-01-25 23:59:59.999999"))] + ) == [ + (to_datetime("2020-01-19 00:00:00"), to_datetime("2020-01-25 23:59:59.999999")) + ] assert context.engine_adapter.fetchall( "select min(start_dt), max(end_dt) from sqlmesh_example__pr_env.daily_model" - ) == [(to_datetime("2020-01-19 00:00:00"), to_datetime("2020-01-31 23:59:59.999999"))] + ) == [ + (to_datetime("2020-01-19 00:00:00"), to_datetime("2020-01-31 23:59:59.999999")) + ] assert context.engine_adapter.fetchall( "select min(start_dt), max(end_dt) from sqlmesh_example__pr_env.hourly_model" - ) == [(to_datetime("2020-01-19 00:00:00"), to_datetime("2020-01-31 23:59:59.999999"))] + ) == [ + (to_datetime("2020-01-19 00:00:00"), to_datetime("2020-01-31 23:59:59.999999")) + ] assert context.engine_adapter.fetchall( "select min(start_dt), max(end_dt) from sqlmesh_example__pr_env.two_hourly_model" - ) == [(to_datetime("2020-01-31 00:00:00"), to_datetime("2020-01-31 23:59:59.999999"))] + ) == [ + (to_datetime("2020-01-31 00:00:00"), to_datetime("2020-01-31 23:59:59.999999")) + ] assert context.engine_adapter.fetchall( "select min(start_dt), max(end_dt) from sqlmesh_example__pr_env.unrelated_monthly_model" - ) == [(to_datetime("2020-01-01 00:00:00"), to_datetime("2020-01-31 23:59:59.999999"))] + ) == [ + (to_datetime("2020-01-01 00:00:00"), to_datetime("2020-01-31 23:59:59.999999")) + ] def test_defaults_pre_post_statements(tmp_path: Path): @@ -4036,8 +4167,7 @@ def test_defaults_pre_post_statements(tmp_path: Path): models_path.mkdir() # Create config with default statements - config_path.write_text( - """ + config_path.write_text(""" model_defaults: dialect: duckdb pre_statements: @@ -4047,21 +4177,18 @@ def test_defaults_pre_post_statements(tmp_path: Path): - ANALYZE @this_model variables: var1: 4 -""" - ) +""") # Create a model model_path = models_path / "test_model.sql" - model_path.write_text( - """ + model_path.write_text(""" MODEL ( name test_model, kind FULL ); SELECT 1 as id, 'test' as status; -""" - ) +""") ctx = Context(paths=[tmp_path]) @@ -4084,16 +4211,14 @@ def test_defaults_pre_post_statements(tmp_path: Path): assert model.render_pre_statements()[1].sql() == 'SET "threads" = 4' # Update config to change pre_statement - config_path.write_text( - """ + config_path.write_text(""" model_defaults: dialect: duckdb pre_statements: - SET memory_limit = '5GB' # Changed value post_statements: - ANALYZE @this_model -""" - ) +""") # Reload context and create new plan ctx = Context(paths=[tmp_path]) @@ -4124,19 +4249,16 @@ def test_model_defaults_statements_with_on_virtual_update(tmp_path: Path): models_path.mkdir() # Create config with on_virtual_update - config_path.write_text( - """ + config_path.write_text(""" model_defaults: dialect: duckdb on_virtual_update: - SELECT 'Model-defailt virtual update' AS message -""" - ) +""") # Create a model with its own on_virtual_update as wel model_path = models_path / "test_model.sql" - model_path.write_text( - """ + model_path.write_text(""" MODEL ( name test_model, kind FULL @@ -4147,8 +4269,7 @@ def test_model_defaults_statements_with_on_virtual_update(tmp_path: Path): ON_VIRTUAL_UPDATE_BEGIN; SELECT 'Model-specific update' AS message; ON_VIRTUAL_UPDATE_END; -""" - ) +""") ctx = Context(paths=[tmp_path]) @@ -4162,8 +4283,13 @@ def test_model_defaults_statements_with_on_virtual_update(tmp_path: Path): assert len(model.on_virtual_update) == 2 # Default statements should come first - assert model.on_virtual_update[0].sql() == "SELECT 'Model-defailt virtual update' AS message" - assert model.on_virtual_update[1].sql() == "SELECT 'Model-specific update' AS message" + assert ( + model.on_virtual_update[0].sql() + == "SELECT 'Model-defailt virtual update' AS message" + ) + assert ( + model.on_virtual_update[1].sql() == "SELECT 'Model-specific update' AS message" + ) def test_uppercase_gateway_external_models(tmp_path): @@ -4222,9 +4348,9 @@ def test_uppercase_gateway_external_models(tmp_path): for model in context_uppercase.models.values() if model.name == "test_db.uppercase_gateway_table" ] - assert len(gateway_specific_models) == 1, ( - f"External model with lowercase gateway name should be found with uppercase gateway. Found {len(gateway_specific_models)} models" - ) + assert ( + len(gateway_specific_models) == 1 + ), f"External model with lowercase gateway name should be found with uppercase gateway. Found {len(gateway_specific_models)} models" # Verify external model without gateway is also found no_gateway_models = [ @@ -4232,13 +4358,15 @@ def test_uppercase_gateway_external_models(tmp_path): for model in context_uppercase.models.values() if model.name == "test_db.no_gateway_table" ] - assert len(no_gateway_models) == 1, ( - f"External model without gateway should be found. Found {len(no_gateway_models)} models" - ) + assert ( + len(no_gateway_models) == 1 + ), f"External model without gateway should be found. Found {len(no_gateway_models)} models" # Check that the column types are properly loaded (not UNKNOWN) external_model = gateway_specific_models[0] - column_types = {name: str(dtype) for name, dtype in external_model.columns_to_types.items()} + column_types = { + name: str(dtype) for name, dtype in external_model.columns_to_types.items() + } assert column_types == { "id": "INT", "name": "TEXT", @@ -4255,9 +4383,9 @@ def test_uppercase_gateway_external_models(tmp_path): if model.name == "test_db.uppercase_gateway_table" ] # This should work but might fail if case sensitivity is not handled correctly - assert len(gateway_specific_models_mixed) == 1, ( - f"External model should be found regardless of gateway parameter case. Found {len(gateway_specific_models_mixed)} models" - ) + assert ( + len(gateway_specific_models_mixed) == 1 + ), f"External model should be found regardless of gateway parameter case. Found {len(gateway_specific_models_mixed)} models" # Test a case that should demonstrate the potential issue: # Create another external model file with uppercase gateway name in the YAML @@ -4291,17 +4419,14 @@ def test_uppercase_gateway_external_models(tmp_path): for model in context_reloaded.models.values() if model.name == "test_db.uppercase_in_yaml" ] - assert len(uppercase_in_yaml_models) == 1, ( - f"External model with uppercase gateway in YAML should be found. Found {len(uppercase_in_yaml_models)} models" - ) + assert ( + len(uppercase_in_yaml_models) == 1 + ), f"External model with uppercase gateway in YAML should be found. Found {len(uppercase_in_yaml_models)} models" def test_plan_no_start_configured(): context = Context(config=Config()) - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name db.xvg, kind INCREMENTAL_BY_TIME_RANGE ( @@ -4314,18 +4439,12 @@ def test_plan_no_start_configured(): ('1', '2020-01-01'), ) data(id, ds) WHERE ds BETWEEN @start_ds AND @end_ds - """ - ) - ) - ) + """))) prod_plan = context.plan(auto_apply=True) assert len(prod_plan.new_snapshots) == 1 - context.upsert_model( - load_sql_based_model( - parse( - """ + context.upsert_model(load_sql_based_model(parse(""" MODEL( name db.xvg, kind INCREMENTAL_BY_TIME_RANGE ( @@ -4339,10 +4458,7 @@ def test_plan_no_start_configured(): ('1', '2020-01-01'), ) data(id, ds) WHERE ds BETWEEN @start_ds AND @end_ds - """ - ) - ) - ) + """))) # This should raise an error because the model has no start configured and the end time is less than the start time which will be calculated from the intervals with pytest.raises( @@ -4363,7 +4479,9 @@ def test_lint_model_projections(tmp_path: Path): ) ) - config_err = "Linter detected errors in the code. Please fix them before proceeding." + config_err = ( + "Linter detected errors in the code. Please fix them before proceeding." + ) with pytest.raises(LinterError, match=config_err): prod_plan = context.plan(no_prompts=True, auto_apply=True) @@ -4398,7 +4516,9 @@ def test_grants_through_plan_apply(sushi_context, mocker): sync_grants_mock.reset_mock() - new_grants = ({"select": ["analyst", "reporter", "manager"], "insert": ["etl_user"]},) + new_grants = ( + {"select": ["analyst", "reporter", "manager"], "insert": ["etl_user"]}, + ) model_updated = model_with_grants.copy( update={ "query": parse_one(model.query.sql() + " LIMIT 1000"), diff --git a/tests/core/test_dialect.py b/tests/core/test_dialect.py index 142b40b31f..73c7024833 100644 --- a/tests/core/test_dialect.py +++ b/tests/core/test_dialect.py @@ -2,28 +2,21 @@ from sqlglot import Dialect, ParseError, exp, parse_one from sqlglot.dialects.dialect import NormalizationStrategy -from sqlmesh.core.dialect import ( - JinjaQuery, - JinjaStatement, - Model, - format_model_expressions, - normalize_model_name, - parse, - select_from_values_for_batch_range, - text_diff, -) import sqlmesh.core.dialect as d -from sqlmesh.core.model import SqlModel, load_sql_based_model from sqlmesh.core.config.connection import DIALECT_TO_TYPE from sqlmesh.core.config.format import FormatConfig +from sqlmesh.core.dialect import (JinjaQuery, JinjaStatement, Model, + format_model_expressions, + normalize_model_name, parse, + select_from_values_for_batch_range, + text_diff) +from sqlmesh.core.model import SqlModel, load_sql_based_model pytestmark = pytest.mark.dialect_isolated def test_format_model_expressions(): - x = format_model_expressions( - parse( - """ + x = format_model_expressions(parse(""" MODEL( name a.b, -- a kind full, -- b @@ -88,12 +81,8 @@ def test_format_model_expressions(): @runtime_stage = 'creating', GRANT SELECT ON foo.bar TO "bla" ) - """ - ) - ) - assert ( - x - == """MODEL ( + """)) + assert x == """MODEL ( name a.b, /* a */ kind FULL, /* b */ references (a, (b, c) AS d), /* c */ @@ -166,7 +155,6 @@ def test_format_model_expressions(): x::INT::INT; @IF(@runtime_stage = 'creating', GRANT SELECT ON foo.bar TO "bla")""" - ) x = format_model_expressions( parse( @@ -175,9 +163,7 @@ def test_format_model_expressions(): JINJA_QUERY_BEGIN; /* comment */ SELECT * FROM x WHERE y = {{ 1 }}; /* comment */ JINJA_END;""" ) ) - assert ( - x - == """MODEL ( + assert x == """MODEL ( name a.b, kind FULL ); @@ -185,20 +171,15 @@ def test_format_model_expressions(): JINJA_QUERY_BEGIN; /* comment */ SELECT * FROM x WHERE y = {{ 1 }}; /* comment */ JINJA_END;""" - ) x = format_model_expressions( - parse( - """ + parse(""" MODEL(name a.b, kind FULL, dialect bigquery); SELECT SAFE_CAST('bla' AS INT64) AS FOO - """ - ), + """), dialect="bigquery", ) - assert ( - x - == """MODEL ( + assert x == """MODEL ( name a.b, kind FULL, dialect bigquery @@ -206,21 +187,16 @@ def test_format_model_expressions(): SELECT SAFE_CAST('bla' AS INT64) AS FOO""" - ) x = format_model_expressions( - parse( - """ + parse(""" MODEL(name a.b, kind FULL, dialect clickhouse); SELECT data.:String AS foo, CAST(1 AS INT) AS bar - """ - ), + """), dialect="clickhouse", ) # JSONCast (e.g. `.:` syntax in ClickHouse) must not be written to `::` - assert ( - x - == """MODEL ( + assert x == """MODEL ( name a.b, kind FULL, dialect clickhouse @@ -229,23 +205,18 @@ def test_format_model_expressions(): SELECT data.:String AS foo, 1::Int32 AS bar""" - ) x = format_model_expressions( - parse( - """ + parse(""" MODEL(name a.b, kind FULL, dialect tsql, allow_partials true); SELECT TRUE AS col, CAST(x AS INT) AS y FROM t - """ - ), + """), dialect="tsql", ) # The MODEL header is SQLMesh DDL and must not be transpiled: a boolean property # such as `allow_partials true` must stay `TRUE`, not become tsql's `(1 = 1)`. # The query body must still transpile to the target dialect. - assert ( - x - == """MODEL ( + assert x == """MODEL ( name a.b, kind FULL, dialect tsql, @@ -256,22 +227,17 @@ def test_format_model_expressions(): 1 AS col, x::INTEGER AS y FROM t""" - ) x = format_model_expressions( - parse( - """ + parse(""" AUDIT(name my_audit, dialect tsql, blocking false); SELECT TRUE AS col, CAST(x AS INT) AS y FROM t WHERE x > 0 - """ - ), + """), dialect="tsql", ) # AUDIT headers are SQLMesh DDL too: a `false` boolean property must stay # `FALSE`, not become tsql's `(1 = 0)`, while the query body still transpiles. - assert ( - x - == """AUDIT ( + assert x == """AUDIT ( name my_audit, dialect tsql, blocking FALSE @@ -283,43 +249,31 @@ def test_format_model_expressions(): FROM t WHERE x > 0""" - ) x = format_model_expressions( - parse( - """ + parse(""" MODEL(name foo); SELECT CAST(1 AS INT) AS bla - """ - ), + """), rewrite_casts=False, ) - assert ( - x - == """MODEL ( + assert x == """MODEL ( name foo ); SELECT CAST(1 AS INT) AS bla""" - ) - x = format_model_expressions( - parse( - """MODEL(name foo); + x = format_model_expressions(parse("""MODEL(name foo); SELECT CAST(1 AS INT) AS bla; on_virtual_update_begin; CREATE OR REPLACE VIEW test_view FROM demo_db.table;GRANT SELECT ON VIEW @this_model TO ROLE owner_name; JINJA_STATEMENT_BEGIN; GRANT SELECT ON VIEW {{this_model}} TO ROLE admin; JINJA_END; GRANT REFERENCES, SELECT ON FUTURE VIEWS IN DATABASE demo_db TO ROLE owner_name; @resolve_parent_name('parent');GRANT SELECT ON VIEW demo_db.table /* sqlglot.meta replace=false */ TO ROLE admin; -ON_VIRTUAL_update_end;""" - ) - ) +ON_VIRTUAL_update_end;""")) - assert ( - x - == """MODEL ( + assert x == """MODEL ( name foo ); @@ -339,7 +293,6 @@ def test_format_model_expressions(): @resolve_parent_name('parent'); GRANT SELECT ON VIEW demo_db.table /* sqlglot.meta replace=false */ TO ROLE admin; ON_VIRTUAL_UPDATE_END;""" - ) def test_format_model_expressions_normalize_functions(): @@ -359,8 +312,7 @@ def test_format_model_expressions_normalize_functions(): ``format_model_expressions``; assertions at the end of this test exercise that path to prevent regression. """ - expressions = parse( - """ + expressions = parse(""" MODEL ( name x, audits ( @@ -370,13 +322,10 @@ def test_format_model_expressions_normalize_functions(): ); SELECT SUM(id), count(id) FROM foo; - """ - ) + """) # Default: audit references preserved lowercase; COUNT/SUM canonicalized uppercase. - assert ( - format_model_expressions(expressions) - == """MODEL ( + assert format_model_expressions(expressions) == """MODEL ( name x, audits ( unique_combination_of_columns(columns := ( @@ -392,12 +341,10 @@ def test_format_model_expressions_normalize_functions(): SUM(id), COUNT(id) FROM foo""" - ) # "upper": audit references uppercased; query functions uppercased. assert ( - format_model_expressions(expressions, normalize_functions="upper") - == """MODEL ( + format_model_expressions(expressions, normalize_functions="upper") == """MODEL ( name x, audits ( UNIQUE_COMBINATION_OF_COLUMNS(columns := ( @@ -417,8 +364,7 @@ def test_format_model_expressions_normalize_functions(): # "lower": audit references preserved lowercase (already lower); query functions lowercased. assert ( - format_model_expressions(expressions, normalize_functions="lower") - == """MODEL ( + format_model_expressions(expressions, normalize_functions="lower") == """MODEL ( name x, audits ( unique_combination_of_columns(columns := ( @@ -439,9 +385,7 @@ def test_format_model_expressions_normalize_functions(): # None: explicit deferral to SQLGlot default → custom/audit names uppercased, # just like "upper". This is distinct from False (preserve) and must be tested # explicitly because None used to be indistinguishable from the missing kwarg. - assert ( - format_model_expressions(expressions, normalize_functions=None) - == """MODEL ( + assert format_model_expressions(expressions, normalize_functions=None) == """MODEL ( name x, audits ( UNIQUE_COMBINATION_OF_COLUMNS(columns := ( @@ -457,12 +401,10 @@ def test_format_model_expressions_normalize_functions(): SUM(id), COUNT(id) FROM foo""" - ) # Single-meta-expression path: normalize_functions must be forwarded. # Without the fix, this path ignored normalize_functions entirely. - single_model = parse( - """ + single_model = parse(""" MODEL ( name x, audits ( @@ -470,8 +412,7 @@ def test_format_model_expressions_normalize_functions(): not_null(columns := (id)) ) ); - """ - ) + """) assert ( format_model_expressions(single_model, normalize_functions="upper") @@ -488,9 +429,7 @@ def test_format_model_expressions_normalize_functions(): )""" ) - assert ( - format_model_expressions(single_model) - == """MODEL ( + assert format_model_expressions(single_model) == """MODEL ( name x, audits ( unique_combination_of_columns(columns := ( @@ -501,12 +440,10 @@ def test_format_model_expressions_normalize_functions(): )) ) )""" - ) # Single-meta path, None: custom audit names are uppercased (explicit SQLGlot default deferral). assert ( - format_model_expressions(single_model, normalize_functions=None) - == """MODEL ( + format_model_expressions(single_model, normalize_functions=None) == """MODEL ( name x, audits ( UNIQUE_COMBINATION_OF_COLUMNS(columns := ( @@ -541,8 +478,7 @@ def test_format_config_normalize_functions_none(): assert "normalize_functions" not in config.generator_options # Confirm the False-default behaviour: custom audit names must be preserved. - expressions = parse( - """ + expressions = parse(""" MODEL ( name x, audits ( @@ -551,8 +487,7 @@ def test_format_config_normalize_functions_none(): ) ); SELECT id FROM foo - """ - ) + """) result = format_model_expressions(expressions, **config.generator_options) assert "unique_combination_of_columns" in result assert "not_null" in result @@ -568,9 +503,7 @@ def test_macro_format(): def test_format_body_macros(): assert ( - format_model_expressions( - parse( - """ + format_model_expressions(parse(""" Model ( name foo , @macro_dialect(), @properties_macro(prop_1 := 'max', prop_2 := 33)); @WITH(TRUE) x AS (SELECT 1) SELECT col::int @@ -579,9 +512,7 @@ def test_format_body_macros(): @ORDER_BY(@include_order_by) @EACH( @columns, item -> @'@iteaoeuatnoehutoenahuoanteuhonateuhaoenthuaoentuhaeotnhaoem'), @'@foo' - """ - ) - ) + """)) == """MODEL ( name foo, @macro_dialect(), @@ -614,8 +545,7 @@ def test_text_diff(): def test_parse(): - expressions = parse( - """ + expressions = parse(""" MODEL ( kind full, dialect "hive", @@ -634,8 +564,7 @@ def test_parse(): {{ side_effect() }}; JINJA_END; - """ - ) + """) assert len(expressions) == 4 assert isinstance(expressions[0], Model) @@ -648,8 +577,7 @@ def test_parse(): assert parse_one("metric") == exp.column("metric") assert parse_one("model(1, 2, 3)") == exp.func("model", 1, 2, 3) - expressions = parse( - """ + expressions = parse(""" MODEL ( kind full, dialect duckdb, @@ -657,8 +585,7 @@ def test_parse(): ); SELECT 1 AS metric - """ - ) + """) assert len(expressions) == 2 assert isinstance(expressions[0], Model) assert isinstance(expressions[1], exp.Select) @@ -679,8 +606,7 @@ def test_parse(): def test_parse_jinja_with_semicolons(): - expressions = parse( - """ + expressions = parse(""" CREATE TABLE a as SELECT 1; CREATE TABLE b as SELECT 1; @@ -694,8 +620,7 @@ def test_parse_jinja_with_semicolons(): DROP TABLE a; DROP TABLE b; - """ - ) + """) assert len(expressions) == 5 assert isinstance(expressions[0], exp.Create) @@ -706,25 +631,20 @@ def test_parse_jinja_with_semicolons(): def test_seed(): - expressions = parse( - """ + expressions = parse(""" MODEL ( kind SEED ( path '..\\..\\..\\data\\data.csv', -- c ), ); - """ - ) + """) assert len(expressions) == 1 assert "../../../data/data.csv" in expressions[0].sql() - assert ( - format_model_expressions(expressions) - == """MODEL ( + assert format_model_expressions(expressions) == """MODEL ( kind SEED ( path '../../../data/data.csv' /* c */ ) )""" - ) def test_select_from_values_for_batch_range_json(): @@ -735,7 +655,9 @@ def test_select_from_values_for_batch_range_json(): "json_col": exp.DataType.build("json"), } - assert select_from_values_for_batch_range(values, columns_to_types, 0, len(values)).sql() == ( + assert select_from_values_for_batch_range( + values, columns_to_types, 0, len(values) + ).sql() == ( """SELECT CAST(id AS INT) AS id, CAST(ds AS TEXT) AS ds, CAST(json_col AS JSON) AS json_col """ """FROM """ """(VALUES (1, '2022-01-01', PARSE_JSON('{"foo":"bar"}')), (2, '2022-01-01', PARSE_JSON('{"foo":"qaz"}'))) """ @@ -755,7 +677,9 @@ def test_select_from_values_that_include_null(): "ts": exp.DataType.build("timestamp", dialect="bigquery"), } - values_expr = select_from_values_for_batch_range(values, columns_to_types, 0, len(values)) + values_expr = select_from_values_for_batch_range( + values, columns_to_types, 0, len(values) + ) assert values_expr.sql(dialect="bigquery") == ( "SELECT CAST(id AS INT64) AS id, CAST(ts AS TIMESTAMP) AS ts FROM " "UNNEST([STRUCT(1 AS id, CAST(NULL AS TIMESTAMP) AS ts)]) AS t" @@ -765,13 +689,25 @@ def test_select_from_values_that_include_null(): @pytest.fixture(params=["mysql", "duckdb", "postgres", "snowflake"]) def normalization_dialect(request): if request.param == "duckdb": - assert Dialect["duckdb"].NORMALIZATION_STRATEGY == NormalizationStrategy.CASE_INSENSITIVE + assert ( + Dialect["duckdb"].NORMALIZATION_STRATEGY + == NormalizationStrategy.CASE_INSENSITIVE + ) elif request.param == "mysql": - assert Dialect["mysql"].NORMALIZATION_STRATEGY == NormalizationStrategy.CASE_SENSITIVE + assert ( + Dialect["mysql"].NORMALIZATION_STRATEGY + == NormalizationStrategy.CASE_SENSITIVE + ) elif request.param == "snowflake": - assert Dialect["snowflake"].NORMALIZATION_STRATEGY == NormalizationStrategy.UPPERCASE + assert ( + Dialect["snowflake"].NORMALIZATION_STRATEGY + == NormalizationStrategy.UPPERCASE + ) elif request.param == "postgres": - assert Dialect["postgres"].NORMALIZATION_STRATEGY == NormalizationStrategy.LOWERCASE + assert ( + Dialect["postgres"].NORMALIZATION_STRATEGY + == NormalizationStrategy.LOWERCASE + ) return request.param @@ -837,7 +773,10 @@ def test_normalize_model_name( uppercase, normalization_dialect, ): - if Dialect[normalization_dialect].NORMALIZATION_STRATEGY == NormalizationStrategy.UPPERCASE: + if ( + Dialect[normalization_dialect].NORMALIZATION_STRATEGY + == NormalizationStrategy.UPPERCASE + ): expected = uppercase elif ( Dialect[normalization_dialect].NORMALIZATION_STRATEGY @@ -851,7 +790,9 @@ def test_normalize_model_name( expected = case_insensitive else: expected = lowercase - assert normalize_model_name(table, default_catalog, normalization_dialect) == expected + assert ( + normalize_model_name(table, default_catalog, normalization_dialect) == expected + ) @pytest.mark.parametrize(normalization_tests_fields, normalization_tests) @@ -864,7 +805,10 @@ def test_multiple_normalization( uppercase, normalization_dialect, ): - if Dialect[normalization_dialect].NORMALIZATION_STRATEGY == NormalizationStrategy.UPPERCASE: + if ( + Dialect[normalization_dialect].NORMALIZATION_STRATEGY + == NormalizationStrategy.UPPERCASE + ): expected = uppercase elif ( Dialect[normalization_dialect].NORMALIZATION_STRATEGY @@ -881,7 +825,8 @@ def test_multiple_normalization( kwargs = {"default_catalog": default_catalog, "dialect": normalization_dialect} assert ( normalize_model_name( - normalize_model_name(normalize_model_name(table, **kwargs), **kwargs), **kwargs + normalize_model_name(normalize_model_name(table, **kwargs), **kwargs), + **kwargs, ) == expected ) @@ -897,7 +842,10 @@ def test_model_normalization_multiple_serde( uppercase, normalization_dialect, ): - if Dialect[normalization_dialect].NORMALIZATION_STRATEGY == NormalizationStrategy.UPPERCASE: + if ( + Dialect[normalization_dialect].NORMALIZATION_STRATEGY + == NormalizationStrategy.UPPERCASE + ): expected = uppercase elif ( Dialect[normalization_dialect].NORMALIZATION_STRATEGY @@ -911,8 +859,7 @@ def test_model_normalization_multiple_serde( expected = case_insensitive else: expected = lowercase - expressions = parse( - f""" + expressions = parse(f""" MODEL ( name {exp.maybe_parse(table, into=exp.Table).sql(dialect=normalization_dialect)}, kind INCREMENTAL_BY_TIME_RANGE( @@ -922,8 +869,7 @@ def test_model_normalization_multiple_serde( ); SELECT col::text, ds::text - """ - ) + """) model = load_sql_based_model( expressions, time_column_format="%Y", default_catalog=default_catalog ) @@ -935,12 +881,16 @@ def test_model_normalization_multiple_serde( def test_model_normalization_quote_flexibility(): assert ( - normalize_model_name("`catalog`.`db`.`table`", default_catalog=None, dialect="spark") + normalize_model_name( + "`catalog`.`db`.`table`", default_catalog=None, dialect="spark" + ) == '"catalog"."db"."table"' ) # This takes advantage of the fact that although double quotes ('"') aren't valid quotes in spark, sqlglot still allows it assert ( - normalize_model_name('"catalog"."db"."table"', default_catalog=None, dialect="spark") + normalize_model_name( + '"catalog"."db"."table"', default_catalog=None, dialect="spark" + ) == '"catalog"."db"."table"' ) @@ -1003,7 +953,10 @@ def test_tsql_alter_column_nullability(): "ALTER TABLE x ALTER COLUMN y INT NOT NULL", "ALTER TABLE x ALTER COLUMN y INTEGER NOT NULL", ), - ("ALTER TABLE x ALTER COLUMN y INT NULL", "ALTER TABLE x ALTER COLUMN y INTEGER NULL"), + ( + "ALTER TABLE x ALTER COLUMN y INT NULL", + "ALTER TABLE x ALTER COLUMN y INTEGER NULL", + ), ("ALTER TABLE x ALTER COLUMN y INT", "ALTER TABLE x ALTER COLUMN y INTEGER"), ]: e = parse_one(sql, read="tsql") @@ -1022,15 +975,16 @@ def test_tsql_alter_column_nullability(): # Statements that don't carry a type are unaffected assert ( - parse_one("ALTER TABLE x ALTER COLUMN y DROP NOT NULL", read="tsql").sql(dialect="tsql") + parse_one("ALTER TABLE x ALTER COLUMN y DROP NOT NULL", read="tsql").sql( + dialect="tsql" + ) == "ALTER TABLE x ALTER COLUMN y DROP NOT NULL" ) def test_model_name_cannot_be_string(): with pytest.raises(ParseError) as parse_error: - parse( - """ + parse(""" MODEL( name 'schema.table', kind FULL @@ -1038,14 +992,15 @@ def test_model_name_cannot_be_string(): SELECT 1 AS c - """ - ) + """) assert "\\'name\\' property cannot be a string value" in str(parse_error) def test_parse_snowflake_create_schema_ddl(): - assert parse_one("CREATE SCHEMA d.s", dialect="snowflake").sql() == "CREATE SCHEMA d.s" + assert ( + parse_one("CREATE SCHEMA d.s", dialect="snowflake").sql() == "CREATE SCHEMA d.s" + ) @pytest.mark.parametrize("dialect", sorted(set(DIALECT_TO_TYPE.values()))) @@ -1066,9 +1021,8 @@ def test_sqlglot_extended_correctly(dialect: str) -> None: def test_format_model_expressions_clustered_by(): # Unquoted AUTO / NONE → formatted without backticks or parens for keyword in ("AUTO", "NONE"): - assert format_model_expressions( - parse( - f""" + assert ( + format_model_expressions(parse(f""" MODEL ( name db.test, kind FULL, @@ -1076,22 +1030,21 @@ def test_format_model_expressions_clustered_by(): clustered_by {keyword} ); SELECT 1 AS a - """ + """)) + == ( + f"MODEL (\n" + f" name db.test,\n" + f" kind FULL,\n" + f" dialect databricks,\n" + f" clustered_by {keyword}\n" + f");\n\nSELECT\n 1 AS a" ) - ) == ( - f"MODEL (\n" - f" name db.test,\n" - f" kind FULL,\n" - f" dialect databricks,\n" - f" clustered_by {keyword}\n" - f");\n\nSELECT\n 1 AS a" ) # Backtick-quoted `auto` / `none` → treated as a column, rendered quoted for name in ("auto", "none"): - assert format_model_expressions( - parse( - f""" + assert ( + format_model_expressions(parse(f""" MODEL ( name db.test, kind FULL, @@ -1099,22 +1052,21 @@ def test_format_model_expressions_clustered_by(): clustered_by `{name}` ); SELECT 1 AS `{name}` - """ + """)) + == ( + f"MODEL (\n" + f" name db.test,\n" + f" kind FULL,\n" + f" dialect databricks,\n" + f' clustered_by "{name}"\n' + f');\n\nSELECT\n 1 AS "{name}"' ) - ) == ( - f"MODEL (\n" - f" name db.test,\n" - f" kind FULL,\n" - f" dialect databricks,\n" - f' clustered_by "{name}"\n' - f');\n\nSELECT\n 1 AS "{name}"' ) # Parens-wrapped (auto) → treated as a column, parens stripped for single column # (same normalisation as partitioned_by (a) → a); quoting happens at model-load time - assert format_model_expressions( - parse( - """ + assert ( + format_model_expressions(parse(""" MODEL ( name db.test, kind FULL, @@ -1122,22 +1074,21 @@ def test_format_model_expressions_clustered_by(): clustered_by (auto) ); SELECT 1 AS auto - """ + """)) + == ( + "MODEL (\n" + " name db.test,\n" + " kind FULL,\n" + " dialect databricks,\n" + " clustered_by auto\n" + ");\n\nSELECT\n 1 AS auto" ) - ) == ( - "MODEL (\n" - " name db.test,\n" - " kind FULL,\n" - " dialect databricks,\n" - " clustered_by auto\n" - ");\n\nSELECT\n 1 AS auto" ) # Multi-column → parens preserved, identifiers as-written # (quoting happens when the model is loaded, not at format time) - assert format_model_expressions( - parse( - """ + assert ( + format_model_expressions(parse(""" MODEL ( name db.test, kind FULL, @@ -1145,15 +1096,15 @@ def test_format_model_expressions_clustered_by(): clustered_by (a, b) ); SELECT 1 AS a, 2 AS b - """ + """)) + == ( + "MODEL (\n" + " name db.test,\n" + " kind FULL,\n" + " dialect databricks,\n" + " clustered_by (a, b)\n" + ");\n\nSELECT\n 1 AS a,\n 2 AS b" ) - ) == ( - "MODEL (\n" - " name db.test,\n" - " kind FULL,\n" - " dialect databricks,\n" - " clustered_by (a, b)\n" - ");\n\nSELECT\n 1 AS a,\n 2 AS b" ) @@ -1161,32 +1112,30 @@ def test_format_model_expressions_clustered_by(): def test_format_model_expressions_clustered_by_non_databricks(keyword: str): """AUTO/NONE without dialect or with a non-Databricks dialect is parsed as a bare identifier.""" # Without dialect — AUTO/NONE parsed as a plain column name (no special keyword handling) - assert format_model_expressions( - parse( - f""" + assert ( + format_model_expressions(parse(f""" MODEL ( name db.test, kind FULL, clustered_by {keyword} ); SELECT 1 AS {keyword.lower()} - """ + """)) + == ( + f"MODEL (\n" + f" name db.test,\n" + f" kind FULL,\n" + f" clustered_by {keyword}\n" + f");\n\nSELECT\n 1 AS {keyword.lower()}" ) - ) == ( - f"MODEL (\n" - f" name db.test,\n" - f" kind FULL,\n" - f" clustered_by {keyword}\n" - f");\n\nSELECT\n 1 AS {keyword.lower()}" ) @pytest.mark.parametrize("keyword", ["AUTO", "NONE"]) def test_format_model_expressions_clustered_by_mixed_list(keyword: str): """AUTO/NONE inside a parenthesised list is treated as a regular column name.""" - assert format_model_expressions( - parse( - f""" + assert ( + format_model_expressions(parse(f""" MODEL ( name db.test, kind FULL, @@ -1194,21 +1143,24 @@ def test_format_model_expressions_clustered_by_mixed_list(keyword: str): clustered_by (a, {keyword}) ); SELECT 1 AS a, 2 AS {keyword.lower()} - """ + """)) + == ( + f"MODEL (\n" + f" name db.test,\n" + f" kind FULL,\n" + f" dialect databricks,\n" + f" clustered_by (a, {keyword})\n" + f");\n\nSELECT\n 1 AS a,\n 2 AS {keyword.lower()}" ) - ) == ( - f"MODEL (\n" - f" name db.test,\n" - f" kind FULL,\n" - f" dialect databricks,\n" - f" clustered_by (a, {keyword})\n" - f");\n\nSELECT\n 1 AS a,\n 2 AS {keyword.lower()}" ) def test_connected_identifier(): ast = d.parse_one("""SELECT ("x"at time zone 'utc')::timestamp as x""", "redshift") - assert ast.sql("redshift") == """SELECT CAST(("x" AT TIME ZONE 'utc') AS TIMESTAMP) AS x""" + assert ( + ast.sql("redshift") + == """SELECT CAST(("x" AT TIME ZONE 'utc') AS TIMESTAMP) AS x""" + ) def test_pipe_syntax(): diff --git a/tests/core/test_environment.py b/tests/core/test_environment.py index 307f220c25..c21cb5df9c 100644 --- a/tests/core/test_environment.py +++ b/tests/core/test_environment.py @@ -66,16 +66,25 @@ def test_lazy_loading(sushi_context): assert all(isinstance(snapshot, SnapshotTableInfo) for snapshot in env.snapshots) assert all(isinstance(s_id, dict) for s_id in env.promoted_snapshot_ids_) assert all(isinstance(s_id, SnapshotId) for s_id in env.promoted_snapshot_ids) - assert all(isinstance(snapshot, dict) for snapshot in env.previous_finalized_snapshots_) assert all( - isinstance(snapshot, SnapshotTableInfo) for snapshot in env.previous_finalized_snapshots + isinstance(snapshot, dict) for snapshot in env.previous_finalized_snapshots_ + ) + assert all( + isinstance(snapshot, SnapshotTableInfo) + for snapshot in env.previous_finalized_snapshots ) - with pytest.raises(ValueError, match="Must be a list of SnapshotTableInfo dicts or objects"): + with pytest.raises( + ValueError, match="Must be a list of SnapshotTableInfo dicts or objects" + ): Environment(**{**env.dict(), **{"snapshots": [1, 2, 3]}}) - with pytest.raises(ValueError, match="Must be a list of SnapshotId dicts or objects"): + with pytest.raises( + ValueError, match="Must be a list of SnapshotId dicts or objects" + ): Environment(**{**env.dict(), **{"promoted_snapshot_ids": [1, 2, 3]}}) - with pytest.raises(ValueError, match="Must be a list of SnapshotTableInfo dicts or objects"): + with pytest.raises( + ValueError, match="Must be a list of SnapshotTableInfo dicts or objects" + ): Environment(**{**env.dict(), **{"previous_finalized_snapshots": [1, 2, 3]}}) diff --git a/tests/core/test_execution_tracker.py b/tests/core/test_execution_tracker.py index 0e58395bee..ef4e8c1598 100644 --- a/tests/core/test_execution_tracker.py +++ b/tests/core/test_execution_tracker.py @@ -2,13 +2,16 @@ from concurrent.futures import ThreadPoolExecutor -from sqlmesh.core.snapshot.execution_tracker import QueryExecutionStats, QueryExecutionTracker -from sqlmesh.core.snapshot import SnapshotIdBatch, SnapshotId +from sqlmesh.core.snapshot import SnapshotId, SnapshotIdBatch +from sqlmesh.core.snapshot.execution_tracker import (QueryExecutionStats, + QueryExecutionTracker) def test_execution_tracker_thread_isolation() -> None: def worker(id: SnapshotId, row_counts: list[int]) -> QueryExecutionStats: - with execution_tracker.track_execution(SnapshotIdBatch(snapshot_id=id, batch_id=0)) as ctx: + with execution_tracker.track_execution( + SnapshotIdBatch(snapshot_id=id, batch_id=0) + ) as ctx: assert execution_tracker.is_tracking() for count in row_counts: @@ -21,8 +24,12 @@ def worker(id: SnapshotId, row_counts: list[int]) -> QueryExecutionStats: with ThreadPoolExecutor() as executor: futures = [ - executor.submit(worker, SnapshotId(name="batch_A", identifier="batch_A"), [10, 5]), - executor.submit(worker, SnapshotId(name="batch_B", identifier="batch_B"), [3, 7]), + executor.submit( + worker, SnapshotId(name="batch_A", identifier="batch_A"), [10, 5] + ), + executor.submit( + worker, SnapshotId(name="batch_B", identifier="batch_B"), [3, 7] + ), ] results = [f.result() for f in futures] diff --git a/tests/core/test_format.py b/tests/core/test_format.py index 5a44e1b381..7044dea8e9 100644 --- a/tests/core/test_format.py +++ b/tests/core/test_format.py @@ -1,14 +1,14 @@ import pathlib +from unittest.mock import call from pytest_mock.plugin import MockerFixture -from sqlmesh.core.config import Config + +from sqlmesh.core.audit import ModelAudit +from sqlmesh.core.config import Config, ModelDefaultsConfig from sqlmesh.core.context import Context from sqlmesh.core.dialect import parse -from sqlmesh.core.audit import ModelAudit from sqlmesh.core.model import SqlModel, load_sql_based_model from tests.utils.test_filesystem import create_temp_file -from unittest.mock import call -from sqlmesh.core.config import ModelDefaultsConfig def test_format_files(tmp_path: pathlib.Path, mocker: MockerFixture): @@ -81,7 +81,9 @@ def test_format_files(tmp_path: pathlib.Path, mocker: MockerFixture): upd1 == "MODEL (\n name this.model,\n dialect 'bigquery'\n);\n\nSELECT\n 1 AS `CaseSensitive`" ) - context.upsert_model(load_sql_based_model(parse(upd1, "bigquery"), default_catalog="memory")) + context.upsert_model( + load_sql_based_model(parse(upd1, "bigquery"), default_catalog="memory") + ) assert context.models['"memory"."this"."model"'].dialect == "bigquery" # Ensure no dialect is added if it's not needed @@ -108,17 +110,26 @@ def test_ignore_formating_files(tmp_path: pathlib.Path): audits_dir = pathlib.Path("audits") # Case 1: Model and Audit are not formatted if the flag is set to false (overriding defaults) - model1_text = "MODEL(name this.model1, dialect 'duckdb', formatting false); SELECT 1 col" - model1 = create_temp_file(tmp_path, pathlib.Path(models_dir, "model_1.sql"), model1_text) + model1_text = ( + "MODEL(name this.model1, dialect 'duckdb', formatting false); SELECT 1 col" + ) + model1 = create_temp_file( + tmp_path, pathlib.Path(models_dir, "model_1.sql"), model1_text + ) audit1_text = "AUDIT(name audit1, dialect 'duckdb', formatting false); SELECT col1 col2 FROM @this_model WHERE foo < 0;" - audit1 = create_temp_file(tmp_path, pathlib.Path(audits_dir, "audit_1.sql"), audit1_text) + audit1 = create_temp_file( + tmp_path, pathlib.Path(audits_dir, "audit_1.sql"), audit1_text + ) audit2_text = "AUDIT(name audit2, dialect 'duckdb', standalone true, formatting false); SELECT col1 col2 FROM @this_model WHERE foo < 0;" - audit2 = create_temp_file(tmp_path, pathlib.Path(audits_dir, "audit_2.sql"), audit2_text) + audit2 = create_temp_file( + tmp_path, pathlib.Path(audits_dir, "audit_2.sql"), audit2_text + ) Context( - paths=tmp_path, config=Config(model_defaults=ModelDefaultsConfig(formatting=True)) + paths=tmp_path, + config=Config(model_defaults=ModelDefaultsConfig(formatting=True)), ).format() assert model1.read_text(encoding="utf-8") == model1_text @@ -127,13 +138,20 @@ def test_ignore_formating_files(tmp_path: pathlib.Path): # Case 2: Model is formatted (or not) based on it's flag and the defaults flag model2_text = "MODEL(name this.model2, dialect 'duckdb'); SELECT 1 col" - model2 = create_temp_file(tmp_path, pathlib.Path(models_dir, "model_2.sql"), model2_text) + model2 = create_temp_file( + tmp_path, pathlib.Path(models_dir, "model_2.sql"), model2_text + ) - model3_text = "MODEL(name this.model3, dialect 'duckdb', formatting true); SELECT 1 col" - model3 = create_temp_file(tmp_path, pathlib.Path(models_dir, "model_3.sql"), model3_text) + model3_text = ( + "MODEL(name this.model3, dialect 'duckdb', formatting true); SELECT 1 col" + ) + model3 = create_temp_file( + tmp_path, pathlib.Path(models_dir, "model_3.sql"), model3_text + ) Context( - paths=tmp_path, config=Config(model_defaults=ModelDefaultsConfig(formatting=False)) + paths=tmp_path, + config=Config(model_defaults=ModelDefaultsConfig(formatting=False)), ).format() # Case 2.1: Model is not formatted if the defaults flag is set to false @@ -158,6 +176,8 @@ def test_format_without_state_load(tmp_path: pathlib.Path, mocker: MockerFixture "MODEL(name local.example, dialect 'duckdb'); SELECT 1 AS col", ) - context = Context(paths=tmp_path, config=Config(project="local_only"), load_state=False) + context = Context( + paths=tmp_path, config=Config(project="local_only"), load_state=False + ) context.format(check=True) mock.assert_not_called() diff --git a/tests/core/test_janitor.py b/tests/core/test_janitor.py index 282336fb06..a255904269 100644 --- a/tests/core/test_janitor.py +++ b/tests/core/test_janitor.py @@ -4,23 +4,17 @@ import pytest from pytest_mock.plugin import MockerFixture -from sqlmesh.core.config import EnvironmentSuffixTarget from sqlmesh.core import constants as c +from sqlmesh.core.config import EnvironmentSuffixTarget from sqlmesh.core.dialect import parse_one, schema_ from sqlmesh.core.engine_adapter import create_engine_adapter from sqlmesh.core.environment import Environment -from sqlmesh.core.model import ( - ModelKindName, - SqlModel, -) +from sqlmesh.core.janitor import (cleanup_expired_views, + delete_expired_snapshots) +from sqlmesh.core.model import ModelKindName, SqlModel from sqlmesh.core.model.definition import ExternalModel -from sqlmesh.core.snapshot import ( - SnapshotChangeCategory, -) -from sqlmesh.core.state_sync import ( - EngineAdapterStateSync, -) -from sqlmesh.core.janitor import cleanup_expired_views, delete_expired_snapshots +from sqlmesh.core.snapshot import SnapshotChangeCategory +from sqlmesh.core.state_sync import EngineAdapterStateSync from sqlmesh.utils.date import now_timestamp pytestmark = pytest.mark.slow @@ -40,13 +34,19 @@ def state_sync(duck_conn, tmp_path): def test_cleanup_expired_views(mocker: MockerFixture, make_snapshot: t.Callable): adapter = mocker.MagicMock() adapter.dialect = None - snapshot_a = make_snapshot(SqlModel(name="catalog.schema.a", query=parse_one("select 1, ds"))) + snapshot_a = make_snapshot( + SqlModel(name="catalog.schema.a", query=parse_one("select 1, ds")) + ) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) - snapshot_b = make_snapshot(SqlModel(name="catalog.schema.b", query=parse_one("select 1, ds"))) + snapshot_b = make_snapshot( + SqlModel(name="catalog.schema.b", query=parse_one("select 1, ds")) + ) snapshot_b.categorize_as(SnapshotChangeCategory.BREAKING) # Make sure that we don't drop schemas from external models snapshot_external_model = make_snapshot( - ExternalModel(name="catalog.external_schema.external_table", kind=ModelKindName.EXTERNAL) + ExternalModel( + name="catalog.external_schema.external_table", kind=ModelKindName.EXTERNAL + ) ) snapshot_external_model.categorize_as(SnapshotChangeCategory.BREAKING) schema_environment = Environment( @@ -63,9 +63,13 @@ def test_cleanup_expired_views(mocker: MockerFixture, make_snapshot: t.Callable) previous_plan_id="test_plan_id", catalog_name_override="catalog_override", ) - snapshot_c = make_snapshot(SqlModel(name="catalog.schema.c", query=parse_one("select 1, ds"))) + snapshot_c = make_snapshot( + SqlModel(name="catalog.schema.c", query=parse_one("select 1, ds")) + ) snapshot_c.categorize_as(SnapshotChangeCategory.BREAKING) - snapshot_d = make_snapshot(SqlModel(name="catalog.schema.d", query=parse_one("select 1, ds"))) + snapshot_d = make_snapshot( + SqlModel(name="catalog.schema.d", query=parse_one("select 1, ds")) + ) snapshot_d.categorize_as(SnapshotChangeCategory.BREAKING) table_environment = Environment( name="test_environment", @@ -101,7 +105,9 @@ def test_cleanup_expired_views(mocker: MockerFixture, make_snapshot: t.Callable) "suffix_target", [EnvironmentSuffixTarget.SCHEMA, EnvironmentSuffixTarget.TABLE] ) def test_cleanup_expired_views_collects_failures( - mocker: MockerFixture, make_snapshot: t.Callable, suffix_target: EnvironmentSuffixTarget + mocker: MockerFixture, + make_snapshot: t.Callable, + suffix_target: EnvironmentSuffixTarget, ): adapter = mocker.MagicMock() adapter.dialect = None @@ -109,7 +115,9 @@ def test_cleanup_expired_views_collects_failures( adapter.drop_view.side_effect = Exception("Failed to drop the view") snapshot = make_snapshot( - SqlModel(name="test_catalog.test_schema.test_model", query=parse_one("select 1, ds")) + SqlModel( + name="test_catalog.test_schema.test_model", query=parse_one("select 1, ds") + ) ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) schema_environment = Environment( @@ -188,9 +196,11 @@ def test_delete_expired_snapshots_common_function_batching( state_sync: EngineAdapterStateSync, make_snapshot: t.Callable, mocker: MockerFixture ): """Test that the common delete_expired_snapshots function properly pages through batches and deletes them.""" - from sqlmesh.core.state_sync.common import ExpiredBatchRange, RowBoundary, LimitBoundary from unittest.mock import MagicMock + from sqlmesh.core.state_sync.common import (ExpiredBatchRange, + LimitBoundary, RowBoundary) + now_ts = now_timestamp() # Create 5 expired snapshots with different timestamps diff --git a/tests/core/test_lineage.py b/tests/core/test_lineage.py index f15abae859..f1c44c52b9 100644 --- a/tests/core/test_lineage.py +++ b/tests/core/test_lineage.py @@ -2,7 +2,8 @@ from sqlmesh.core.config import Config from sqlmesh.core.config.model import ModelDefaultsConfig from sqlmesh.core.context import Context -from sqlmesh.core.lineage import column_dependencies, column_description, lineage +from sqlmesh.core.lineage import (column_dependencies, column_description, + lineage) from sqlmesh.core.model import load_sql_based_model @@ -20,19 +21,19 @@ def test_column_description(sushi_context_pre_scheduling): def test_lineage(): - context = Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="snowflake"))) + context = Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="snowflake")) + ) model = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL (name db.model1); SELECT "A" FROM ( SELECT 1 a ) x - """ - ), + """), ) context.upsert_model(model) diff --git a/tests/core/test_loader.py b/tests/core/test_loader.py index 14a20ec09a..5f8a0bf15a 100644 --- a/tests/core/test_loader.py +++ b/tests/core/test_loader.py @@ -1,5 +1,7 @@ -import pytest from pathlib import Path + +import pytest + from sqlmesh.cli.project_init import init_example_project from sqlmesh.core.config import Config, ModelDefaultsConfig from sqlmesh.core.context import Context @@ -85,7 +87,8 @@ def test_duplicate_model_names_different_kind(tmp_path: Path, sample_models): path_3.write_text(model_3["contents"]) with pytest.raises( - ConfigError, match=r'Duplicate model name\(s\) found: "memory"."test_schema"."test_model".' + ConfigError, + match=r'Duplicate model name\(s\) found: "memory"."test_schema"."test_model".', ): Context(paths=tmp_path, config=config) diff --git a/tests/core/test_macros.py b/tests/core/test_macros.py index 0b3bcf70ee..cfa9032766 100644 --- a/tests/core/test_macros.py +++ b/tests/core/test_macros.py @@ -1,16 +1,17 @@ import typing as t -from datetime import datetime, date +from datetime import date, datetime import pytest from sqlglot import MappingSchema, ParseError, exp, parse_one -from sqlmesh.core import constants as c, dialect as d +from sqlmesh.core import constants as c +from sqlmesh.core import dialect as d from sqlmesh.core.dialect import StagedFilePath -from sqlmesh.core.macros import SQL, MacroEvalError, MacroEvaluator, macro -from sqlmesh.utils.date import to_datetime, to_date +from sqlmesh.core.macros import (SQL, MacroEvalError, MacroEvaluator, + RuntimeStage, macro) +from sqlmesh.utils.date import to_date, to_datetime from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.metaprogramming import Executable -from sqlmesh.core.macros import RuntimeStage @pytest.fixture @@ -19,7 +20,9 @@ def macro_evaluator() -> MacroEvaluator: def filter_country( evaluator: MacroEvaluator, expression: exp.Condition, country: exp.Literal ) -> exp.Condition: - return t.cast(exp.Condition, exp.and_(expression, exp.column("country").eq(country))) + return t.cast( + exp.Condition, exp.and_(expression, exp.column("country").eq(country)) + ) @macro("UPPER") def upper_case(evaluator: MacroEvaluator, expression: exp.Condition) -> str: @@ -34,12 +37,16 @@ def bitshift_square(evaluator: MacroEvaluator, x: int, y: int) -> int: return (x >> y) ** 2 @macro() - def prefix_db(evaluator: MacroEvaluator, table: exp.Table, prefix: str) -> exp.Table: + def prefix_db( + evaluator: MacroEvaluator, table: exp.Table, prefix: str + ) -> exp.Table: table.set("db", prefix + table.db) return table @macro() - def repeated(evaluator: MacroEvaluator, expr: str, times: int = 2, multi: bool = False): + def repeated( + evaluator: MacroEvaluator, expr: str, times: int = 2, multi: bool = False + ): if multi is True: return (expr,) * times return expr * times @@ -73,7 +80,9 @@ def suffix_idents(evaluator: MacroEvaluator, items: t.List[str], suffix: str): return [item + suffix for item in items] @macro() - def suffix_idents_2(evaluator: MacroEvaluator, items: t.Tuple[str, ...], suffix: str): + def suffix_idents_2( + evaluator: MacroEvaluator, items: t.Tuple[str, ...], suffix: str + ): return [item + suffix for item in items] @macro() @@ -97,7 +106,9 @@ def test_select_macro(evaluator): return "SELECT 1 AS col" @macro() - def test_literal_type(evaluator, a: t.Literal["test_literal_a", "test_literal_b", 1, True]): + def test_literal_type( + evaluator, a: t.Literal["test_literal_a", "test_literal_b", 1, True] + ): if isinstance(a, exp.Expr): raise SQLMeshError("Coercion failed") return f"'{a}'" @@ -121,7 +132,9 @@ def test_star(assert_exp_eq) -> None: dialect="tsql", ) evaluator = MacroEvaluator(schema=schema, dialect="tsql") - assert_exp_eq(evaluator.transform(parse_one(sql, read="tsql")), expected_sql, dialect="tsql") + assert_exp_eq( + evaluator.transform(parse_one(sql, read="tsql")), expected_sql, dialect="tsql" + ) sql = "SELECT @STAR(foo, exclude := [SomeColumn]) FROM foo" expected_sql = "SELECT CAST(`foo`.`a` AS STRING) AS `a` FROM foo" @@ -153,12 +166,12 @@ def test_star(assert_exp_eq) -> None: dialect="tsql", ) evaluator = MacroEvaluator(schema=schema, dialect="tsql") - assert_exp_eq(evaluator.transform(parse_one(sql, read="tsql")), expected_sql, dialect="tsql") + assert_exp_eq( + evaluator.transform(parse_one(sql, read="tsql")), expected_sql, dialect="tsql" + ) sql = """SELECT @STAR(foo) FROM foo""" - expected_sql = ( - """SELECT CAST("FOO"."A" AS DATE) AS "A", CAST("FOO"."B" AS INTEGER) AS "B" FROM foo""" - ) + expected_sql = """SELECT CAST("FOO"."A" AS DATE) AS "A", CAST("FOO"."B" AS INTEGER) AS "B" FROM foo""" schema = MappingSchema( { "foo": { @@ -170,13 +183,13 @@ def test_star(assert_exp_eq) -> None: ) evaluator = MacroEvaluator(schema=schema, dialect="snowflake") assert_exp_eq( - evaluator.transform(parse_one(sql, read="snowflake")), expected_sql, dialect="snowflake" + evaluator.transform(parse_one(sql, read="snowflake")), + expected_sql, + dialect="snowflake", ) sql = """SELECT @STAR("foo") FROM "foo" """ - expected_sql = ( - """SELECT CAST("foo"."A" AS DATE) AS "A", CAST("foo"."B" AS INTEGER) AS "B" FROM "foo" """ - ) + expected_sql = """SELECT CAST("foo"."A" AS DATE) AS "A", CAST("foo"."B" AS INTEGER) AS "B" FROM "foo" """ schema = MappingSchema( { '"foo"': { @@ -188,7 +201,9 @@ def test_star(assert_exp_eq) -> None: ) evaluator = MacroEvaluator(schema=schema, dialect="snowflake") assert_exp_eq( - evaluator.transform(parse_one(sql, read="snowflake")), expected_sql, dialect="snowflake" + evaluator.transform(parse_one(sql, read="snowflake")), + expected_sql, + dialect="snowflake", ) sql = """SELECT @STAR(foo, alias := "bar") FROM foo "bar" """ @@ -204,16 +219,22 @@ def test_star(assert_exp_eq) -> None: ) evaluator = MacroEvaluator(schema=schema, dialect="snowflake") assert_exp_eq( - evaluator.transform(parse_one(sql, read="snowflake")), expected_sql, dialect="snowflake" + evaluator.transform(parse_one(sql, read="snowflake")), + expected_sql, + dialect="snowflake", ) def test_start_no_column_types(assert_exp_eq) -> None: sql = """SELECT @STAR(foo) FROM foo""" expected_sql = """SELECT [foo].[a] AS [a] FROM foo""" - schema = MappingSchema({"foo": {"a": exp.DataType.build("UNKNOWN")}}, dialect="tsql") + schema = MappingSchema( + {"foo": {"a": exp.DataType.build("UNKNOWN")}}, dialect="tsql" + ) evaluator = MacroEvaluator(schema=schema, dialect="tsql") - assert_exp_eq(evaluator.transform(parse_one(sql, read="tsql")), expected_sql, dialect="tsql") + assert_exp_eq( + evaluator.transform(parse_one(sql, read="tsql")), expected_sql, dialect="tsql" + ) def test_case(macro_evaluator: MacroEvaluator) -> None: @@ -239,7 +260,10 @@ def test_macro_var(macro_evaluator): macro_evaluator.dialect = "snowflake" assert e.find(StagedFilePath) is not None - assert macro_evaluator.transform(e).sql(dialect="snowflake") == "SELECT a FROM @path, t2" + assert ( + macro_evaluator.transform(e).sql(dialect="snowflake") + == "SELECT a FROM @path, t2" + ) # Referencing a var that doesn't exist in the evaluator's scope should raise macro_evaluator.locals = {} @@ -396,13 +420,21 @@ def test_ast_correctness(macro_evaluator): ), ("SELECT @EACH([1], a -> [@a])", "SELECT ARRAY(1)", {}), ("SELECT @EACH([1, 2], a -> [@a])", "SELECT ARRAY(1), ARRAY(2)", {}), - ("SELECT @REDUCE(@EACH([1], a -> [@a]), (x, y) -> x + y)", "SELECT ARRAY(1)", {}), + ( + "SELECT @REDUCE(@EACH([1], a -> [@a]), (x, y) -> x + y)", + "SELECT ARRAY(1)", + {}, + ), ( "SELECT @REDUCE(@EACH([1, 2], a -> [@a]), (x, y) -> x + y)", "SELECT ARRAY(1) + ARRAY(2)", {}, ), - ("SELECT @REDUCE([[1],[2]], (x, y) -> x + y)", "SELECT ARRAY(1) + ARRAY(2)", {}), + ( + "SELECT @REDUCE([[1],[2]], (x, y) -> x + y)", + "SELECT ARRAY(1) + ARRAY(2)", + {}, + ), ( """@WITH(@do_with) all_cities as (select * from city) select all_cities""", "WITH all_cities AS (SELECT * FROM city) SELECT all_cities", @@ -646,7 +678,9 @@ def test_ast_correctness(macro_evaluator): ), ], ) -def test_macro_functions(macro_evaluator: MacroEvaluator, assert_exp_eq, sql, expected, args): +def test_macro_functions( + macro_evaluator: MacroEvaluator, assert_exp_eq, sql, expected, args +): macro_evaluator.locals = args or {} assert_exp_eq(macro_evaluator.transform(parse_one(sql)), expected) @@ -667,22 +701,28 @@ def test_macro_coercion(macro_evaluator: MacroEvaluator, assert_exp_eq): assert coerce(exp.Literal.number(1.1), float) == 1.1 assert coerce(exp.Literal.string("Hi mom"), str) == "Hi mom" assert coerce(exp.true(), bool) is True - assert coerce(exp.Literal.string("2020-01-01"), datetime) == to_datetime("2020-01-01") + assert coerce(exp.Literal.string("2020-01-01"), datetime) == to_datetime( + "2020-01-01" + ) assert coerce(exp.Literal.string("2020-01-01"), date) == to_date("2020-01-01") # Coercing a string literal to a column should return a column with the same name assert_exp_eq(coerce(exp.Literal.string("order"), exp.Column), exp.column("order")) # Not possible to coerce this string literal Cast to an exp.Column node -- so it should just return the input assert_exp_eq( - coerce(exp.Literal.string("order::date"), exp.Column), exp.Literal.string("order::date") + coerce(exp.Literal.string("order::date"), exp.Column), + exp.Literal.string("order::date"), ) # This however, is correctly coercible since it's a cast assert_exp_eq( - coerce(exp.Literal.string("order::date"), exp.Cast), exp.cast(exp.column("order"), "DATE") + coerce(exp.Literal.string("order::date"), exp.Cast), + exp.cast(exp.column("order"), "DATE"), ) # Here we resolve ambiguity via the user type hint - assert_exp_eq(coerce(exp.Literal.string("order"), exp.Identifier), exp.to_identifier("order")) + assert_exp_eq( + coerce(exp.Literal.string("order"), exp.Identifier), exp.to_identifier("order") + ) assert_exp_eq(coerce(exp.Literal.string("order"), exp.Table), exp.table_("order")) # Resolve a union type hint by choosing the first one that works @@ -699,7 +739,8 @@ def test_macro_coercion(macro_evaluator: MacroEvaluator, assert_exp_eq): # From a string literal to a Select should parse the string literal, and the inverse operation works as well assert_exp_eq( - coerce(exp.Literal.string("SELECT 1 FROM a"), exp.Select), parse_one("SELECT 1 FROM a") + coerce(exp.Literal.string("SELECT 1 FROM a"), exp.Select), + parse_one("SELECT 1 FROM a"), ) assert coerce(parse_one("SELECT 1 FROM a"), SQL) == "SELECT 1 FROM a" @@ -749,10 +790,14 @@ def test_positional_follows_kwargs(macro_evaluator): def test_macro_parameter_resolution(macro_evaluator): - with pytest.raises(MacroEvalError, match=".*missing a required argument: 'pos_only'"): + with pytest.raises( + MacroEvalError, match=".*missing a required argument: 'pos_only'" + ): macro_evaluator.evaluate(parse_one("@test_arg_resolution()")) - with pytest.raises(MacroEvalError, match=".*missing a required argument: 'pos_only'"): + with pytest.raises( + MacroEvalError, match=".*missing a required argument: 'pos_only'" + ): macro_evaluator.evaluate(parse_one("@test_arg_resolution(a1 := 1)")) with pytest.raises(MacroEvalError, match=".*missing a required argument: 'a1'"): @@ -799,9 +844,7 @@ def test_macro_first_value_ignore_respect_nulls(assert_exp_eq) -> None: ) assert_exp_eq(evaluator.transform(actual_expr), expected_sql, dialect="duckdb") - expected_sql = ( - "SELECT FIRST_VALUE(x RESPECT NULLS) OVER (ORDER BY y NULLS FIRST) AS column_test" - ) + expected_sql = "SELECT FIRST_VALUE(x RESPECT NULLS) OVER (ORDER BY y NULLS FIRST) AS column_test" actual_expr = d.parse_one( "SELECT FIRST_VALUE(@test(x) RESPECT NULLS) OVER (ORDER BY y) AS column_test" ) @@ -879,14 +922,18 @@ def test_deduplicate_error_handling(macro_evaluator): SQLMeshError, match="partition_by must be a list of columns: \\[, cast\\( as \\)\\]", ): - macro_evaluator.evaluate(parse_one("@deduplicate(my_table, user_id, ['timestamp DESC'])")) + macro_evaluator.evaluate( + parse_one("@deduplicate(my_table, user_id, ['timestamp DESC'])") + ) # Test error handling: non-list order_by with pytest.raises( SQLMeshError, match="order_by must be a list of strings, optional - nulls ordering: \\[' nulls '\\]", ): - macro_evaluator.evaluate(parse_one("@deduplicate(my_table, [user_id], 'timestamp DESC')")) + macro_evaluator.evaluate( + parse_one("@deduplicate(my_table, [user_id], 'timestamp DESC')") + ) # Test error handling: empty order_by with pytest.raises( @@ -1043,7 +1090,9 @@ def test_date_spine(assert_exp_eq, dialect, date_part): FROM _generated_dates ) AS _generated_dates """ - assert_exp_eq(evaluator.transform(parse_one(date_spine_macro)), expected_sql, dialect=dialect) + assert_exp_eq( + evaluator.transform(parse_one(date_spine_macro)), expected_sql, dialect=dialect + ) def test_date_spine_error_handling(macro_evaluator): @@ -1052,28 +1101,36 @@ def test_date_spine_error_handling(macro_evaluator): MacroEvalError, match=".*Invalid datepart 'invalid'. Expected: 'day', 'week', 'month', 'quarter', or 'year'", ): - macro_evaluator.evaluate(parse_one("@date_spine('invalid', '2022-01-01', '2024-12-31')")) + macro_evaluator.evaluate( + parse_one("@date_spine('invalid', '2022-01-01', '2024-12-31')") + ) # Test error handling: invalid start_date format with pytest.raises( MacroEvalError, match=".*Invalid date format - start_date and end_date must be in format: YYYY-MM-DD", ): - macro_evaluator.evaluate(parse_one("@date_spine('day', '2022/01/01', '2024-12-31')")) + macro_evaluator.evaluate( + parse_one("@date_spine('day', '2022/01/01', '2024-12-31')") + ) # Test error handling: invalid end_date format with pytest.raises( MacroEvalError, match=".*Invalid date format - start_date and end_date must be in format: YYYY-MM-DD", ): - macro_evaluator.evaluate(parse_one("@date_spine('day', '2022-01-01', '2024/12/31')")) + macro_evaluator.evaluate( + parse_one("@date_spine('day', '2022-01-01', '2024/12/31')") + ) # Test error handling: start_date after end_date with pytest.raises( MacroEvalError, match=".*Invalid date range - start_date '2024-12-31' is after end_date '2022-01-01'.", ): - macro_evaluator.evaluate(parse_one("@date_spine('day', '2024-12-31', '2022-01-01')")) + macro_evaluator.evaluate( + parse_one("@date_spine('day', '2024-12-31', '2022-01-01')") + ) def test_macro_union(assert_exp_eq, macro_evaluator: MacroEvaluator): @@ -1103,7 +1160,11 @@ def test_resolve_template_literal(): evaluator.transform(parsed_sql) evaluator.locals.update( - {"this_model": exp.to_table("test_catalog.sqlmesh__test.test__test_model__2517971505")} + { + "this_model": exp.to_table( + "test_catalog.sqlmesh__test.test__test_model__2517971505" + ) + } ) assert ( @@ -1114,7 +1175,11 @@ def test_resolve_template_literal(): # Evaluating evaluator = MacroEvaluator(runtime_stage=RuntimeStage.EVALUATING) evaluator.locals.update( - {"this_model": exp.to_table("test_catalog.sqlmesh__test.test__test_model__2517971505")} + { + "this_model": exp.to_table( + "test_catalog.sqlmesh__test.test__test_model__2517971505" + ) + } ) assert ( evaluator.transform(parsed_sql).sql() @@ -1129,7 +1194,11 @@ def test_resolve_template_table(): evaluator = MacroEvaluator(runtime_stage=RuntimeStage.CREATING) evaluator.locals.update( - {"this_model": exp.to_table("test_catalog.sqlmesh__test.test__test_model__2517971505")} + { + "this_model": exp.to_table( + "test_catalog.sqlmesh__test.test__test_model__2517971505" + ) + } ) assert ( @@ -1149,7 +1218,9 @@ def test_resolve_template_subquery(): evaluator.locals.update( { "this_model": exp.select("*") - .from_(exp.to_table("test_catalog.sqlmesh__test.test__test_model__2517971505")) + .from_( + exp.to_table("test_catalog.sqlmesh__test.test__test_model__2517971505") + ) .where(exp.column("ds").between("2020-01-01", "2020-01-02")) .subquery() } @@ -1165,10 +1236,14 @@ def test_resolve_template_subquery(): evaluator.locals.update( { "this_model": exp.select("*") - .from_(exp.to_table("test_catalog.sqlmesh__test.test__test_model__2517971505")) + .from_( + exp.to_table("test_catalog.sqlmesh__test.test__test_model__2517971505") + ) .where( exp.column("ds").isin( - query=exp.select("ds").from_(exp.to_table("other_catalog.other_schema.other")) + query=exp.select("ds").from_( + exp.to_table("other_catalog.other_schema.other") + ) ) ) .subquery() @@ -1257,7 +1332,9 @@ def test_generate_surrogate_key_hash_semantics() -> None: def render(dialect: str, hash_function: str) -> str: sql = f"SELECT @GENERATE_SURROGATE_KEY(a, hash_function := '{hash_function}') FROM foo" - rendered = MacroEvaluator(dialect=dialect).transform(parse_one(sql, dialect=dialect)) + rendered = MacroEvaluator(dialect=dialect).transform( + parse_one(sql, dialect=dialect) + ) assert isinstance(rendered, exp.Expr) return rendered.sql(dialect) diff --git a/tests/core/test_model.py b/tests/core/test_model.py index 2f8ff49f7d..adb5254ba6 100644 --- a/tests/core/test_model.py +++ b/tests/core/test_model.py @@ -1,78 +1,57 @@ # ruff: noqa: F811 import json -import typing as t import re +import typing as t from datetime import date, datetime from pathlib import Path -from unittest.mock import patch, PropertyMock +from unittest.mock import PropertyMock, patch -import time_machine import pandas as pd # noqa: TID253 import pytest +import time_machine +from pydantic import ValidationError, model_validator from pytest_mock.plugin import MockerFixture from sqlglot import exp, parse_one from sqlglot.errors import ParseError from sqlglot.schema import MappingSchema -from sqlmesh.cli.project_init import init_example_project, ProjectTemplate -from sqlmesh.core.environment import EnvironmentNamingInfo -from sqlmesh.core.model.kind import TimeColumn, ModelKindName, SeedKind -from sqlmesh import CustomMaterialization, CustomKind -from pydantic import model_validator, ValidationError +from sqlmesh import CustomKind, CustomMaterialization +from sqlmesh.cli.project_init import ProjectTemplate, init_example_project from sqlmesh.core import constants as c from sqlmesh.core import dialect as d -from sqlmesh.core.console import get_console from sqlmesh.core.audit import ModelAudit, load_audit -from sqlmesh.core.model.common import ParsableSql -from sqlmesh.core.config import ( - Config, - DuckDBConnectionConfig, - GatewayConfig, - NameInferenceConfig, - ModelDefaultsConfig, - LinterConfig, -) -from sqlmesh.core import constants as c +from sqlmesh.core.config import (Config, DuckDBConnectionConfig, GatewayConfig, + LinterConfig, ModelDefaultsConfig, + NameInferenceConfig) +from sqlmesh.core.console import get_console from sqlmesh.core.context import Context, ExecutionContext from sqlmesh.core.dialect import parse -from sqlmesh.core.engine_adapter.base import MERGE_SOURCE_ALIAS, MERGE_TARGET_ALIAS +from sqlmesh.core.engine_adapter.base import (MERGE_SOURCE_ALIAS, + MERGE_TARGET_ALIAS) from sqlmesh.core.engine_adapter.duckdb import DuckDBEngineAdapter from sqlmesh.core.engine_adapter.shared import DataObjectType -from sqlmesh.core.macros import MacroEvaluator, macro -from sqlmesh.core.model import ( - CustomKind, - PythonModel, - FullKind, - IncrementalByTimeRangeKind, - IncrementalUnmanagedKind, - IncrementalByUniqueKeyKind, - ModelCache, - ModelMeta, - SeedKind, - SqlModel, - TimeColumn, - ExternalKind, - ViewKind, - EmbeddedKind, - SCDType2ByTimeKind, - create_external_model, - create_seed_model, - create_sql_model, - load_sql_based_model, - load_sql_based_models, - model, -) -from sqlmesh.core.model.common import parse_expression -from sqlmesh.core.model.kind import _ModelKind, ModelKindName, _model_kind_validator +from sqlmesh.core.environment import EnvironmentNamingInfo +from sqlmesh.core.macros import MacroEvaluator, RuntimeStage, macro +from sqlmesh.core.model import (CustomKind, EmbeddedKind, ExternalKind, + FullKind, IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, + IncrementalUnmanagedKind, ModelCache, + ModelMeta, PythonModel, SCDType2ByTimeKind, + SeedKind, SqlModel, TimeColumn, ViewKind, + create_external_model, create_seed_model, + create_sql_model, load_sql_based_model, + load_sql_based_models, model) +from sqlmesh.core.model.common import ParsableSql, parse_expression +from sqlmesh.core.model.kind import (ModelKindName, SeedKind, TimeColumn, + _model_kind_validator, _ModelKind) from sqlmesh.core.model.seed import CsvSettings -from sqlmesh.core.node import IntervalUnit, _Node, DbtNodeInfo +from sqlmesh.core.node import DbtNodeInfo, IntervalUnit, _Node from sqlmesh.core.signal import signal from sqlmesh.core.snapshot import Snapshot, SnapshotChangeCategory from sqlmesh.utils.date import TimeLike, to_datetime, to_ds, to_timestamp -from sqlmesh.utils.errors import ConfigError, SQLMeshError, LinterError -from sqlmesh.utils.jinja import JinjaMacroRegistry, MacroInfo, MacroExtractor +from sqlmesh.utils.errors import ConfigError, LinterError, SQLMeshError +from sqlmesh.utils.jinja import JinjaMacroRegistry, MacroExtractor, MacroInfo from sqlmesh.utils.metaprogramming import Executable, SqlValue -from sqlmesh.core.macros import RuntimeStage from tests.utils.test_helpers import use_terminal_console @@ -86,8 +65,7 @@ def missing_schema_warning_msg(model, deps): def test_load(assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -127,8 +105,7 @@ def test_load(assert_exp_eq): t1.a = t2.a; DROP TABLE x; - """ - ) + """) model = load_sql_based_model(expressions) assert model.name == "db.table" @@ -149,7 +126,10 @@ def test_load(assert_exp_eq): } assert model.annotated assert model.view_name == "table" - assert model.macro_definitions == [d.parse_one("@DEF(x, 1)"), d.parse_one("@DEF(y, @x + 1)")] + assert model.macro_definitions == [ + d.parse_one("@DEF(x, 1)"), + d.parse_one("@DEF(y, @x + 1)"), + ] assert list(model.pre_statements) == [ d.parse_one("@DEF(x, 1)"), d.parse_one("@DEF(y, @x + 1)"), @@ -161,9 +141,7 @@ def test_load(assert_exp_eq): ] assert model.depends_on == {'"db"."other_table"'} - assert ( - model.render_query().sql(pretty=True, dialect="spark") - == """SELECT + assert model.render_query().sql(pretty=True, dialect="spark") == """SELECT CAST(1 AS INT) AS `a`, CAST(2 AS DOUBLE) AS `b`, CAST(`c` AS BOOLEAN) AS `c`, @@ -174,7 +152,6 @@ def test_load(assert_exp_eq): FROM `db`.`other_table` AS `t1` LEFT JOIN `db`.`table` AS `t2` ON `t1`.`a` = `t2`.`a`""" - ) assert model.tags == ["tag_foo", "tag_bar"] assert [r.dict() for r in model.all_references] == [ @@ -186,8 +163,7 @@ def test_load(assert_exp_eq): def test_model_multiple_select_statements(): # Make sure the load_model raises an exception for model with multiple select statements. - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -196,15 +172,13 @@ def test_model_multiple_select_statements(): SELECT 1, ds; SELECT 2, ds; - """ - ) + """) with pytest.raises(ConfigError, match=r"^Only one SELECT.*"): load_sql_based_model(expressions) def test_model_validation(tmp_path): - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, kind FULL, @@ -214,11 +188,12 @@ def test_model_validation(tmp_path): y::int, x::int AS y FROM db.ext - """ - ) + """) ctx = Context( - config=Config(linter=LinterConfig(enabled=True, rules=["noambiguousprojections"])), + config=Config( + linter=LinterConfig(enabled=True, rules=["noambiguousprojections"]) + ), paths=tmp_path, ) ctx.upsert_model(load_sql_based_model(expressions, default_catalog="memory")) @@ -227,16 +202,14 @@ def test_model_validation(tmp_path): assert errors, "Expected NoAmbiguousProjections violation" assert errors[0].violation_msg == "Found duplicate outer select name 'y'" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind FULL, ); SELECT a, a UNION SELECT c, c - """ - ) + """) ctx.upsert_model(load_sql_based_model(expressions, default_catalog="memory")) @@ -244,16 +217,14 @@ def test_model_validation(tmp_path): assert errors, "Expected NoAmbiguousProjections violation" assert errors[0].violation_msg == "Found duplicate outer select name 'a'" - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, kind FULL, ); SELECT * FROM db.table - """ - ) + """) model = load_sql_based_model(expressions) with pytest.raises(ConfigError) as ex: @@ -263,30 +234,28 @@ def test_model_validation(tmp_path): def test_model_union_query(sushi_context, assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind FULL, ); SELECT a, b UNION SELECT c, c - """ - ) + """) load_sql_based_model(expressions) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name sushi.test, kind FULL, ); @union('all', sushi.marketing, sushi.marketing) - """ + """) + sushi_context.upsert_model( + load_sql_based_model(expressions, default_catalog="memory") ) - sushi_context.upsert_model(load_sql_based_model(expressions, default_catalog="memory")) assert_exp_eq( sushi_context.get_model("sushi.test").render_query(), """SELECT @@ -407,7 +376,13 @@ def test_model_union_query(sushi_context, assert_exp_eq): ], ) def test_model_union_conditional( - sushi_context, assert_exp_eq, test_id, condition, union_type, table_count, expected_result + sushi_context, + assert_exp_eq, + test_id, + condition, + union_type, + table_count, + expected_result, ): @macro() def get_date(evaluator): @@ -430,17 +405,17 @@ def get_date(evaluator): # Handle the missing union_type case union_type_arg = f", {union_type}" if union_type else "" - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name sushi.{test_id}, kind FULL, ); @union({condition}{union_type_arg}, {tables}) - """ + """) + sushi_context.upsert_model( + load_sql_based_model(expressions, default_catalog="memory") ) - sushi_context.upsert_model(load_sql_based_model(expressions, default_catalog="memory")) assert_exp_eq( sushi_context.get_model(f"sushi.{test_id}").render_query(), @@ -451,19 +426,18 @@ def get_date(evaluator): @use_terminal_console def test_model_qualification(tmp_path: Path): with patch.object(get_console(), "log_warning") as mock_logger: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind FULL, ); SELECT a - """ - ) + """) ctx = Context( - config=Config(linter=LinterConfig(enabled=True, warn_rules=["ALL"])), paths=tmp_path + config=Config(linter=LinterConfig(enabled=True, warn_rules=["ALL"])), + paths=tmp_path, ) ctx.upsert_model(load_sql_based_model(expressions)) ctx.plan_builder("dev") @@ -477,19 +451,19 @@ def test_model_qualification(tmp_path: Path): @use_terminal_console def test_model_missing_audits(tmp_path: Path): with patch.object(get_console(), "log_warning") as mock_logger: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind FULL, ); SELECT a - """ - ) + """) ctx = Context( - config=Config(linter=LinterConfig(enabled=True, warn_rules=["nomissingaudits"])), + config=Config( + linter=LinterConfig(enabled=True, warn_rules=["nomissingaudits"]) + ), paths=tmp_path, ) ctx.upsert_model(load_sql_based_model(expressions)) @@ -550,8 +524,7 @@ def test_project_is_set_in_standalone_audit(tmp_path: Path) -> None: def test_partitioned_by( partition_by_input, partition_by_output, output_dialect, expected_exception ): - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, dialect bigquery, @@ -564,8 +537,7 @@ def test_partitioned_by( ); SELECT 1::int AS a, 2::int AS b, 3 AS c, 4 as d; - """ - ) + """) model = load_sql_based_model(expressions) assert model.clustered_by == [exp.to_column('"c"'), exp.to_column('"d"')] @@ -580,8 +552,7 @@ def test_partitioned_by( def test_opt_out_of_time_column_in_partitioned_by(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect bigquery, @@ -593,23 +564,20 @@ def test_opt_out_of_time_column_in_partitioned_by(): ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) model = load_sql_based_model(expressions) assert model.partitioned_by == [exp.to_column('"b"')] def test_model_no_name(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( dialect bigquery, ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) with pytest.raises(ConfigError) as ex: load_sql_based_model(expressions) @@ -621,16 +589,14 @@ def test_model_no_name(): def test_model_field_name_suggestions(): # top-level field - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialects bigquery, ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) with pytest.raises(ConfigError) as ex: load_sql_based_model(expressions) @@ -640,8 +606,7 @@ def test_model_field_name_suggestions(): ) # kind field - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -651,8 +616,7 @@ def test_model_field_name_suggestions(): ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) with pytest.raises(ConfigError) as ex: load_sql_based_model(expressions) @@ -662,8 +626,7 @@ def test_model_field_name_suggestions(): ) # multiple fields - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialects bigquery, @@ -672,8 +635,7 @@ def test_model_field_name_suggestions(): ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) with pytest.raises(ConfigError) as ex: load_sql_based_model(expressions) @@ -689,16 +651,14 @@ def test_model_field_name_suggestions(): def test_model_required_field_missing(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE (), ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) with pytest.raises(ConfigError) as ex: load_sql_based_model(expressions) @@ -736,8 +696,7 @@ def test_no_model_statement(tmp_path: Path): def test_unordered_model_statements(): - expressions = d.parse( - """ + expressions = d.parse(""" SELECT 1 AS x; MODEL ( @@ -745,8 +704,7 @@ def test_unordered_model_statements(): dialect spark, owner owner_name ); - """ - ) + """) with pytest.raises(ConfigError) as ex: load_sql_based_model(expressions) @@ -754,8 +712,7 @@ def test_unordered_model_statements(): def test_no_query(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -763,14 +720,15 @@ def test_no_query(): ); @DEF(x, 1) - """ - ) + """) with pytest.raises(ConfigError) as ex: model = load_sql_based_model(expressions, path=Path("test_location")) model.validate_definition() - assert "Model query needs to be a SELECT or a UNION, got @DEF(x, 1)." in str(ex.value) + assert "Model query needs to be a SELECT or a UNION, got @DEF(x, 1)." in str( + ex.value + ) def test_single_macro_as_query(assert_exp_eq): @@ -778,15 +736,13 @@ def test_single_macro_as_query(assert_exp_eq): def select_query(evaluator, *projections): return exp.select(*[f'{p} AS "{p}"' for p in projections]) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test ); @SELECT_QUERY(1, 2, 3) - """ - ) + """) model = load_sql_based_model(expressions) assert_exp_eq( model.render_query(), @@ -800,8 +756,7 @@ def select_query(evaluator, *projections): def test_partition_key_is_missing_in_query(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -813,8 +768,7 @@ def test_partition_key_is_missing_in_query(): ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) model = load_sql_based_model(expressions) with pytest.raises(ConfigError) as ex: @@ -823,8 +777,7 @@ def test_partition_key_is_missing_in_query(): def test_cluster_key_is_missing_in_query(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -836,8 +789,7 @@ def test_cluster_key_is_missing_in_query(): ); SELECT 1::int AS a, 2::int AS b; - """ - ) + """) model = load_sql_based_model(expressions) with pytest.raises(ConfigError) as ex: @@ -846,8 +798,7 @@ def test_cluster_key_is_missing_in_query(): def test_partition_key_and_select_star(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -859,15 +810,13 @@ def test_partition_key_and_select_star(): ); SELECT * FROM tbl; - """ - ) + """) load_sql_based_model(expressions) def test_json_serde(): - expressions = parse( - """ + expressions = parse(""" MODEL ( name test_model, kind INCREMENTAL_BY_TIME_RANGE( @@ -889,8 +838,7 @@ def test_json_serde(): @DEF(key, 'value'); SELECT a, ds FROM `tbl` - """ - ) + """) model = load_sql_based_model(expressions) @@ -904,8 +852,7 @@ def test_json_serde(): assert deserialized_model.dict() == model.dict() - expressions = parse( - """ + expressions = parse(""" MODEL ( name test_model, kind FULL, @@ -914,8 +861,7 @@ def test_json_serde(): SELECT x ~ y AS c - """ - ) + """) model = load_sql_based_model(expressions) model_json = model.json() @@ -971,8 +917,7 @@ def test_column_descriptions(sushi_context, assert_exp_eq): "event_date": "Date", } - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind FULL, @@ -983,8 +928,7 @@ def test_column_descriptions(sushi_context, assert_exp_eq): id::int, -- primary key foo::int, -- bar FROM table - """ - ) + """) model = load_sql_based_model(expressions, default_catalog="memory") assert_exp_eq( @@ -1005,8 +949,7 @@ def test_model_jinja_macro_reference_extraction(): def test_macro(**kwargs) -> None: pass - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -1018,8 +961,7 @@ def test_macro(**kwargs) -> None: JINJA_END; SELECT 1 AS x; - """ - ) + """) model = load_sql_based_model(expressions) assert "test_macro" in model.python_env @@ -1031,8 +973,7 @@ def test_model_pre_post_statements(): def foo(**kwargs) -> None: pass - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -1050,8 +991,7 @@ def foo(**kwargs) -> None: @foo(bar='x', val=@this); DROP TABLE x2; - """ - ) + """) model = load_sql_based_model(expressions) expected_pre = [ @@ -1067,17 +1007,18 @@ def foo(**kwargs) -> None: @macro() def multiple_statements(evaluator, t1_value=exp.Literal.number(1)): - return [f"CREATE TABLE t1 AS SELECT {t1_value} AS c", "CREATE TABLE t2 AS SELECT 2 AS c"] + return [ + f"CREATE TABLE t1 AS SELECT {t1_value} AS c", + "CREATE TABLE t2 AS SELECT 2 AS c", + ] - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); SELECT 1 AS col; @multiple_statements() - """ - ) + """) model = load_sql_based_model(expressions) expected_post = d.parse( @@ -1100,8 +1041,7 @@ def foo(evaluator: MacroEvaluator, start: str, end: str) -> str: def bar(evaluator: MacroEvaluator, start: int, end: int) -> str: return f"'{start}, {end}'" - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, kind {model_kind}, @@ -1112,8 +1052,7 @@ def bar(evaluator: MacroEvaluator, start: int, end: int) -> str: SELECT 1 AS x; @bar(@start_millis, @end_millis); - """ - ) + """) model = load_sql_based_model(expressions) start = "2025-01-01" @@ -1128,8 +1067,7 @@ def bar(evaluator: MacroEvaluator, start: int, end: int) -> str: def test_seed_hydration(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1137,10 +1075,11 @@ def test_seed_hydration(): batch_size 100, ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert model.is_hydrated assert not model.derived_columns_to_types @@ -1163,8 +1102,7 @@ def test_seed_hydration(): def test_seed(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1172,10 +1110,11 @@ def test_seed(): batch_size 100, ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert isinstance(model.kind, SeedKind) assert model.kind.path == "../seeds/waiter_names.csv" @@ -1192,23 +1131,20 @@ def test_seed(): def test_seed_model_creation_error(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( path 'gibberish', ) ); - """ - ) + """) with pytest.raises(FileNotFoundError, match="No such file or directory"): load_sql_based_model(expressions) def test_seed_provided_columns(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1220,10 +1156,11 @@ def test_seed_provided_columns(): alias varchar ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert isinstance(model.kind, SeedKind) assert model.kind.path == "../seeds/waiter_names.csv" @@ -1247,8 +1184,7 @@ def test_seed_case_sensitive_columns(tmp_path): """ ) - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.seed, dialect postgres, @@ -1262,10 +1198,11 @@ def test_seed_case_sensitive_columns(tmp_path): "camelCaseTimestamp" timestamp ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert isinstance(model.kind, SeedKind) assert model.seed is not None @@ -1295,8 +1232,7 @@ def test_seed_case_sensitive_columns(tmp_path): def test_seed_csv_settings(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1314,10 +1250,11 @@ def test_seed_csv_settings(): alias varchar ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert isinstance(model.kind, SeedKind) assert model.kind.csv_settings == CsvSettings( @@ -1334,8 +1271,7 @@ def test_seed_csv_settings(): "False", ] - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1345,10 +1281,11 @@ def test_seed_csv_settings(): ), ), ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert isinstance(model.kind, SeedKind) assert model.kind.csv_settings == CsvSettings(na_values=["#N/A", "other"]) @@ -1356,8 +1293,7 @@ def test_seed_csv_settings(): def test_seed_marker_substitution(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1365,8 +1301,7 @@ def test_seed_marker_substitution(): batch_size 100, ) ); - """ - ) + """) model = load_sql_based_model( expressions, @@ -1387,8 +1322,7 @@ def test_seed_pre_post_statements(): def bar(**kwargs) -> None: pass - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1408,10 +1342,11 @@ def bar(**kwargs) -> None: @bar(foo='x', val=@this); DROP TABLE x2; - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) expected_pre = [ *d.parse("@bar()"), @@ -1427,8 +1362,7 @@ def bar(**kwargs) -> None: def test_seed_pre_statements_only(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1442,10 +1376,11 @@ def test_seed_pre_statements_only(): JINJA_END; DROP TABLE x2; - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) expected_pre = [ d.jinja_statement("CREATE TABLE x{{ 1 + 1 }};"), @@ -1456,8 +1391,7 @@ def test_seed_pre_statements_only(): def test_seed_on_virtual_update_statements(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -1477,10 +1411,11 @@ def test_seed_on_virtual_update_statements(): DROP TABLE x2; ON_VIRTUAL_UPDATE_END; - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert model.pre_statements == [d.jinja_statement("CREATE TABLE x{{ 1 + 1 }};")] assert model.on_virtual_update == [ @@ -1493,11 +1428,9 @@ def test_seed_model_custom_types(tmp_path): model_csv_path = (tmp_path / "model.csv").absolute() with open(model_csv_path, "w", encoding="utf-8") as fd: - fd.write( - """key,ds_date,ds_timestamp,b_a,b_b,i,i_str,empty_date + fd.write("""key,ds_date,ds_timestamp,b_a,b_b,i,i_str,empty_date 123,2022-01-01,2022-01-01,false,0,321,321, -""" - ) +""") model = create_seed_model( "test_db.test_model", @@ -1587,8 +1520,7 @@ def test_seed_with_special_characters_in_column(tmp_path, assert_exp_eq): with open(model_csv_path, "w", encoding="utf-8") as fd: fd.write("col.\tcol!@#$\n123\tfoo") - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name memory.test_db.test_model, kind SEED ( @@ -1598,8 +1530,7 @@ def test_seed_with_special_characters_in_column(tmp_path, assert_exp_eq): ) ), ); - """ - ) + """) context.upsert_model(load_sql_based_model(expressions)) assert_exp_eq( @@ -1643,7 +1574,10 @@ def model_with_statements(context, **kwargs): ) python_model = model.get_registry()["db.test_model"].model( - module_path=Path("."), path=Path("."), dialect="duckdb", jinja_macros=jinja_macros + module_path=Path("."), + path=Path("."), + dialect="duckdb", + jinja_macros=jinja_macros, ) assert len(jinja_macros.root_macros) == 2 @@ -1657,7 +1591,9 @@ def model_with_statements(context, **kwargs): ), ] assert python_model.pre_statements == expected_pre - assert python_model.render_pre_statements()[0].sql() == 'CREATE OR REPLACE TABLE "x2"' + assert ( + python_model.render_pre_statements()[0].sql() == 'CREATE OR REPLACE TABLE "x2"' + ) expected_post = [ d.jinja_statement("CREATE INDEX {{test_macro('idx')}} ON db.test_model(id);"), @@ -1672,8 +1608,7 @@ def model_with_statements(context, **kwargs): def test_audits(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, audits ( @@ -1684,12 +1619,12 @@ def test_audits(): tags (foo) ); SELECT 1, ds; - """ - ) + """) audit_definitions = { audit_name: load_audit( - d.parse(f"AUDIT (name {audit_name}); SELECT 1 WHERE FALSE"), dialect="duckdb" + d.parse(f"AUDIT (name {audit_name}); SELECT 1 WHERE FALSE"), + dialect="duckdb", ) for audit_name in ("audit_a", "audit_b", "audit_c") } @@ -1709,15 +1644,13 @@ def test_audits(): def test_custom_audit_arg_changes_affect_fingerprint(): def make_model(min_val: int) -> t.Any: - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.model, audits (check_count(min := {min_val})) ); SELECT 1 AS id; - """ - ) + """) audit_definitions = { "check_count": load_audit( d.parse( @@ -1739,8 +1672,7 @@ def make_model(min_val: int) -> t.Any: def test_enable_audits_from_model_defaults(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.audit_model, ); @@ -1753,10 +1685,11 @@ def test_enable_audits_from_model_defaults(): FROM @this_model WHERE id < 0; - """ - ) + """) - model_defaults = ModelDefaultsConfig(dialect="duckdb", audits=["assert_positive_order_ids"]) + model_defaults = ModelDefaultsConfig( + dialect="duckdb", audits=["assert_positive_order_ids"] + ) model = load_sql_based_model( expressions, @@ -1767,7 +1700,11 @@ def test_enable_audits_from_model_defaults(): assert len(model.audits) == 1 config = Config(model_defaults=model_defaults) - assert config.model_defaults.audits[0] == ("assert_positive_order_ids", {}) == model.audits[0] + assert ( + config.model_defaults.audits[0] + == ("assert_positive_order_ids", {}) + == model.audits[0] + ) audits_with_args = model.audits_with_args assert len(audits_with_args) == 1 @@ -1778,7 +1715,10 @@ def test_enable_audits_from_model_defaults(): def test_description(sushi_context): - assert sushi_context.models['"memory"."sushi"."orders"'].description == "Table of sushi orders." + assert ( + sushi_context.models['"memory"."sushi"."orders"'].description + == "Table of sushi orders." + ) def test_model_defaults_statements_merge(): @@ -1796,8 +1736,7 @@ def test_model_defaults_statements_merge(): ) # Create a model with its own statements as well - expressions = parse( - """ + expressions = parse(""" MODEL ( name test_model, kind FULL @@ -1812,8 +1751,7 @@ def test_model_defaults_statements_merge(): ON_VIRTUAL_UPDATE_BEGIN; UPDATE stats_table SET last_update = CURRENT_TIMESTAMP; ON_VIRTUAL_UPDATE_END; - """ - ) + """) model = load_sql_based_model( expressions, @@ -1824,13 +1762,21 @@ def test_model_defaults_statements_merge(): # Check that pre_statements contains both default and model-specific statements assert len(model.pre_statements) == 3 assert model.pre_statements[0].sql() == "SET enable_progress_bar = TRUE" - assert model.pre_statements[1].sql() == "CREATE TEMPORARY TABLE default_temp AS SELECT 1" - assert model.pre_statements[2].sql() == "CREATE TEMPORARY TABLE model_temp AS SELECT 2" + assert ( + model.pre_statements[1].sql() + == "CREATE TEMPORARY TABLE default_temp AS SELECT 1" + ) + assert ( + model.pre_statements[2].sql() == "CREATE TEMPORARY TABLE model_temp AS SELECT 2" + ) # Check that post_statements contains both default and model-specific statements assert len(model.post_statements) == 3 assert model.post_statements[0].sql() == "DROP TABLE IF EXISTS default_temp" - assert model.post_statements[1].sql() == "GRANT SELECT ON @this_model TO GROUP reporter" + assert ( + model.post_statements[1].sql() + == "GRANT SELECT ON @this_model TO GROUP reporter" + ) assert model.post_statements[2].sql() == "DROP TABLE IF EXISTS model_temp" # Check that the query is rendered correctly with @this_model resolved to table name @@ -1858,16 +1804,14 @@ def test_model_defaults_statements_integration(): ) ) - expressions = parse( - """ + expressions = parse(""" MODEL ( name test_model, kind FULL ); SELECT * FROM source_table; - """ - ) + """) model = load_sql_based_model( expressions, @@ -1883,7 +1827,10 @@ def test_model_defaults_statements_integration(): assert isinstance(model.post_statements[0], exp.Command) assert len(model.on_virtual_update) == 1 - assert model.on_virtual_update[0].sql() == "GRANT SELECT ON @this_model TO GROUP public" + assert ( + model.on_virtual_update[0].sql() + == "GRANT SELECT ON @this_model TO GROUP public" + ) assert ( model.render_on_virtual_update()[0].sql() == 'GRANT SELECT ON "test_model" TO GROUP "public"' @@ -1891,8 +1838,7 @@ def test_model_defaults_statements_integration(): def test_render_definition(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, owner owner_name, @@ -1938,13 +1884,14 @@ def test_render_definition(): t1.a = t2.a; @IF( @runtime_stage = 'creating', create index db_table_idx on db.table(a) ); - """ - ) + """) model = load_sql_based_model( expressions, python_env={ - "test_macro": Executable(payload="def test_macro(evaluator, v):\n return v"), + "test_macro": Executable( + payload="def test_macro(evaluator, v):\n return v" + ), }, default_catalog="catalog", ) @@ -1955,14 +1902,15 @@ def test_render_definition(): ) == d.format_model_expressions(expressions) # Should include the macro implementation. - assert "def test_macro(evaluator, v):" in d.format_model_expressions(model.render_definition()) + assert "def test_macro(evaluator, v):" in d.format_model_expressions( + model.render_definition() + ) def test_tsql_alter_column_post_statement(make_snapshot: t.Callable) -> None: # Issue #5932: the trailing NOT NULL made this parse as a Command, which left @this_model # unresolved and sent the macro to the engine verbatim. - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test.test_model, dialect tsql, @@ -1971,8 +1919,7 @@ def test_tsql_alter_column_post_statement(make_snapshot: t.Callable) -> None: SELECT 1 AS id; @IF(@runtime_stage = 'creating', ALTER TABLE @SQL('@this_model') ALTER COLUMN id INT NOT NULL); - """ - ) + """) model = load_sql_based_model(expressions, default_catalog="catalog") @@ -2010,8 +1957,7 @@ def test_render_definition_with_defaults(): t1.a = t2.a """ - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, owner owner_name, @@ -2020,16 +1966,14 @@ def test_render_definition_with_defaults(): ); {query} - """ - ) + """) model = load_sql_based_model( expressions, default_catalog="catalog", ) - expected_expressions = d.parse( - f""" + expected_expressions = d.parse(f""" MODEL ( name db.table, owner owner_name, @@ -2043,8 +1987,7 @@ def test_render_definition_with_defaults(): ); {query} - """ - ) + """) # Should not include the macro implementation. assert d.format_model_expressions( @@ -2055,8 +1998,7 @@ def test_render_definition_with_defaults(): def test_render_definition_with_grants(): from sqlmesh.core.model.meta import GrantsTargetLayer - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test.grants_model, kind FULL, @@ -2068,8 +2010,7 @@ def test_render_definition_with_grants(): grants_target_layer all, ); SELECT 1 as id - """ - ) + """) model = load_sql_based_model(expressions) assert model.grants_target_layer == GrantsTargetLayer.ALL assert model.grants == { @@ -2096,7 +2037,10 @@ def test_render_definition_with_grants(): grants={"select": ["user1", "user2"], "insert": ["admin"]}, grants_target_layer=GrantsTargetLayer.ALL, ) - assert model_with_grants.grants == {"select": ["user1", "user2"], "insert": ["admin"]} + assert model_with_grants.grants == { + "select": ["user1", "user2"], + "insert": ["admin"], + } assert model_with_grants.grants_target_layer == GrantsTargetLayer.ALL rendered_text = d.format_model_expressions( model_with_grants.render_definition(include_defaults=True) @@ -2110,37 +2054,33 @@ def test_render_definition_with_grants(): rendered_text, ) - virtual_expressions = d.parse( - """ + virtual_expressions = d.parse(""" MODEL ( name test.virtual_grants_model, kind FULL, grants_target_layer virtual ); SELECT 1 as id - """ - ) + """) virtual_model = load_sql_based_model(virtual_expressions) assert virtual_model.grants_target_layer == GrantsTargetLayer.VIRTUAL - default_expressions = d.parse( - """ + default_expressions = d.parse(""" MODEL ( name test.default_grants_model, kind FULL ); SELECT 1 as id - """ - ) + """) default_model = load_sql_based_model(default_expressions) - assert default_model.grants_target_layer == GrantsTargetLayer.VIRTUAL # default value + assert ( + default_model.grants_target_layer == GrantsTargetLayer.VIRTUAL + ) # default value def test_render_definition_partitioned_by(): # no parenthesis in definition, no parenthesis when rendered - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.table, kind FULL, @@ -2148,24 +2088,17 @@ def test_render_definition_partitioned_by(): ); select 1 as a; - """ - ) - ) + """)) assert model.partitioned_by == [exp.column("a", quoted=True)] - assert ( - model.render_definition()[0].sql(pretty=True) - == """MODEL ( + assert model.render_definition()[0].sql(pretty=True) == """MODEL ( name db.table, kind FULL, partitioned_by "a" )""" - ) # single column wrapped in parenthesis in defintion, no parenthesis in rendered - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.table, kind FULL, @@ -2173,24 +2106,17 @@ def test_render_definition_partitioned_by(): ); select 1 as a; - """ - ) - ) + """)) assert model.partitioned_by == [exp.column("a", quoted=True)] - assert ( - model.render_definition()[0].sql(pretty=True) - == """MODEL ( + assert model.render_definition()[0].sql(pretty=True) == """MODEL ( name db.table, kind FULL, partitioned_by "a" )""" - ) # multiple columns wrapped in parenthesis in definition, parenthesis in rendered - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.table, kind FULL, @@ -2198,25 +2124,21 @@ def test_render_definition_partitioned_by(): ); select 1 as a, 2 as b; - """ - ) - ) + """)) - assert model.partitioned_by == [exp.column("a", quoted=True), exp.column("b", quoted=True)] - assert ( - model.render_definition()[0].sql(pretty=True) - == """MODEL ( + assert model.partitioned_by == [ + exp.column("a", quoted=True), + exp.column("b", quoted=True), + ] + assert model.render_definition()[0].sql(pretty=True) == """MODEL ( name db.table, kind FULL, partitioned_by ("a", "b") )""" - ) # multiple columns not wrapped in parenthesis in the definition is an error with pytest.raises(ParseError, match=r"keyword: 'value' missing"): - load_sql_based_model( - d.parse( - f""" + load_sql_based_model(d.parse(f""" MODEL ( name db.table, kind FULL, @@ -2224,14 +2146,11 @@ def test_render_definition_partitioned_by(): ); select 1 as a, 2 as b; - """ - ) - ) + """)) # Iceberg transforms / functions model = load_sql_based_model( - d.parse( - f""" + d.parse(f""" MODEL ( name db.table, kind FULL, @@ -2239,8 +2158,7 @@ def test_render_definition_partitioned_by(): ); select 1 as a, 2 as b, 3 as c; - """ - ), + """), dialect="trino", ) @@ -2253,23 +2171,18 @@ def test_render_definition_partitioned_by(): this=exp.column("c", quoted=True), expression=exp.Literal.number(3) ), ] - assert ( - model.render_definition()[0].sql(pretty=True) - == """MODEL ( + assert model.render_definition()[0].sql(pretty=True) == """MODEL ( name db.table, dialect trino, kind FULL, partitioned_by (DAY("a"), TRUNCATE("b", 4), BUCKET("c", 3)) )""" - ) def test_render_definition_clustered_by(): # Unquoted AUTO keyword → rendered without backticks or parens for keyword in ("AUTO", "NONE"): - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.test, kind FULL, @@ -2277,9 +2190,7 @@ def test_render_definition_clustered_by(): clustered_by {keyword} ); SELECT 1 AS a - """ - ) - ) + """)) assert model.render_definition()[0].sql(pretty=True) == ( f"MODEL (\n" f" name db.test,\n" @@ -2291,9 +2202,7 @@ def test_render_definition_clustered_by(): # Backtick-quoted `auto` / `none` → treated as a real column name, rendered quoted for name in ("auto", "none"): - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.test, kind FULL, @@ -2301,9 +2210,7 @@ def test_render_definition_clustered_by(): clustered_by `{name}` ); SELECT 1 AS `{name}` - """ - ) - ) + """)) assert model.render_definition()[0].sql(pretty=True) == ( f"MODEL (\n" f" name db.test,\n" @@ -2314,9 +2221,7 @@ def test_render_definition_clustered_by(): ) # Parens-wrapped (AUTO) → treated as a real column name, rendered quoted - model = load_sql_based_model( - d.parse( - """ + model = load_sql_based_model(d.parse(""" MODEL ( name db.test, kind FULL, @@ -2324,17 +2229,13 @@ def test_render_definition_clustered_by(): clustered_by (auto) ); SELECT 1 AS auto - """ - ) - ) + """)) assert model.render_definition()[0].sql(pretty=True) == ( 'MODEL (\n name db.test,\n dialect databricks,\n kind FULL,\n clustered_by "auto"\n)' ) # Multi-column → rendered with parens, unchanged - model = load_sql_based_model( - d.parse( - """ + model = load_sql_based_model(d.parse(""" MODEL ( name db.test, kind FULL, @@ -2342,9 +2243,7 @@ def test_render_definition_clustered_by(): clustered_by (a, b) ); SELECT 1 AS a, 2 AS b - """ - ) - ) + """)) assert model.render_definition()[0].sql(pretty=True) == ( "MODEL (\n" " name db.test,\n" @@ -2357,9 +2256,7 @@ def test_render_definition_clustered_by(): def test_render_definition_with_virtual_update_statements(): # model has virtual update statements - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.table, kind FULL @@ -2370,9 +2267,7 @@ def test_render_definition_with_virtual_update_statements(): ON_VIRTUAL_UPDATE_BEGIN; GRANT SELECT ON VIEW @this_model TO ROLE role_name ON_VIRTUAL_UPDATE_END; - """ - ) - ) + """)) assert model.on_virtual_update == [ exp.Grant( @@ -2380,69 +2275,81 @@ def test_render_definition_with_virtual_update_statements(): kind="VIEW", securable=exp.Table(this=d.MacroVar(this="this_model")), principals=[ - exp.GrantPrincipal(this=exp.Identifier(this="role_name", quoted=False), kind="ROLE") + exp.GrantPrincipal( + this=exp.Identifier(this="role_name", quoted=False), kind="ROLE" + ) ], ) ] - assert ( - model.render_definition()[-1].sql(pretty=True) - == """ON_VIRTUAL_UPDATE_BEGIN; + assert model.render_definition()[-1].sql(pretty=True) == """ON_VIRTUAL_UPDATE_BEGIN; GRANT SELECT ON VIEW @this_model TO ROLE role_name; ON_VIRTUAL_UPDATE_END;""" - ) def test_render_definition_dbt_node_info(): node_info = DbtNodeInfo(unique_id="model.db.table", name="table", fqn="db.table") model = load_sql_based_model( - d.parse( - f""" + d.parse(f""" MODEL ( name db.table, kind FULL ); select 1 as a; - """ - ), + """), dbt_node_info=node_info, ) assert model.dbt_node_info - assert ( - model.render_definition()[0].sql(pretty=True) - == """MODEL ( + assert model.render_definition()[0].sql(pretty=True) == """MODEL ( name db.table, dbt_node_info (fqn := 'db.table', name := 'table', unique_id := 'model.db.table'), kind FULL )""" - ) def test_cron(): daily = _Node(name="x", cron="@daily") assert to_datetime(daily.cron_prev("2020-01-01")) == to_datetime("2019-12-31") assert to_datetime(daily.cron_floor("2020-01-01")) == to_datetime("2020-01-01") - assert to_timestamp(daily.cron_floor("2020-01-01 10:00:00")) == to_timestamp("2020-01-01") - assert to_timestamp(daily.cron_next("2020-01-01 10:00:00")) == to_timestamp("2020-01-02") + assert to_timestamp(daily.cron_floor("2020-01-01 10:00:00")) == to_timestamp( + "2020-01-01" + ) + assert to_timestamp(daily.cron_next("2020-01-01 10:00:00")) == to_timestamp( + "2020-01-02" + ) interval = daily.interval_unit assert to_datetime(interval.cron_prev("2020-01-01")) == to_datetime("2019-12-31") assert to_datetime(interval.cron_floor("2020-01-01")) == to_datetime("2020-01-01") - assert to_timestamp(interval.cron_floor("2020-01-01 10:00:00")) == to_timestamp("2020-01-01") - assert to_timestamp(interval.cron_next("2020-01-01 10:00:00")) == to_timestamp("2020-01-02") + assert to_timestamp(interval.cron_floor("2020-01-01 10:00:00")) == to_timestamp( + "2020-01-01" + ) + assert to_timestamp(interval.cron_next("2020-01-01 10:00:00")) == to_timestamp( + "2020-01-02" + ) offset = _Node(name="x", cron="1 0 * * *") - assert to_datetime(offset.cron_prev("2020-01-01")) == to_datetime("2019-12-31 00:01") - assert to_datetime(offset.cron_floor("2020-01-01")) == to_datetime("2019-12-31 00:01") + assert to_datetime(offset.cron_prev("2020-01-01")) == to_datetime( + "2019-12-31 00:01" + ) + assert to_datetime(offset.cron_floor("2020-01-01")) == to_datetime( + "2019-12-31 00:01" + ) assert to_timestamp(offset.cron_floor("2020-01-01 10:00:00")) == to_timestamp( "2020-01-01 00:01" ) - assert to_timestamp(offset.cron_next("2020-01-01 10:00:00")) == to_timestamp("2020-01-02 00:01") + assert to_timestamp(offset.cron_next("2020-01-01 10:00:00")) == to_timestamp( + "2020-01-02 00:01" + ) interval = offset.interval_unit assert to_datetime(interval.cron_prev("2020-01-01")) == to_datetime("2019-12-31") assert to_datetime(interval.cron_floor("2020-01-01")) == to_datetime("2020-01-01") - assert to_timestamp(interval.cron_floor("2020-01-01 10:00:00")) == to_timestamp("2020-01-01") - assert to_timestamp(interval.cron_next("2020-01-01 10:00:00")) == to_timestamp("2020-01-02") + assert to_timestamp(interval.cron_floor("2020-01-01 10:00:00")) == to_timestamp( + "2020-01-01" + ) + assert to_timestamp(interval.cron_next("2020-01-01 10:00:00")) == to_timestamp( + "2020-01-02" + ) hourly = _Node(name="x", cron="1 * * * *") assert to_timestamp(hourly.cron_prev("2020-01-01 10:00:00")) == to_timestamp( @@ -2519,26 +2426,40 @@ def test_cron(): def test_lookback(): model = ModelMeta( - name="x", cron="@hourly", kind=IncrementalByTimeRangeKind(time_column="ts", lookback=2) + name="x", + cron="@hourly", + kind=IncrementalByTimeRangeKind(time_column="ts", lookback=2), ) assert to_timestamp(model.lookback_start("Jan 8 2020 04:00:00")) == to_timestamp( "Jan 8 2020 02:00:00" ) model = ModelMeta( - name="x", cron="@daily", kind=IncrementalByTimeRangeKind(time_column="ds", lookback=2) + name="x", + cron="@daily", + kind=IncrementalByTimeRangeKind(time_column="ds", lookback=2), + ) + assert to_timestamp(model.lookback_start("Jan 8 2020")) == to_timestamp( + "Jan 6 2020" ) - assert to_timestamp(model.lookback_start("Jan 8 2020")) == to_timestamp("Jan 6 2020") model = ModelMeta( - name="x", cron="0 0 1 * *", kind=IncrementalByTimeRangeKind(time_column="ds", lookback=2) + name="x", + cron="0 0 1 * *", + kind=IncrementalByTimeRangeKind(time_column="ds", lookback=2), + ) + assert to_timestamp(model.lookback_start("April 1 2020")) == to_timestamp( + "Feb 1 2020" ) - assert to_timestamp(model.lookback_start("April 1 2020")) == to_timestamp("Feb 1 2020") model = ModelMeta( - name="x", cron="0 0 1 1 *", kind=IncrementalByTimeRangeKind(time_column="ds", lookback=2) + name="x", + cron="0 0 1 1 *", + kind=IncrementalByTimeRangeKind(time_column="ds", lookback=2), + ) + assert to_timestamp(model.lookback_start("Jan 1 2020")) == to_timestamp( + "Jan 1 2018" ) - assert to_timestamp(model.lookback_start("Jan 1 2020")) == to_timestamp("Jan 1 2018") def test_render_query(assert_exp_eq, sushi_context): @@ -2546,15 +2467,13 @@ def test_render_query(assert_exp_eq, sushi_context): name="test", cron="1 0 * * *", kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="y")), - query=d.parse_one( - """ + query=d.parse_one(""" SELECT y FROM x WHERE y BETWEEN @start_date and @end_date AND y BETWEEN @start_ds and @end_ds - """ - ), + """), ) assert_exp_eq( model.render_query(start="2020-10-28", end="2020-10-28"), @@ -2568,7 +2487,9 @@ def test_render_query(assert_exp_eq, sushi_context): """, ) assert_exp_eq( - model.render_query(start="2020-10-28", end="2020-10-28", table_mapping={"x": "x_mapped"}), + model.render_query( + start="2020-10-28", end="2020-10-28", table_mapping={"x": "x_mapped"} + ), """ SELECT "y" AS "y" @@ -2579,8 +2500,7 @@ def test_render_query(assert_exp_eq, sushi_context): """, ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name dummy.model, kind FULL, @@ -2588,8 +2508,7 @@ def test_render_query(assert_exp_eq, sushi_context): ); SELECT COUNT(DISTINCT a) FILTER (WHERE b > 0) AS c FROM x - """ - ) + """) model = load_sql_based_model(expressions, dialect="postgres") assert_exp_eq( model.render_query(), @@ -2608,8 +2527,7 @@ def test_render_query(assert_exp_eq, sushi_context): """, ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name dummy.model, kind FULL, @@ -2619,13 +2537,11 @@ def test_render_query(assert_exp_eq, sushi_context): @DEF(x, ['1', '2', '3']); SELECT @x AS "x" - """ - ) + """) model = load_sql_based_model(expressions, dialect="duckdb") assert model.render_query().sql("duckdb") == '''SELECT ['1', '2', '3'] AS "x"''' - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name dummy.model, kind FULL @@ -2634,16 +2550,14 @@ def test_render_query(assert_exp_eq, sushi_context): @DEF(area, r -> pi() * r * r); SELECT route, centroid, @area(route_radius) AS area - """ - ) + """) model = load_sql_based_model(expressions) assert ( model.render_query().sql() == 'SELECT "route" AS "route", "centroid" AS "centroid", PI() * "route_radius" * "route_radius" AS "area"' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name dummy.model, kind FULL @@ -2653,16 +2567,14 @@ def test_render_query(assert_exp_eq, sushi_context): @DEF(container_volume, (r, h) -> @area(@r) * h); SELECT container_id, @container_volume((cont_di / 2), cont_hi) AS area - """ - ) + """) model = load_sql_based_model(expressions) assert ( model.render_query().sql() == 'SELECT "container_id" AS "container_id", PI() * ("cont_di" / 2) * ("cont_di" / 2) * "cont_hi" AS "area"' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name dummy.model, kind FULL @@ -2671,8 +2583,7 @@ def test_render_query(assert_exp_eq, sushi_context): @DEF(times2, x -> x * 2); SELECT @times4(10) AS "i dont exist" - """ - ) + """) model = load_sql_based_model(expressions) with pytest.raises(SQLMeshError, match=r"Macro 'times4' does not exist.*"): model.render_query() @@ -2705,8 +2616,7 @@ def test_render_query(assert_exp_eq, sushi_context): def test_time_column(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2715,15 +2625,13 @@ def test_time_column(): ); SELECT col::text, ds::text - """ - ) + """) model = load_sql_based_model(expressions) assert model.time_column.column == exp.to_column("ds", quoted=True) assert model.time_column.format == "%Y-%m-%d" assert model.time_column.expression == parse_one("(\"ds\", '%Y-%m-%d')") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2732,15 +2640,13 @@ def test_time_column(): ); SELECT col::text, ds::text - """ - ) + """) model = load_sql_based_model(expressions) assert model.time_column.column == exp.to_column("ds", quoted=True) assert model.time_column.format == "%Y-%m-%d" assert model.time_column.expression == d.parse_one("(\"ds\", '%Y-%m-%d')") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect 'hive', @@ -2750,15 +2656,13 @@ def test_time_column(): ); SELECT col::text, ds::text - """ - ) + """) model = load_sql_based_model(expressions) assert model.time_column.column == exp.to_column("ds", quoted=True) assert model.time_column.format == "%Y-%m" assert model.time_column.expression == d.parse_one("(\"ds\", '%Y-%m')") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2767,15 +2671,13 @@ def test_time_column(): ); SELECT col::text, ds::text - """ - ) + """) with pytest.raises(ConfigError, match="Time Column cannot be empty."): load_sql_based_model(expressions) def test_default_time_column(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2784,13 +2686,11 @@ def test_default_time_column(): ); SELECT col::text, ds::text - """ - ) + """) model = load_sql_based_model(expressions, time_column_format="%Y") assert model.time_column.format == "%Y" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2799,13 +2699,11 @@ def test_default_time_column(): ); SELECT col::text, ds::text - """ - ) + """) model = load_sql_based_model(expressions, time_column_format="%m") assert model.time_column.format == "%Y" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect hive, @@ -2815,15 +2713,13 @@ def test_default_time_column(): ); SELECT col::text, ds::text - """ - ) + """) model = load_sql_based_model(expressions, dialect="duckdb", time_column_format="%Y") assert model.time_column.format == "%d" def test_convert_to_time_column(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2832,14 +2728,14 @@ def test_convert_to_time_column(): ); SELECT ds::text - """ - ) + """) model = load_sql_based_model(expressions) assert model.convert_to_time_column("2022-01-01") == d.parse_one("'2022-01-01'") - assert model.convert_to_time_column(to_datetime("2022-01-01")) == d.parse_one("'2022-01-01'") + assert model.convert_to_time_column(to_datetime("2022-01-01")) == d.parse_one( + "'2022-01-01'" + ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2848,13 +2744,11 @@ def test_convert_to_time_column(): ); SELECT ds::text - """ - ) + """) model = load_sql_based_model(expressions) assert model.convert_to_time_column("2022-01-01") == d.parse_one("'01/01/2022'") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2863,13 +2757,11 @@ def test_convert_to_time_column(): ); SELECT di::int - """ - ) + """) model = load_sql_based_model(expressions) assert model.convert_to_time_column("2022-01-01") == d.parse_one("20220101") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2878,13 +2770,13 @@ def test_convert_to_time_column(): ); SELECT ds::date - """ - ) + """) model = load_sql_based_model(expressions) - assert model.convert_to_time_column("2022-01-01") == d.parse_one("CAST('2022-01-01' AS DATE)") + assert model.convert_to_time_column("2022-01-01") == d.parse_one( + "CAST('2022-01-01' AS DATE)" + ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -2893,8 +2785,7 @@ def test_convert_to_time_column(): ); SELECT ds::timestamp - """ - ) + """) model = load_sql_based_model(expressions) assert model.convert_to_time_column("2022-01-01") == d.parse_one( "CAST('2022-01-01 00:00:00' AS TIMESTAMP)" @@ -2902,8 +2793,7 @@ def test_convert_to_time_column(): def test_parse(assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name sushi.items, kind INCREMENTAL_BY_TIME_RANGE( @@ -2921,8 +2811,7 @@ def test_parse(assert_exp_eq): WHERE ds BETWEEN '{{ start_ds }}' AND @end_ds; JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions, dialect="hive") assert model.columns_to_types == { "ds": exp.DataType.build("unknown"), @@ -2970,116 +2859,82 @@ def assert_match(test_sql: str, expected_value: t.Optional[str] = "duckdb"): assert dialect_str == expected_value # single-quoted dialect - assert_match( - make_test_sql( - """ + assert_match(make_test_sql(""" dialect 'duckdb', description 'there's a dialect foo in here too!' - """ - ) - ) + """)) # bare dialect - assert_match( - make_test_sql( - """ + assert_match(make_test_sql(""" dialect duckdb, description 'there's a dialect foo in here too!' - """ - ) - ) + """)) # double-quoted dialect (allowed in BQ) - assert_match( - make_test_sql( - """ + assert_match(make_test_sql(""" dialect "duckdb", description 'there's a dialect foo in here too!' - """ - ) - ) + """)) # no dialect specified, "dialect" in description - test_sql = make_test_sql( - """ + test_sql = make_test_sql(""" description 'there's a dialect foo in here too!' - """ - ) + """) matches = list(d.DIALECT_PATTERN.finditer(test_sql)) assert not matches # line comment between properties - assert_match( - make_test_sql( - """ + assert_match(make_test_sql(""" tag my_tag, -- comment dialect duckdb - """ - ) - ) + """)) # block comment between properties - assert_match( - make_test_sql( - """ + assert_match(make_test_sql(""" tag my_tag, /* comment */ dialect duckdb - """ - ) - ) + """)) # quoted empty dialect assert_match( - make_test_sql( - """ + make_test_sql(""" dialect '', tag my_tag - """ - ), + """), None, ) # double-quoted empty dialect assert_match( - make_test_sql( - """ + make_test_sql(""" dialect "", tag my_tag - """ - ), + """), None, ) # trailing comment after dialect value - assert_match( - make_test_sql( - """ + assert_match(make_test_sql(""" dialect duckdb -- trailing comment - """ - ) - ) + """)) # dialect value isn't terminated by ',' or ')' - test_sql = make_test_sql( - """ + test_sql = make_test_sql(""" dialect duckdb -- trailing comment tag my_tag - """ - ) + """) matches = list(d.DIALECT_PATTERN.finditer(test_sql)) assert not matches # dialect first - assert_match( - """ + assert_match(""" MODEL( dialect duckdb, name my_name ); - """ - ) + """) # full parse sql = """ @@ -3196,7 +3051,11 @@ def test_python_model_with_properties(make_snapshot): name="python_model_prop", kind="full", columns={"some_col": "int"}, - session_properties={"some_string": "string_prop", "some_bool": True, "some_float": 1.0}, + session_properties={ + "some_string": "string_prop", + "some_bool": True, + "some_float": 1.0, + }, physical_properties={"partition_expiration_days": 7}, virtual_properties={"creatable_type": None}, ) @@ -3243,7 +3102,9 @@ def python_model_prop(context, **kwargs): snapshot.categorize_as(SnapshotChangeCategory.BREAKING) # Rendering the properties will result to a TRANSIENT creatable_type and the removal of the conditional prop - assert m.render_physical_properties(snapshots={m.fqn: snapshot}, python_env=m.python_env) == { + assert m.render_physical_properties( + snapshots={m.fqn: snapshot}, python_env=m.python_env + ) == { "partition_expiration_days": exp.convert(7), "creatable_type": exp.convert("TRANSIENT"), } @@ -3264,7 +3125,9 @@ def test_python_models_returning_sql(assert_exp_eq) -> None: post_statements=["PUT file:///dir/tmp.csv @%table"], ) def model1_entrypoint(evaluator: MacroEvaluator) -> exp.Select: - return exp.select("x", "y").from_(exp.values([("1", 2), ("2", 3)], "_v", ["x", "y"])) + return exp.select("x", "y").from_( + exp.values([("1", 2), ("2", 3)], "_v", ["x", "y"]) + ) @model(name="model2", is_sql=True, kind="full", dialect="snowflake") def model2_entrypoint(evaluator: MacroEvaluator) -> str: @@ -3299,7 +3162,9 @@ def model2_entrypoint(evaluator: MacroEvaluator) -> str: context.render( "model2", expand=[ - d.normalize_model_name("model1", context.default_catalog, context.config.dialect) + d.normalize_model_name( + "model1", context.default_catalog, context.config.dialect + ) ], ), """ @@ -3346,7 +3211,9 @@ def my_model(context): pass # error if kind dict with no `name` key - with pytest.raises(ConfigError, match="`kind` dictionary must contain a `name` key"): + with pytest.raises( + ConfigError, match="`kind` dictionary must contain a `name` key" + ): python_model = model.get_registry()["kind_empty_dict"].model( module_path=Path("."), path=Path("."), @@ -3428,7 +3295,11 @@ def my_model(context): def test_python_model_decorator_col_descriptions() -> None: # `columns` and `column_descriptions` column names are different cases, but name normalization makes both lower - @model("col_descriptions", columns={"col": "int"}, column_descriptions={"COL": "a column"}) + @model( + "col_descriptions", + columns={"col": "int"}, + column_descriptions={"COL": "a column"}, + ) def a_model(context): pass @@ -3475,7 +3346,8 @@ def the_kind(context): pass with pytest.raises( - SQLMeshError, match=r".*Cannot create Python model.*doesn't support Python models" + SQLMeshError, + match=r".*Cannot create Python model.*doesn't support Python models", ): model.get_registry()[f"kind_{kindname}"].model( module_path=Path("."), @@ -3487,8 +3359,7 @@ def test_star_expansion(assert_exp_eq) -> None: context = Context(config=Config()) model1 = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL (name db.model1, kind full); SELECT @@ -3506,30 +3377,25 @@ def test_star_expansion(assert_exp_eq) -> None: (6, 1, '2020-01-06'), (7, 1, '2020-01-07') ) AS t (id, item_id, ds) - """ - ), + """), default_catalog=context.default_catalog, ) model2 = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL (name db.model2, kind full); SELECT * FROM db.model1 AS model1 - """ - ), + """), default_catalog=context.default_catalog, ) model3 = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL(name db.model3, kind full); SELECT * FROM db.model2 AS model2 - """ - ), + """), default_catalog=context.default_catalog, ) @@ -3625,26 +3491,22 @@ def test_case_sensitivity(assert_exp_eq): context = Context(config=config) source = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL (name example.source, kind EMBEDDED); SELECT 'id' AS "id", 'name' AS "name", 'payload' AS "payload" - """ - ), + """), dialect="snowflake", default_catalog=context.default_catalog, ) # Ensure that when manually specifying dependencies, they're normalized correctly downstream = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL (name example.model, kind FULL, depends_on [ExAmPlE.SoUrCe]); SELECT JSON_EXTRACT_PATH_TEXT("payload", 'field') AS "new_field", * FROM example.source - """ - ), + """), dialect="snowflake", default_catalog=context.default_catalog, ) @@ -3677,8 +3539,7 @@ def test_case_sensitivity(assert_exp_eq): def test_batch_size_validation(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -3691,40 +3552,37 @@ def test_batch_size_validation(): ), batch_size 100, ); - """ - ) + """) with pytest.raises(ConfigError): - load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) def test_model_cache(tmp_path: Path, mocker: MockerFixture): cache = ModelCache(tmp_path) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.model_sql, ); SELECT 1, ds; - """ - ) + """) model = load_sql_based_model([e for e in expressions if e]) assert cache.put([model], "test_model", "test_entry_a") assert cache.get("test_model", "test_entry_a")[0].dict() == model.dict() - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.model_seed, kind SEED ( path '../seeds/waiter_names.csv', ), ); - """ - ) + """) seed_model = load_sql_based_model( expressions, path=Path("./examples/sushi/models/test_model.sql") @@ -3775,68 +3633,57 @@ def test_model_cache_default_catalog(tmp_path: Path, mocker: MockerFixture): def test_model_ctas_query(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name `a-b-c.table`, kind FULL, dialect bigquery); SELECT 1 as a FROM x WHERE TRUE LIMIT 2 - """ - ) + """) assert ( load_sql_based_model(expressions, dialect="bigquery").ctas_query().sql() == 'SELECT 1 AS "a" FROM "x" AS "x" WHERE TRUE AND FALSE LIMIT 0' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table, kind FULL); SELECT 1 as a FROM b - """ - ) + """) assert ( load_sql_based_model(expressions).ctas_query().sql() == 'SELECT 1 AS "a" FROM "b" AS "b" WHERE FALSE LIMIT 0' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name `a-b-c.table`, kind FULL, dialect bigquery); SELECT 1 AS a FROM t UNION ALL SELECT 2 AS a FROM t UNION ALL SELECT 3 AS a FROM t UNION ALL SELECT 4 AS a FROM t - """ - ) + """) assert ( load_sql_based_model(expressions, dialect="bigquery").ctas_query().sql() == 'SELECT 1 AS "a" FROM "t" AS "t" WHERE FALSE UNION ALL SELECT 2 AS "a" FROM "t" AS "t" WHERE FALSE UNION ALL SELECT 3 AS "a" FROM "t" AS "t" WHERE FALSE UNION ALL SELECT 4 AS "a" FROM "t" AS "t" WHERE FALSE LIMIT 0' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name `a-b-c.table`, kind FULL, dialect bigquery); SELECT 1 AS a FROM t UNION ALL SELECT 2 AS a FROM t - """ - ) + """) assert ( load_sql_based_model(expressions, dialect="bigquery").ctas_query().sql() == 'SELECT 1 AS "a" FROM "t" AS "t" WHERE FALSE UNION ALL SELECT 2 AS "a" FROM "t" AS "t" WHERE FALSE LIMIT 0' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name `a-b-c.table`, kind FULL, dialect bigquery); SELECT 1 AS a FROM t UNION ALL SELECT 2 AS a FROM t ORDER BY 1 - """ - ) + """) assert ( load_sql_based_model(expressions, dialect="bigquery").ctas_query().sql() == 'SELECT 1 AS "a" FROM "t" AS "t" WHERE FALSE UNION ALL SELECT 2 AS "a" FROM "t" AS "t" WHERE FALSE ORDER BY 1 LIMIT 0' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name `a-b-c.table`, kind FULL, dialect bigquery); WITH RECURSIVE a AS ( SELECT * FROM x @@ -3845,16 +3692,14 @@ def test_model_ctas_query(): ) SELECT * FROM b - """ - ) + """) assert ( load_sql_based_model(expressions, dialect="bigquery").ctas_query().sql() == 'WITH RECURSIVE "a" AS (SELECT * FROM "x" AS "x" WHERE FALSE), "b" AS (SELECT * FROM "a" AS "a" WHERE FALSE UNION ALL SELECT * FROM "a" AS "a" WHERE FALSE) SELECT * FROM "b" AS "b" WHERE FALSE LIMIT 0' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name `a-b-c.table`, kind FULL, dialect bigquery); WITH RECURSIVE a AS ( SELECT * FROM (SELECT * FROM (SELECT * FROM x)) @@ -3863,16 +3708,14 @@ def test_model_ctas_query(): ) SELECT * FROM b - """ - ) + """) assert ( load_sql_based_model(expressions, dialect="bigquery").ctas_query().sql() == 'WITH RECURSIVE "a" AS (SELECT * FROM (SELECT * FROM (SELECT * FROM "x" AS "x" WHERE FALSE) AS "_0" WHERE FALSE) AS "_1" WHERE FALSE), "b" AS (SELECT * FROM "a" AS "a" WHERE FALSE UNION ALL SELECT * FROM "a" AS "a" WHERE FALSE) SELECT * FROM "b" AS "b" WHERE FALSE LIMIT 0' ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name `a-b-c.table`, kind FULL, dialect bigquery); WITH RECURSIVE a AS ( WITH nested_a AS ( @@ -3884,8 +3727,7 @@ def test_model_ctas_query(): ) SELECT * FROM b - """ - ) + """) assert ( load_sql_based_model(expressions, dialect="bigquery").ctas_query().sql() @@ -3895,10 +3737,18 @@ def test_model_ctas_query(): def test_is_breaking_change(): model = create_external_model("a", columns={"a": "int", "limit": "int"}) - assert model.is_breaking_change(create_external_model("a", columns={"a": "int"})) is False - assert model.is_breaking_change(create_external_model("a", columns={"a": "text"})) is None assert ( - model.is_breaking_change(create_external_model("a", columns={"a": "int", "limit": "int"})) + model.is_breaking_change(create_external_model("a", columns={"a": "int"})) + is False + ) + assert ( + model.is_breaking_change(create_external_model("a", columns={"a": "text"})) + is None + ) + assert ( + model.is_breaking_change( + create_external_model("a", columns={"a": "int", "limit": "int"}) + ) is False ) assert ( @@ -3924,15 +3774,13 @@ def runtime_macro(evaluator, **kwargs) -> None: raise ParsetimeAdapterCallError("") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); JINJA_QUERY_BEGIN; SELECT {{ runtime_macro() }} as a FROM b JINJA_QUERY_END; - """ - ) + """) model = load_sql_based_model(expressions) with pytest.raises( @@ -3944,13 +3792,11 @@ def runtime_macro(evaluator, **kwargs) -> None: @use_terminal_console def test_update_schema(tmp_path: Path): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); SELECT * FROM table_a JOIN table_b - """ - ) + """) model = load_sql_based_model(expressions) schema = MappingSchema(normalize=False) @@ -3961,7 +3807,8 @@ def test_update_schema(tmp_path: Path): assert model.mapping_schema == {'"table_a"': {"a": "INT"}} ctx = Context( - config=Config(linter=LinterConfig(enabled=True, warn_rules=["ALL"])), paths=tmp_path + config=Config(linter=LinterConfig(enabled=True, warn_rules=["ALL"])), + paths=tmp_path, ) with patch.object(get_console(), "log_warning") as mock_logger: ctx.upsert_model(model) @@ -3999,36 +3846,51 @@ def test_missing_schema_warnings(tmp_path: Path): console = get_console() ctx = Context( - config=Config(linter=LinterConfig(enabled=True, warn_rules=["ALL"])), paths=tmp_path + config=Config(linter=LinterConfig(enabled=True, warn_rules=["ALL"])), + paths=tmp_path, ) # star, no schema, no deps with patch.object(console, "log_warning") as mock_logger: - model = load_sql_based_model(d.parse("MODEL (name test); SELECT * FROM (SELECT 1 a) x")) + model = load_sql_based_model( + d.parse("MODEL (name test); SELECT * FROM (SELECT 1 a) x") + ) model.render_query(needs_optimization=True) mock_logger.assert_not_called() # star, full schema with patch.object(console, "log_warning") as mock_logger: - model = load_sql_based_model(d.parse("MODEL (name test); SELECT * FROM a CROSS JOIN b")) + model = load_sql_based_model( + d.parse("MODEL (name test); SELECT * FROM a CROSS JOIN b") + ) model.update_schema(full_schema) model.render_query(needs_optimization=True) mock_logger.assert_not_called() # star, partial schema with patch.object(console, "log_warning") as mock_logger: - model = load_sql_based_model(d.parse("MODEL (name test); SELECT * FROM a CROSS JOIN b")) + model = load_sql_based_model( + d.parse("MODEL (name test); SELECT * FROM a CROSS JOIN b") + ) model.update_schema(partial_schema) ctx.upsert_model(model) ctx.plan_builder("dev") - assert missing_schema_warning_msg('"test"', ('"b"',)) in mock_logger.call_args[0][0] + assert ( + missing_schema_warning_msg('"test"', ('"b"',)) + in mock_logger.call_args[0][0] + ) # star, no schema with patch.object(console, "log_warning") as mock_logger: - model = load_sql_based_model(d.parse("MODEL (name test); SELECT * FROM b JOIN a")) + model = load_sql_based_model( + d.parse("MODEL (name test); SELECT * FROM b JOIN a") + ) ctx.upsert_model(model) ctx.plan_builder("dev") - assert missing_schema_warning_msg('"test"', ('"a"', '"b"')) in mock_logger.call_args[0][0] + assert ( + missing_schema_warning_msg('"test"', ('"a"', '"b"')) + in mock_logger.call_args[0][0] + ) # no star, full schema with patch.object(console, "log_warning") as mock_logger: @@ -4059,26 +3921,25 @@ def test_missing_schema_warnings(tmp_path: Path): def test_user_provided_depends_on(): for l_delim, r_delim in (("(", ")"), ("[", "]")): - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL (name db.table, depends_on {l_delim}table_b{r_delim}); SELECT a FROM table_a - """ - ) + """) model = load_sql_based_model(expressions) - assert model.depends_on == {'"table_a"', '"table_b"'}, f"Delimiters {l_delim}, {r_delim}" + assert model.depends_on == { + '"table_a"', + '"table_b"', + }, f"Delimiters {l_delim}, {r_delim}" def test_check_schema_mapping_when_rendering_at_runtime(assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table, depends_on [table_b]); SELECT * FROM table_a JOIN table_b - """ - ) + """) model = load_sql_based_model(expressions) @@ -4097,13 +3958,13 @@ def test_check_schema_mapping_when_rendering_at_runtime(assert_exp_eq): # Simulate rendering at runtime. assert_exp_eq( - model.render_query(), """SELECT * FROM "table_a" AS "table_a", "table_b" AS "table_b" """ + model.render_query(), + """SELECT * FROM "table_a" AS "table_a", "table_b" AS "table_b" """, ) def test_model_normalization(): - expr = d.parse( - """ + expr = d.parse(""" MODEL ( name `project-1.db.tbl`, kind FULL, @@ -4113,10 +3974,11 @@ def test_model_normalization(): grain [id, ds] ); SELECT * FROM project-1.db.raw - """ - ) + """) - model = SqlModel.parse_raw(load_sql_based_model(expr, depends_on={"project-2.db.raw"}).json()) + model = SqlModel.parse_raw( + load_sql_based_model(expr, depends_on={"project-2.db.raw"}).json() + ) assert model.name == "`project-1.db.tbl`" assert model.columns_to_types["a"].sql(dialect="bigquery") == "STRUCT<`a` INT64>" assert model.partitioned_by[0].sql(dialect="bigquery") == "foo(`ds`)" @@ -4124,8 +3986,7 @@ def test_model_normalization(): # since it normalizes values according to the target dialect but not the quotes assert model.depends_on == {'"project-1"."db"."raw"', '"project-2"."db"."raw"'} - expr = d.parse( - """ + expr = d.parse(""" MODEL ( name foo, kind INCREMENTAL_BY_TIME_RANGE ( @@ -4139,8 +4000,7 @@ def test_model_normalization(): clustered_by a ); SELECT * FROM bla - """ - ) + """) model = SqlModel.parse_raw(load_sql_based_model(expr).json()) assert model.name == "foo" @@ -4154,8 +4014,7 @@ def test_model_normalization(): # Check possible variations of unique_key definitions for key in ("""[a, COALESCE(b, ''), "c"]""", """(a, COALESCE(b, ''), "c")"""): - expr = d.parse( - f""" + expr = d.parse(f""" MODEL ( name foo, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -4169,8 +4028,7 @@ def test_model_normalization(): x.b AS b, x."c" AS c FROM test.x AS x - """ - ) + """) model = SqlModel.parse_raw(load_sql_based_model(expr).json()) assert model.unique_key == [ exp.column("A", quoted=True), @@ -4178,8 +4036,7 @@ def test_model_normalization(): exp.column("c", quoted=True), ] - expr = d.parse( - """ + expr = d.parse(""" MODEL ( name foo, dialect snowflake, @@ -4190,8 +4047,7 @@ def test_model_normalization(): ); SELECT x.a AS a FROM test.x AS x - """ - ) + """) model = SqlModel.parse_raw(load_sql_based_model(expr).json()) assert model.unique_key == [exp.column("A", quoted=True)] # we should never normalize the model meta, additionally, we should force lower case @@ -4201,7 +4057,9 @@ def test_model_normalization(): "foo", parse_one("SELECT * FROM bla"), columns={"a": "int"}, - kind=IncrementalByTimeRangeKind(time_column=exp.column("a"), dialect="snowflake"), + kind=IncrementalByTimeRangeKind( + time_column=exp.column("a"), dialect="snowflake" + ), dialect="snowflake", grain=[exp.to_column("id"), exp.to_column("ds")], tags=["pii", "fact"], @@ -4240,8 +4098,7 @@ def test_model_normalization(): @pytest.mark.parametrize("keyword", ["AUTO", "NONE"]) def test_clustered_by_keyword(keyword: str): # Via SQL DDL - expr = d.parse( - f""" + expr = d.parse(f""" MODEL ( name db.test, kind FULL, @@ -4249,8 +4106,7 @@ def test_clustered_by_keyword(keyword: str): clustered_by {keyword} ); SELECT 1 AS a - """ - ) + """) model = load_sql_based_model(expr) assert len(model.clustered_by) == 1 assert model.clustered_by[0].sql(dialect="databricks").upper() == keyword @@ -4285,8 +4141,7 @@ def test_clustered_by_keyword(keyword: str): def test_clustered_by_quoted_keyword_column(): """A backtick-quoted column named `auto` or `none` is a real column, not a keyword.""" for name in ("auto", "none"): - expr = d.parse( - f""" + expr = d.parse(f""" MODEL ( name db.test, kind FULL, @@ -4294,8 +4149,7 @@ def test_clustered_by_quoted_keyword_column(): clustered_by `{name}` ); SELECT 1 AS `{name}` - """ - ) + """) model = load_sql_based_model(expr) assert len(model.clustered_by) == 1 # Must be a Column (quoted identifier), not treated as a keyword @@ -4308,9 +4162,7 @@ def test_clustered_by_quoted_keyword_column(): def test_clustered_by_keyword_non_databricks_dialect(keyword: str): """AUTO/NONE should be rejected for non-Databricks dialects as they are meaningless there.""" with pytest.raises(ConfigError): - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.test, kind FULL, @@ -4318,17 +4170,14 @@ def test_clustered_by_keyword_non_databricks_dialect(keyword: str): clustered_by {keyword} ); SELECT 1 AS a - """ - ) - ) + """)) model.validate_definition() @pytest.mark.parametrize("keyword", ["AUTO", "NONE"]) def test_clustered_by_mixed_list_pins_behaviour(keyword: str): """clustered_by (a, AUTO) — AUTO alongside a real column is treated as a column named AUTO.""" - expr = d.parse( - f""" + expr = d.parse(f""" MODEL ( name db.test, kind FULL, @@ -4336,8 +4185,7 @@ def test_clustered_by_mixed_list_pins_behaviour(keyword: str): clustered_by (a, {keyword}) ); SELECT 1 AS a, 2 AS {keyword.lower()} - """ - ) + """) model = load_sql_based_model(expr) # Both entries are real columns (AUTO/NONE inside parens is a column, not a keyword) assert len(model.clustered_by) == 2 @@ -4348,9 +4196,7 @@ def test_clustered_by_mixed_list_pins_behaviour(keyword: str): @pytest.mark.parametrize("keyword", ["AUTO", "NONE"]) def test_clustered_by_keyword_serialisation_round_trip(keyword: str): """exp.Var(AUTO/NONE) must survive JSON serialisation and deserialisation unchanged.""" - model = load_sql_based_model( - d.parse( - f""" + model = load_sql_based_model(d.parse(f""" MODEL ( name db.test, kind FULL, @@ -4358,9 +4204,7 @@ def test_clustered_by_keyword_serialisation_round_trip(keyword: str): clustered_by {keyword} ); SELECT 1 AS a - """ - ) - ) + """)) model_json = model.json() deserialized = SqlModel.parse_raw(model_json) assert deserialized.clustered_by == model.clustered_by @@ -4387,24 +4231,21 @@ def test_incremental_unmanaged_validation(): def test_incremental_unmanaged(): - expr = d.parse( - """ + expr = d.parse(""" MODEL ( name foo, kind INCREMENTAL_UNMANAGED ); SELECT x.a AS a FROM test.x AS x - """ - ) + """) model = load_sql_based_model(expressions=expr) assert isinstance(model.kind, IncrementalUnmanagedKind) assert not model.kind.insert_overwrite - expr = d.parse( - """ + expr = d.parse(""" MODEL ( name foo, kind INCREMENTAL_UNMANAGED ( @@ -4414,8 +4255,7 @@ def test_incremental_unmanaged(): ); SELECT x.a AS a FROM test.x AS x - """ - ) + """) model = load_sql_based_model(expressions=expr) assert isinstance(model.kind, IncrementalUnmanagedKind) @@ -4425,7 +4265,9 @@ def test_incremental_unmanaged(): def test_custom_interval_unit(): assert ( load_sql_based_model( - d.parse("MODEL (name db.table, interval_unit FIVE_MINUTE); SELECT a FROM tbl;") + d.parse( + "MODEL (name db.table, interval_unit FIVE_MINUTE); SELECT a FROM tbl;" + ) ).interval_unit == IntervalUnit.FIVE_MINUTE ) @@ -4439,7 +4281,9 @@ def test_custom_interval_unit(): assert ( load_sql_based_model( - d.parse("MODEL (name db.table, interval_unit Hour, cron '@daily'); SELECT a FROM tbl;") + d.parse( + "MODEL (name db.table, interval_unit Hour, cron '@daily'); SELECT a FROM tbl;" + ) ).interval_unit == IntervalUnit.HOUR ) @@ -4468,7 +4312,8 @@ def test_custom_interval_unit(): ) with pytest.raises( - ConfigError, match=r"Cron '@daily' cannot be more frequent than interval unit 'month'." + ConfigError, + match=r"Cron '@daily' cannot be more frequent than interval unit 'month'.", ): load_sql_based_model( d.parse("MODEL (name db.table, interval_unit month); SELECT a FROM tbl;") @@ -4479,7 +4324,9 @@ def test_custom_interval_unit(): match=r"Cron '@hourly' cannot be more frequent than interval unit 'day'. If this is intentional, set allow_partials to True.", ): load_sql_based_model( - d.parse("MODEL (name db.table, interval_unit Day, cron '@hourly'); SELECT a FROM tbl;") + d.parse( + "MODEL (name db.table, interval_unit Day, cron '@hourly'); SELECT a FROM tbl;" + ) ) @@ -4499,7 +4346,9 @@ def test_interval_unit_larger_than_cron_period(): match=r"Cron '@hourly' cannot be more frequent than interval unit 'day'. If this is intentional, set allow_partials to True.", ): load_sql_based_model( - d.parse("MODEL (name db.table, interval_unit day, cron '@hourly'); SELECT a FROM tbl;") + d.parse( + "MODEL (name db.table, interval_unit day, cron '@hourly'); SELECT a FROM tbl;" + ) ) with pytest.raises( @@ -4542,9 +4391,7 @@ def my_model(context, **kwargs): } # Validate a tuple. - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, physical_properties ( @@ -4555,9 +4402,7 @@ def my_model(context, **kwargs): ) ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == { "key_a": exp.convert("value_a"), "key_b": exp.convert(1), @@ -4569,17 +4414,13 @@ def my_model(context, **kwargs): ) # Validate a tuple with one item. - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, physical_properties (key_a = 'value_a') ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == {"key_a": exp.convert("value_a")} assert ( sql_model.physical_properties_.sql() # type: ignore @@ -4587,9 +4428,7 @@ def my_model(context, **kwargs): ) # Validate an array. - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, physical_properties [ @@ -4598,33 +4437,27 @@ def my_model(context, **kwargs): ] ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == { "key_a": exp.convert("value_a"), "key_b": exp.convert(1), } - assert sql_model.physical_properties_ == d.parse_one("""(key_a = 'value_a', 'key_b' = 1)""") + assert sql_model.physical_properties_ == d.parse_one( + """(key_a = 'value_a', 'key_b' = 1)""" + ) # Validate empty. - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == {} assert sql_model.physical_properties_ is None # Validate sql expression. - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, physical_properties [ @@ -4632,11 +4465,11 @@ def my_model(context, **kwargs): ] ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == {"key": d.parse_one("['value']")} - assert sql_model.physical_properties_ == exp.Tuple(expressions=[d.parse_one("key = ['value']")]) + assert sql_model.physical_properties_ == exp.Tuple( + expressions=[d.parse_one("key = ['value']")] + ) # Validate dict parsing. sql_model = create_sql_model( @@ -4663,9 +4496,7 @@ def my_model(context, **kwargs): ConfigError, match=r"Invalid property 'invalid'. Properties must be specified as key-value pairs = . ", ): - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, physical_properties [ @@ -4673,15 +4504,11 @@ def my_model(context, **kwargs): ] ); SELECT a FROM tbl; - """ - ) - ) + """)) def test_model_physical_properties_labels() -> None: - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, physical_properties [ @@ -4689,16 +4516,14 @@ def test_model_physical_properties_labels() -> None: ] ); SELECT a FROM tbl; - """ - ) - ) - assert sql_model.physical_properties == {"labels": exp.array("('test-label', 'label-value')")} + """)) + assert sql_model.physical_properties == { + "labels": exp.array("('test-label', 'label-value')") + } def test_physical_and_virtual_table_properties() -> None: - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, physical_properties ( @@ -4710,9 +4535,7 @@ def test_physical_and_virtual_table_properties() -> None: ) ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == { "partition_expiration_days": exp.convert(7), "labels": exp.array("('test-physical-label', 'label-physical-value')"), @@ -4725,9 +4548,7 @@ def test_physical_and_virtual_table_properties() -> None: def test_model_table_properties() -> None: # Ensure backward compatibility to table_properties. - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, table_properties ( @@ -4738,9 +4559,7 @@ def test_model_table_properties() -> None: ) ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == { "key_a": exp.convert("value_a"), "key_b": exp.convert(1), @@ -4751,9 +4570,7 @@ def test_model_table_properties() -> None: """(key_a = 'value_a', 'key_b' = 1, key_c = TRUE, "key_d" = 2.0)""" ) - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, table_properties ( @@ -4761,21 +4578,19 @@ def test_model_table_properties() -> None: ) ); SELECT a FROM tbl; - """ - ) - ) + """)) assert sql_model.physical_properties == { "partition_expiration_days": exp.convert(7), } - assert sql_model.physical_properties_ == d.parse_one("""(partition_expiration_days = 7,)""") + assert sql_model.physical_properties_ == d.parse_one( + """(partition_expiration_days = 7,)""" + ) def test_model_table_properties_conflicts() -> None: # Throw an error on conflicting usage of table_properties and physical_properties. with pytest.raises(ConfigError, match=r"Cannot use argument 'table_properties'*"): - sql_model = load_sql_based_model( - d.parse( - """ + sql_model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, table_properties ( @@ -4787,9 +4602,7 @@ def test_model_table_properties_conflicts() -> None: physical_properties (key_a = 'value_a') ); SELECT a FROM tbl; - """ - ) - ) + """)) sql_model.physical_properties @@ -4973,7 +4786,8 @@ def test_conditional_physical_properties(make_snapshot): == full_model.physical_properties == { "creatable_type": exp.maybe_parse( - "@IF(@model_kind_name != 'VIEW', 'TRANSIENT', NULL)", dialect="snowflake" + "@IF(@model_kind_name != 'VIEW', 'TRANSIENT', NULL)", + dialect="snowflake", ) } ) @@ -5229,7 +5043,11 @@ def python_model_prop_macro(context, **kwargs): path=Path("."), dialect="duckdb", defaults=model_defaults, - variables={"gateway": "local", "create_type": "SECURE", "cron_macro_expr": "0 */2 * * *"}, + variables={ + "gateway": "local", + "create_type": "SECURE", + "cron_macro_expr": "0 */2 * * *", + }, ) # Even if in the project wide defaults this is ignored for python models @@ -5268,7 +5086,9 @@ def python_model_prop_macro(context, **kwargs): } assert m.session_properties == { - "spark.executor.cores": exp.maybe_parse("@IF(@gateway = 'dev', 1, 2)", dialect="duckdb"), + "spark.executor.cores": exp.maybe_parse( + "@IF(@gateway = 'dev', 1, 2)", dialect="duckdb" + ), "spark.executor.memory": "1G", } @@ -5401,7 +5221,9 @@ def test_model_session_properties(sushi_context): ) ) assert model.session_properties == { - "query_label": parse_one("[('key1', 'value1'), ('key2', 'value2')]", dialect="bigquery") + "query_label": parse_one( + "[('key1', 'value1'), ('key2', 'value2')]", dialect="bigquery" + ) } model = load_sql_based_model( @@ -5420,7 +5242,9 @@ def test_model_session_properties(sushi_context): default_dialect="bigquery", ) ) - assert model.session_properties == {"query_label": parse_one("(('key1', 'value1'))")} + assert model.session_properties == { + "query_label": parse_one("(('key1', 'value1'))") + } with pytest.raises( ConfigError, @@ -5646,8 +5470,7 @@ def test_session_properties_query_tags_validation(): def test_model_jinja_macro_rendering(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -5660,12 +5483,13 @@ def test_model_jinja_macro_rendering(): JINJA_END; SELECT 1 AS x; - """ - ) + """) jinja_macros = JinjaMacroRegistry( packages={ - "test_package": {"macro_a": MacroInfo(definition="macro_a_body", depends_on=[])}, + "test_package": { + "macro_a": MacroInfo(definition="macro_a_body", depends_on=[]) + }, }, root_macros={"macro_b": MacroInfo(definition="macro_b_body", depends_on=[])}, global_objs={"test_int": 1, "test_str": "value"}, @@ -5679,19 +5503,16 @@ def test_model_jinja_macro_rendering(): def test_view_model_data_hash(): - view_model_expressions = d.parse( - """ + view_model_expressions = d.parse(""" MODEL ( name db.table, kind VIEW, ); SELECT 1; - """ - ) + """) view_model_hash = load_sql_based_model(view_model_expressions).data_hash - materialized_view_model_expressions = d.parse( - """ + materialized_view_model_expressions = d.parse(""" MODEL ( name db.table, kind VIEW ( @@ -5699,8 +5520,7 @@ def test_view_model_data_hash(): ), ); SELECT 1; - """ - ) + """) materialized_view_model_hash = load_sql_based_model( materialized_view_model_expressions ).data_hash @@ -5709,8 +5529,7 @@ def test_view_model_data_hash(): def test_view_materialized_partition_by_clustered_by(): - materialized_view_model_expressions = d.parse( - """ + materialized_view_model_expressions = d.parse(""" MODEL ( name db.table, kind VIEW ( @@ -5720,60 +5539,56 @@ def test_view_materialized_partition_by_clustered_by(): clustered_by a ); SELECT 1; - """ - ) + """) materialized_view_model = load_sql_based_model(materialized_view_model_expressions) assert materialized_view_model.partitioned_by == [exp.column("ds", quoted=True)] assert materialized_view_model.clustered_by == [exp.to_column('"a"')] def test_view_non_materialized_partition_by(): - view_model_expressions = d.parse( - """ + view_model_expressions = d.parse(""" MODEL ( name db.table, kind VIEW, partitioned_by ds, ); SELECT 1; - """ - ) - with pytest.raises(ValidationError, match=r".*partitioned_by field cannot be set for VIEW.*"): + """) + with pytest.raises( + ValidationError, match=r".*partitioned_by field cannot be set for VIEW.*" + ): load_sql_based_model(view_model_expressions) def test_view_non_materialized_clustered_by(): - view_model_expressions = d.parse( - """ + view_model_expressions = d.parse(""" MODEL ( name db.table, kind VIEW, clustered_by ds, ); SELECT 1; - """ - ) - with pytest.raises(ValidationError, match=r".*clustered_by field cannot be set for VIEW.*"): + """) + with pytest.raises( + ValidationError, match=r".*clustered_by field cannot be set for VIEW.*" + ): load_sql_based_model(view_model_expressions) def test_seed_model_data_hash(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( path '../seeds/waiter_names.csv', ) ); - """ - ) + """) seed_model = load_sql_based_model( expressions, path=Path("./examples/sushi/models/test_model.sql") ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -5783,8 +5598,7 @@ def test_seed_model_data_hash(): ) ) ); - """ - ) + """) new_seed_model = load_sql_based_model( expressions, path=Path("./examples/sushi/models/test_model.sql") ) @@ -5822,8 +5636,7 @@ def test_interval_unit_validation(): def test_scd_type_2_by_time_defaults(): - model_def = d.parse( - """ + model_def = d.parse(""" MODEL ( name db.table, kind SCD_TYPE_2 ( @@ -5837,8 +5650,7 @@ def test_scd_type_2_by_time_defaults(): '2020-01-01' as test_valid_from, '2020-01-01' as test_valid_to ; - """ - ) + """) scd_type_2_model = load_sql_based_model(model_def) assert scd_type_2_model.unique_key == [ parse_one("""COALESCE("ID", '') || '|' || COALESCE("ds", '')"""), @@ -5857,8 +5669,12 @@ def test_scd_type_2_by_time_defaults(): "valid_from": exp.DataType.build("TIMESTAMP"), "valid_to": exp.DataType.build("TIMESTAMP"), } - assert scd_type_2_model.kind.updated_at_name == exp.column("updated_at", quoted=True) - assert scd_type_2_model.kind.valid_from_name == exp.column("valid_from", quoted=True) + assert scd_type_2_model.kind.updated_at_name == exp.column( + "updated_at", quoted=True + ) + assert scd_type_2_model.kind.valid_from_name == exp.column( + "valid_from", quoted=True + ) assert scd_type_2_model.kind.valid_to_name == exp.column("valid_to", quoted=True) assert not scd_type_2_model.kind.updated_at_as_valid_from assert scd_type_2_model.kind.is_scd_type_2_by_time @@ -5869,8 +5685,7 @@ def test_scd_type_2_by_time_defaults(): def test_scd_type_2_by_time_overrides(): - model_def = d.parse( - """ + model_def = d.parse(""" MODEL ( name db.table, kind SCD_TYPE_2_BY_TIME ( @@ -5893,8 +5708,7 @@ def test_scd_type_2_by_time_overrides(): '2020-01-01' as test_valid_from, '2020-01-01' as test_valid_to ; - """ - ) + """) scd_type_2_model = load_sql_based_model(model_def) assert scd_type_2_model.unique_key == [ exp.column("iD", quoted=True), @@ -5904,9 +5718,15 @@ def test_scd_type_2_by_time_overrides(): "TEST_VALID_FROM": exp.DataType.build("TIMESTAMPTZ"), "TEST_VALID_TO": exp.DataType.build("TIMESTAMPTZ"), } - assert scd_type_2_model.kind.updated_at_name == exp.column("TEST_UPDATED_AT", quoted=True) - assert scd_type_2_model.kind.valid_from_name == exp.column("TEST_VALID_FROM", quoted=True) - assert scd_type_2_model.kind.valid_to_name == exp.column("TEST_VALID_TO", quoted=True) + assert scd_type_2_model.kind.updated_at_name == exp.column( + "TEST_UPDATED_AT", quoted=True + ) + assert scd_type_2_model.kind.valid_from_name == exp.column( + "TEST_VALID_FROM", quoted=True + ) + assert scd_type_2_model.kind.valid_to_name == exp.column( + "TEST_VALID_TO", quoted=True + ) assert scd_type_2_model.kind.updated_at_as_valid_from assert scd_type_2_model.kind.is_scd_type_2_by_time assert scd_type_2_model.kind.is_scd_type_2 @@ -5920,8 +5740,7 @@ def test_scd_type_2_by_time_overrides(): def test_scd_type_2_by_column_defaults(): - model_def = d.parse( - """ + model_def = d.parse(""" MODEL ( name db.table, kind SCD_TYPE_2_BY_COLUMN ( @@ -5934,11 +5753,12 @@ def test_scd_type_2_by_column_defaults(): 2 as "value_to_track", '2020-01-01' as ds, ; - """ - ) + """) scd_type_2_model = load_sql_based_model(model_def) assert scd_type_2_model.unique_key == [exp.to_column("ID", quoted=True)] - assert scd_type_2_model.kind.columns == [exp.to_column("value_to_track", quoted=True)] + assert scd_type_2_model.kind.columns == [ + exp.to_column("value_to_track", quoted=True) + ] assert scd_type_2_model.columns_to_types == { "ID": exp.DataType.build("int"), "value_to_track": exp.DataType.build("int"), @@ -5950,7 +5770,9 @@ def test_scd_type_2_by_column_defaults(): "valid_from": exp.DataType.build("TIMESTAMP"), "valid_to": exp.DataType.build("TIMESTAMP"), } - assert scd_type_2_model.kind.valid_from_name == exp.column("valid_from", quoted=True) + assert scd_type_2_model.kind.valid_from_name == exp.column( + "valid_from", quoted=True + ) assert scd_type_2_model.kind.valid_to_name == exp.column("valid_to", quoted=True) assert not scd_type_2_model.kind.execution_time_as_valid_from assert scd_type_2_model.kind.is_scd_type_2_by_column @@ -5961,8 +5783,7 @@ def test_scd_type_2_by_column_defaults(): def test_scd_type_2_by_column_overrides(): - model_def = d.parse( - """ + model_def = d.parse(""" MODEL ( name db.table, kind SCD_TYPE_2_BY_COLUMN ( @@ -5983,8 +5804,7 @@ def test_scd_type_2_by_column_overrides(): 2 as "value_to_track", '2020-01-01' as ds, ; - """ - ) + """) scd_type_2_model = load_sql_based_model(model_def) assert scd_type_2_model.unique_key == [ exp.column("iD", quoted=True), @@ -5994,8 +5814,12 @@ def test_scd_type_2_by_column_overrides(): "test_valid_from": exp.DataType.build("TIMESTAMPTZ"), "test_valid_to": exp.DataType.build("TIMESTAMPTZ"), } - assert scd_type_2_model.kind.valid_from_name == exp.column("test_valid_from", quoted=True) - assert scd_type_2_model.kind.valid_to_name == exp.column("test_valid_to", quoted=True) + assert scd_type_2_model.kind.valid_from_name == exp.column( + "test_valid_from", quoted=True + ) + assert scd_type_2_model.kind.valid_to_name == exp.column( + "test_valid_to", quoted=True + ) assert scd_type_2_model.kind.execution_time_as_valid_from assert scd_type_2_model.kind.is_scd_type_2_by_column assert scd_type_2_model.kind.is_scd_type_2 @@ -6071,8 +5895,7 @@ def scd_type_2_model(context, **kwargs): ], ) def test_check_column_variants(input_columns, expected_columns): - model_def = d.parse( - f""" + model_def = d.parse(f""" MODEL ( name db.table, kind SCD_TYPE_2_BY_COLUMN ( @@ -6082,22 +5905,19 @@ def test_check_column_variants(input_columns, expected_columns): ); SELECT 1 ; - """ - ) + """) scd_type_2_model = load_sql_based_model(model_def) assert scd_type_2_model.kind.columns == expected_columns def test_model_dialect_name(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name `project-1`.`db`.`tbl1`, dialect bigquery ); SELECT 1; - """ - ) + """) model = load_sql_based_model(expressions) assert model.fqn == '"project-1"."db"."tbl1"' @@ -6105,22 +5925,22 @@ def test_model_dialect_name(): model = create_external_model( "`project-1`.`db`.`tbl1`", columns={"x": "STRING"}, dialect="bigquery" ) - assert "name `project-1`.`db`.`tbl1`" in model.render_definition()[0].sql(dialect="bigquery") + assert "name `project-1`.`db`.`tbl1`" in model.render_definition()[0].sql( + dialect="bigquery" + ) # This used to fail due to the dialect regex picking up `DIALECT_TEST` as the model's dialect expressions = d.parse("MODEL(name DIALECT_TEST.foo); SELECT 1") def test_model_allow_partials(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, allow_partials true, ); SELECT 1; - """ - ) + """) model = load_sql_based_model(expressions) @@ -6130,8 +5950,7 @@ def test_model_allow_partials(): def test_signals(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, signals [ @@ -6139,8 +5958,7 @@ def test_signals(): ], ); SELECT 1; - """ - ) + """) model = load_sql_based_model(expressions) assert model.signals[0][1] == {"arg": exp.Literal.number(1)} @@ -6149,8 +5967,7 @@ def test_signals(): def my_signal(batch): return True - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, signals [ @@ -6173,8 +5990,7 @@ def my_signal(batch): ], ); SELECT 1; - """ - ) + """) model = load_sql_based_model( expressions, @@ -6213,7 +6029,9 @@ def my_signal(batch): ), ] - rendered_signals = model.render_signals(start="2023-01-01", end="2023-01-02 15:00:00") + rendered_signals = model.render_signals( + start="2023-01-01", end="2023-01-02 15:00:00" + ) assert rendered_signals == [ {"table_name": "table_a", "ds": "2023-01-02"}, {"table_name": "table_b", "ds": "2023-01-02", "hour": 14}, @@ -6251,8 +6069,7 @@ def model_with_signal(context, **kwargs): def test_null_column_type(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test_db.test_model, columns ( @@ -6265,8 +6082,7 @@ def test_null_column_type(): id::INT AS id, ds FROM x - """ - ) + """) model = load_sql_based_model(expressions, dialect="hive") assert model.columns_to_types == { "ds": exp.DataType.build("null"), @@ -6276,8 +6092,7 @@ def test_null_column_type(): def test_when_matched(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -6286,8 +6101,7 @@ def test_when_matched(): ) ); SELECT 'name' AS name, 1 AS salary; - """ - ) + """) expected_when_matched = "(WHEN MATCHED THEN UPDATE SET `__MERGE_TARGET__`.`salary` = COALESCE(`__MERGE_SOURCE__`.`salary`, `__MERGE_TARGET__`.`salary`))" @@ -6297,8 +6111,7 @@ def test_when_matched(): model = SqlModel.parse_raw(model.json()) assert model.kind.when_matched.sql(dialect="hive") == expected_when_matched - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name @{macro_val}.test, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -6313,12 +6126,10 @@ def test_when_matched(): SELECT purchase_order_id FROM @{macro_val}.upstream - """ - ) + """) model = SqlModel.parse_raw(load_sql_based_model(expressions).json()) - assert d.format_model_expressions(model.render_definition()) == ( - """MODEL ( + assert d.format_model_expressions(model.render_definition()) == ("""MODEL ( name @{macro_val}.test, kind INCREMENTAL_BY_UNIQUE_KEY ( unique_key ("purchase_order_id"), @@ -6337,8 +6148,7 @@ def test_when_matched(): SELECT purchase_order_id -FROM @{macro_val}.upstream""" - ) +FROM @{macro_val}.upstream""") @macro() def fingerprint_merge( @@ -6346,15 +6156,18 @@ def fingerprint_merge( fingerprint_column: exp.Column, update_columns: list[exp.Column], ) -> exp.Whens: - fingerprint_evaluation = f"source.{fingerprint_column} <> target.{fingerprint_column}" - column_update = [f"target.{column} = source.{column}" for column in update_columns] + fingerprint_evaluation = ( + f"source.{fingerprint_column} <> target.{fingerprint_column}" + ) + column_update = [ + f"target.{column} = source.{column}" for column in update_columns + ] return exp.maybe_parse( f"WHEN MATCHED AND {fingerprint_evaluation} THEN UPDATE SET {column_update}", into=exp.Whens, ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -6367,12 +6180,10 @@ def fingerprint_merge( 1 AS purchase_order_id, 1 AS salary, CAST('2020-01-01 12:05:01' AS DATETIME) AS update_datetime - """ - ) + """) model = SqlModel.parse_raw(load_sql_based_model(expressions).json()) - assert d.format_model_expressions(model.render_definition()) == ( - """MODEL ( + assert d.format_model_expressions(model.render_definition()) == ("""MODEL ( name test, kind INCREMENTAL_BY_UNIQUE_KEY ( unique_key ("purchase_order_id"), @@ -6391,13 +6202,11 @@ def fingerprint_merge( SELECT 1 AS purchase_order_id, 1 AS salary, - '2020-01-01 12:05:01'::DATETIME AS update_datetime""" - ) + '2020-01-01 12:05:01'::DATETIME AS update_datetime""") def test_when_matched_multiple(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name @{schema}.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -6408,15 +6217,16 @@ def test_when_matched_multiple(): ) ); SELECT 'name' AS name, 1 AS salary; - """ - ) + """) expected_when_matched = [ "WHEN MATCHED AND `__MERGE_SOURCE__`.`x` = 1 THEN UPDATE SET `__MERGE_TARGET__`.`salary` = COALESCE(`__MERGE_SOURCE__`.`salary`, `__MERGE_TARGET__`.`salary`)", "WHEN MATCHED THEN UPDATE SET `__MERGE_TARGET__`.`salary` = COALESCE(`__MERGE_SOURCE__`.`salary`, `__MERGE_TARGET__`.`salary`)", ] - model = load_sql_based_model(expressions, dialect="hive", variables={"schema": "db"}) + model = load_sql_based_model( + expressions, dialect="hive", variables={"schema": "db"} + ) whens = model.kind.when_matched assert len(whens.expressions) == 2 assert whens.expressions[0].sql(dialect="hive") == expected_when_matched[0] @@ -6430,8 +6240,7 @@ def test_when_matched_multiple(): def test_when_matched_merge_filter_multi_part_columns(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name @{schema}.records_model, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -6450,8 +6259,7 @@ def test_when_matched_merge_filter_multi_part_columns(): ) AS record FROM @{schema}.seed_model; - """ - ) + """) expected_when_matched = [ "WHEN MATCHED AND `__MERGE_SOURCE__`.`record`.`nested_record`.`field` = 1 THEN UPDATE SET `__MERGE_TARGET__`.`repeated_record`.`sub_repeated_record`.`sub_field` = COALESCE(`__MERGE_SOURCE__`.`repeated_record`.`sub_repeated_record`.`sub_field`, `__MERGE_TARGET__`.`repeated_record`.`sub_repeated_record`.`sub_field`)", @@ -6463,7 +6271,9 @@ def test_when_matched_merge_filter_multi_part_columns(): "`__MERGE_TARGET__`.`repeated_record`.`sub_repeated_record`.`sub_field` > `__MERGE_SOURCE__`.`repeated_record`.`sub_repeated_record`.`sub_field`" ) - model = load_sql_based_model(expressions, dialect="bigquery", variables={"schema": "db"}) + model = load_sql_based_model( + expressions, dialect="bigquery", variables={"schema": "db"} + ) whens = model.kind.when_matched assert len(whens.expressions) == 2 assert whens.expressions[0].sql(dialect="bigquery") == expected_when_matched[0] @@ -6480,8 +6290,7 @@ def test_when_matched_merge_filter_multi_part_columns(): def test_when_matched_normalization() -> None: # unquoted should be normalized and quoted - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -6494,8 +6303,7 @@ def test_when_matched_normalization() -> None: ) ); SELECT 'name' AS name, 1 AS key_a, 2 AS key_b; - """ - ) + """) model = load_sql_based_model(expressions, dialect="snowflake") assert isinstance(model.kind, IncrementalByUniqueKeyKind) @@ -6508,8 +6316,7 @@ def test_when_matched_normalization() -> None: ) # quoted should be preserved - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -6522,8 +6329,7 @@ def test_when_matched_normalization() -> None: ) ); SELECT 'name' AS name, 1 AS "kEy_A", 2 AS "kEY_b"; - """ - ) + """) model = load_sql_based_model(expressions, dialect="snowflake") assert isinstance(model.kind, IncrementalByUniqueKeyKind) @@ -6545,15 +6351,13 @@ def test_default_catalog_sql(assert_exp_eq): HASH_WITH_CATALOG = "2768215345" # Test setting default catalog doesn't change hash if it matches existing logic - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name catalog.db.table ); SELECT x FROM catalog.db.source - """ - ) + """) model = load_sql_based_model(expressions, default_catalog="catalog") assert model.default_catalog == "catalog" @@ -6572,15 +6376,13 @@ def test_default_catalog_sql(assert_exp_eq): assert model.data_hash == HASH_WITH_CATALOG - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name catalog.db.table, ); SELECT x FROM catalog.db.source - """ - ) + """) model = load_sql_based_model(expressions) assert model.default_catalog is None @@ -6600,15 +6402,13 @@ def test_default_catalog_sql(assert_exp_eq): assert model.data_hash == HASH_WITH_CATALOG # Test setting default catalog to a different catalog but everything if fully qualified then no hash change - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name catalog.db.table ); SELECT x FROM catalog.db.source - """ - ) + """) model = load_sql_based_model(expressions, default_catalog="other_catalog") assert model.default_catalog == "other_catalog" @@ -6628,15 +6428,13 @@ def test_default_catalog_sql(assert_exp_eq): assert model.data_hash == HASH_WITH_CATALOG # test that hash changes if model contains a non-fully-qualified reference - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name catalog.db.table ); SELECT x FROM db.source - """ - ) + """) model = load_sql_based_model(expressions, default_catalog="other_catalog") assert model.default_catalog == "other_catalog" @@ -6649,15 +6447,13 @@ def test_default_catalog_sql(assert_exp_eq): # test that hash is the same but the fqn is different so the snapshot is different so this is # a new snapshot but with the same hash as before - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, ); SELECT x FROM catalog.db.source - """ - ) + """) model = load_sql_based_model(expressions) assert model.default_catalog is None @@ -6669,15 +6465,13 @@ def test_default_catalog_sql(assert_exp_eq): # This will also have the same hash but the fqn is different so the snapshot is different so this is # a new snapshot but with the same hash as before - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table ); SELECT x FROM catalog.db.source - """ - ) + """) model = load_sql_based_model(expressions, default_catalog="catalog") assert model.default_catalog == "catalog" @@ -6688,15 +6482,13 @@ def test_default_catalog_sql(assert_exp_eq): assert model.data_hash == HASH_WITH_CATALOG # Query is different since default catalog does not apply and therefore the hash is different - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name table ); SELECT x FROM source - """ - ) + """) model = load_sql_based_model(expressions, default_catalog="catalog") assert model.default_catalog == "catalog" @@ -6819,7 +6611,9 @@ def test_default_catalog_external_model(): assert model.data_hash == EXPECTED_HASH model = create_external_model( - "catalog.db.table", columns={"a": "int", "limit": "int"}, default_catalog="other_catalog" + "catalog.db.table", + columns={"a": "int", "limit": "int"}, + default_catalog="other_catalog", ) assert model.default_catalog == "other_catalog" assert model.name == "catalog.db.table" @@ -6840,29 +6634,41 @@ def test_default_catalog_external_model(): def test_user_cannot_set_default_catalog(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, default_catalog some_catalog ); SELECT 1::int AS a, 2::int AS b, 3 AS c, 4 as d; - """ - ) + """) - with pytest.raises(ConfigError, match="`default_catalog` cannot be set on a per-model basis"): + with pytest.raises( + ConfigError, match="`default_catalog` cannot be set on a per-model basis" + ): load_sql_based_model(expressions) - with pytest.raises(ConfigError, match="`default_catalog` cannot be set on a per-model basis"): + with pytest.raises( + ConfigError, match="`default_catalog` cannot be set on a per-model basis" + ): - @model(name="db.table", kind="full", columns={'"COL"': "int"}, default_catalog="catalog") + @model( + name="db.table", + kind="full", + columns={'"COL"': "int"}, + default_catalog="catalog", + ) def my_model(context, **kwargs): context.resolve_table("dependency.table") def test_depends_on_default_catalog_python(): - @model(name="some.table", kind="full", columns={'"COL"': "int"}, depends_on={"other.table"}) + @model( + name="some.table", + kind="full", + columns={'"COL"': "int"}, + depends_on={"other.table"}, + ) def my_model(context, **kwargs): context.resolve_table("dependency.table") @@ -6877,8 +6683,7 @@ def my_model(context, **kwargs): def test_end_date(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -6889,18 +6694,17 @@ def test_end_date(): ); SELECT 1::int AS a, 2::int AS b, now::timestamp as ts - """ - ) + """) model = load_sql_based_model(expressions) assert model.start == "2023-01-01" assert model.end == "2023-06-01" assert model.interval_unit == IntervalUnit.DAY - with pytest.raises(ValidationError, match=".*Start date.+can't be greater than end date.*"): - load_sql_based_model( - d.parse( - """ + with pytest.raises( + ValidationError, match=".*Start date.+can't be greater than end date.*" + ): + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -6911,14 +6715,11 @@ def test_end_date(): ); SELECT 1::int AS a, 2::int AS b, now::timestamp as ts - """ - ) - ) + """)) def test_end_no_start(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -6928,9 +6729,10 @@ def test_end_no_start(): ); SELECT 1::int AS a, 2::int AS b, now::timestamp as ts - """ - ) - with pytest.raises(ConfigError, match="Must define a start date if an end date is defined"): + """) + with pytest.raises( + ConfigError, match="Must define a start date if an end date is defined" + ): load_sql_based_model(expressions) load_sql_based_model(expressions, defaults={"start": "2023-01-01"}) @@ -6977,7 +6779,9 @@ def test_macro_var(evaluator) -> exp.Expr: == "SELECT 'test_value' AS `a`, 'default_value' AS `b`, NULL AS `c`, 11 AS `d`, 'foo_4' AS `e`, `foo_5` AS `f`, 'foo_@{test_var_unused}' AS `g`" ) - with pytest.raises(ConfigError, match=r"Macro VAR requires at least one argument.*"): + with pytest.raises( + ConfigError, match=r"Macro VAR requires at least one argument.*" + ): expressions = parse( """ MODEL( @@ -6991,7 +6795,8 @@ def test_macro_var(evaluator) -> exp.Expr: load_sql_based_model(expressions) with pytest.raises( - ConfigError, match=r"The variable name must be a string literal, '123' was given instead.*" + ConfigError, + match=r"The variable name must be a string literal, '123' was given instead.*", ): expressions = parse( """ @@ -7025,13 +6830,11 @@ def test_macro_var(evaluator) -> exp.Expr: def test_named_variable_macros() -> None: model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_gateway_macro); @DEF(overridden_var, 'overridden_value'); SELECT @gateway AS gateway, @TEST_VAR_A AS test_var_a, @overridden_var AS overridden_var - """ - ), + """), variables={ c.GATEWAY: "in_memory", "test_var_a": "test_value", @@ -7041,7 +6844,11 @@ def test_named_variable_macros() -> None: ) assert model.python_env[c.SQLMESH_VARS] == Executable.value( - {c.GATEWAY: "in_memory", "test_var_a": "test_value", "overridden_var": "initial_value"}, + { + c.GATEWAY: "in_memory", + "test_var_a": "test_value", + "overridden_var": "initial_value", + }, sort_root_dict=True, ) assert ( @@ -7052,13 +6859,11 @@ def test_named_variable_macros() -> None: def test_variables_in_templates() -> None: model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_gateway_macro); @DEF(overridden_var, overridden_value); SELECT 'gateway' AS col_@gateway, 'test_var_a' AS @{test_var_a}_col, 'overridden_var' AS col_@{overridden_var}_col - """ - ), + """), variables={ c.GATEWAY: "in_memory", "test_var_a": "test_value", @@ -7068,7 +6873,11 @@ def test_variables_in_templates() -> None: ) assert model.python_env[c.SQLMESH_VARS] == Executable.value( - {c.GATEWAY: "in_memory", "test_var_a": "test_value", "overridden_var": "initial_value"}, + { + c.GATEWAY: "in_memory", + "test_var_a": "test_value", + "overridden_var": "initial_value", + }, sort_root_dict=True, ) assert ( @@ -7077,13 +6886,11 @@ def test_variables_in_templates() -> None: ) model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_gateway_macro); @DEF(overridden_var, overridden_value); SELECT 'combo' AS col_@{test_var_a}_@{overridden_var}_col_@gateway - """ - ), + """), variables={ c.GATEWAY: "in_memory", "test_var_a": "test_value", @@ -7093,7 +6900,11 @@ def test_variables_in_templates() -> None: ) assert model.python_env[c.SQLMESH_VARS] == Executable.value( - {c.GATEWAY: "in_memory", "test_var_a": "test_value", "overridden_var": "initial_value"}, + { + c.GATEWAY: "in_memory", + "test_var_a": "test_value", + "overridden_var": "initial_value", + }, sort_root_dict=True, ) assert ( @@ -7102,16 +6913,14 @@ def test_variables_in_templates() -> None: ) model = load_sql_based_model( - parse( - """ + parse(""" MODEL( name @{some_var}.bar, dialect snowflake ); SELECT 1 AS c - """ - ), + """), variables={ "some_var": "foo", }, @@ -7182,11 +6991,15 @@ def model_with_variables(context, **kwargs): ) assert python_model.name == "foo_suffix" - assert python_model.python_env[c.SQLMESH_VARS] == Executable.value({"test_var_a": "test_value"}) + assert python_model.python_env[c.SQLMESH_VARS] == Executable.value( + {"test_var_a": "test_value"} + ) context = ExecutionContext(mocker.Mock(), {}, None, None) df = list(python_model.render(context=context))[0] - assert df.to_dict(orient="records") == [{"a": "test_value", "b": "default_value", "c": None}] + assert df.to_dict(orient="records") == [ + {"a": "test_value", "b": "default_value", "c": None} + ] def test_load_external_model_python(sushi_context) -> None: @@ -7206,7 +7019,9 @@ def external_model_python(context, **kwargs): path=Path("."), ) - context = ExecutionContext(sushi_context.engine_adapter, sushi_context.snapshots, None, None) + context = ExecutionContext( + sushi_context.engine_adapter, sushi_context.snapshots, None, None + ) df = list(python_model.render(context=context))[0] assert df.to_dict(orient="records") == [{"customer_id": 1, "zip": "00000"}] @@ -7249,7 +7064,9 @@ def test_macros_python_model(mocker: MockerFixture) -> None: @model( "foo_macro_model_@{bar}", columns={"a": "string"}, - kind=dict(name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, time_column="@{time_col}"), + kind=dict( + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, time_column="@{time_col}" + ), stamp="@{stamp}", cron="@some_cron_var", owner="@IF(@gateway = 'dev', @{dev_owner}, @{prod_owner})", @@ -7284,7 +7101,9 @@ def model_with_macros(context, **kwargs): ) assert python_model.name == "foo_macro_model_suffix" - assert python_model.python_env[c.SQLMESH_VARS] == Executable.value({"test_var_a": "test_value"}) + assert python_model.python_env[c.SQLMESH_VARS] == Executable.value( + {"test_var_a": "test_value"} + ) assert not python_model.enabled assert python_model.start == "2024-01-01" assert python_model.owner == "pr_1" @@ -7361,8 +7180,7 @@ def model_with_macros(evaluator, **kwargs): def test_unrendered_macros_sql_model(mocker: MockerFixture) -> None: model = load_sql_based_model( - parse( - """ + parse(""" MODEL ( name db.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -7388,8 +7206,7 @@ def test_unrendered_macros_sql_model(mocker: MockerFixture) -> None: ); SELECT * FROM src; - """ - ), + """), variables={ "gateway": "dev", "key": "a", # Not included in python_env because kind is rendered at load time @@ -7422,7 +7239,9 @@ def test_unrendered_macros_sql_model(mocker: MockerFixture) -> None: "spark.executor.memory": "1G", "baz": exp.maybe_parse("@session_var"), } - assert model.virtual_properties["creatable_type"] == exp.maybe_parse("@{create_type}") + assert model.virtual_properties["creatable_type"] == exp.maybe_parse( + "@{create_type}" + ) assert ( model.physical_properties["location1"].sql() @@ -7470,7 +7289,9 @@ def model_with_macros(evaluator, **kwargs): exp.convert(evaluator.var("TEST_VAR_A")).as_("a"), ) - python_sql_model = model.get_registry()["test_unrendered_macros_python_model_@{bar}"].model( + python_sql_model = model.get_registry()[ + "test_unrendered_macros_python_model_@{bar}" + ].model( module_path=Path("."), path=Path("."), macros=macro.get_registry(), @@ -7572,7 +7393,11 @@ def test_named_variables_python_model(mocker: MockerFixture) -> None: columns={"a": "string", "b": "string", "c": "string"}, ) def model_with_named_variables( - context, start: TimeLike, test_var_a: str, test_var_b: t.Optional[str] = None, **kwargs + context, + start: TimeLike, + test_var_a: str, + test_var_b: t.Optional[str] = None, + **kwargs, ): return pd.DataFrame( [{"a": test_var_a, "b": test_var_b, "start": start.strftime("%Y-%m-%d")}] # type: ignore @@ -7595,7 +7420,9 @@ def model_with_named_variables( context = ExecutionContext(mocker.Mock(), {}, None, None) df = list(python_model.render(context=context))[0] - assert df.to_dict(orient="records") == [{"a": "test_value", "b": None, "start": to_ds(c.EPOCH)}] + assert df.to_dict(orient="records") == [ + {"a": "test_value", "b": None, "start": to_ds(c.EPOCH)} + ] def test_named_variables_kw_only_python_model(mocker: MockerFixture) -> None: @@ -7617,7 +7444,9 @@ def model_with_named_kw_only_variables( variables={"test_var_a": "test_value"}, ) - assert python_model.python_env[c.SQLMESH_VARS] == Executable.value({"test_var_a": "test_value"}) + assert python_model.python_env[c.SQLMESH_VARS] == Executable.value( + {"test_var_a": "test_value"} + ) context = ExecutionContext(mocker.Mock(), {}, None, None) df = list(python_model.render(context=context))[0] @@ -7626,16 +7455,16 @@ def model_with_named_kw_only_variables( def test_gateway_macro() -> None: model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_gateway_macro); SELECT @gateway AS gateway - """ - ), + """), variables={c.GATEWAY: "in_memory"}, ) - assert model.python_env[c.SQLMESH_VARS] == Executable.value({c.GATEWAY: "in_memory"}) + assert model.python_env[c.SQLMESH_VARS] == Executable.value( + {c.GATEWAY: "in_memory"} + ) assert model.render_query_or_raise().sql() == "SELECT 'in_memory' AS \"gateway\"" @macro() @@ -7643,16 +7472,16 @@ def macro_uses_gateway(evaluator) -> exp.Expr: return exp.convert(evaluator.gateway + "_from_macro") model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_gateway_macro); SELECT @macro_uses_gateway() AS gateway_from_macro - """ - ), + """), variables={c.GATEWAY: "in_memory"}, ) - assert model.python_env[c.SQLMESH_VARS] == Executable.value({c.GATEWAY: "in_memory"}) + assert model.python_env[c.SQLMESH_VARS] == Executable.value( + {c.GATEWAY: "in_memory"} + ) assert ( model.render_query_or_raise().sql() == "SELECT 'in_memory_from_macro' AS \"gateway_from_macro\"" @@ -7661,19 +7490,21 @@ def macro_uses_gateway(evaluator) -> exp.Expr: def test_gateway_macro_jinja() -> None: model = load_sql_based_model( - parse( - """ + parse(""" MODEL(name sushi.test_gateway_macro_jinja); JINJA_QUERY_BEGIN; SELECT '{{ gateway() }}' AS gateway_jinja; JINJA_END; - """ - ), + """), variables={c.GATEWAY: "in_memory"}, ) - assert model.python_env[c.SQLMESH_VARS] == Executable.value({c.GATEWAY: "in_memory"}) - assert model.render_query_or_raise().sql() == "SELECT 'in_memory' AS \"gateway_jinja\"" + assert model.python_env[c.SQLMESH_VARS] == Executable.value( + {c.GATEWAY: "in_memory"} + ) + assert ( + model.render_query_or_raise().sql() == "SELECT 'in_memory' AS \"gateway_jinja\"" + ) def test_gateway_python_model(mocker: MockerFixture) -> None: @@ -7691,7 +7522,9 @@ def model_with_variables(context, **kwargs): variables={c.GATEWAY: "in_memory"}, ) - assert python_model.python_env[c.SQLMESH_VARS] == Executable.value({c.GATEWAY: "in_memory"}) + assert python_model.python_env[c.SQLMESH_VARS] == Executable.value( + {c.GATEWAY: "in_memory"} + ) context = ExecutionContext(mocker.Mock(), {}, None, None) df = list(python_model.render(context=context))[0] @@ -7700,15 +7533,13 @@ def model_with_variables(context, **kwargs): @pytest.mark.parametrize("dialect", ["spark", "trino"]) def test_view_render_no_quote_identifiers(dialect: str) -> None: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind VIEW, ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, dialect=dialect) assert ( model.render_query_or_raise().sql(dialect=dialect) @@ -7726,15 +7557,13 @@ def test_view_render_no_quote_identifiers(dialect: str) -> None: ], ) def test_render_quote_identifiers(dialect: str, kind: str) -> None: - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.table, kind {kind}, ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, dialect=dialect) assert ( model.render_query_or_raise().sql(dialect="duckdb") @@ -7743,8 +7572,7 @@ def test_render_quote_identifiers(dialect: str, kind: str) -> None: def test_this_model() -> None: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name `project-1.table`, dialect bigquery, @@ -7761,8 +7589,7 @@ def test_this_model() -> None: JINJA_STATEMENT_BEGIN; VACUUM {{ this_model }} TO 'b'; JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) assert ( @@ -7815,12 +7642,12 @@ def this_model_resolves_to_quoted_table(evaluator): return not this_model or ( isinstance(this_model, exp.Table) - and this_model.sql(dialect=evaluator.dialect, comments=False) == expected_name + and this_model.sql(dialect=evaluator.dialect, comments=False) + == expected_name and evaluator.this_model == expected_name ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table, dialect snowflake); SELECT @@ -7828,33 +7655,33 @@ def this_model_resolves_to_quoted_table(evaluator): @this_model_resolves_to_quoted_table() AS this_model_resolves_to_quoted_table; CREATE TABLE db.other AS SELECT * FROM @this_model AS x; - """ - ) + """) model = load_sql_based_model(expressions) - expected_post = d.parse('CREATE TABLE "DB"."OTHER" AS SELECT * FROM "DB"."TABLE" AS "X";') + expected_post = d.parse( + 'CREATE TABLE "DB"."OTHER" AS SELECT * FROM "DB"."TABLE" AS "X";' + ) assert model.render_post_statements() == expected_post snapshot = Snapshot.from_node(model, nodes={}) assert ( - model.render_query_or_raise(snapshots={snapshot.name: snapshot}, start="2020-01-01").sql( - dialect="snowflake" - ) + model.render_query_or_raise( + snapshots={snapshot.name: snapshot}, start="2020-01-01" + ).sql(dialect="snowflake") == 'SELECT 1 AS "COL", TRUE AS "THIS_MODEL_RESOLVES_TO_QUOTED_TABLE"' ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) assert ( - model.render_query_or_raise(snapshots={snapshot.name: snapshot}, start="2021-01-01").sql( - dialect="snowflake" - ) + model.render_query_or_raise( + snapshots={snapshot.name: snapshot}, start="2021-01-01" + ).sql(dialect="snowflake") == 'SELECT 1 AS "COL", TRUE AS "THIS_MODEL_RESOLVES_TO_QUOTED_TABLE"' ) def test_macros_in_physical_properties(make_snapshot): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test.test_model, kind FULL, @@ -7871,8 +7698,7 @@ def test_macros_in_physical_properties(make_snapshot): ); SELECT 1; - """ - ) + """) model = load_sql_based_model( expressions, variables={"gateway": "dev"}, default_catalog="unit_test" @@ -7922,8 +7748,7 @@ def session_properties(evaluator, value): value=exp.convert([exp.convert("foo").eq(exp.var(f"bar_{value}"))]), ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name @{gateway}__@{gateway}.test_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -7935,8 +7760,7 @@ def session_properties(evaluator, value): ); SELECT a, b UNION SELECT c, c - """ - ) + """) model = load_sql_based_model( expressions, variables={"gateway": "test_gateway", "time_column": "a"} @@ -7965,8 +7789,7 @@ def not_loaded_macro(evaluator: MacroEvaluator) -> int: def max_value(evaluator: MacroEvaluator) -> int: return 1000 - audit_expression = parse( - """ + audit_expression = parse(""" AUDIT ( name assert_max_value, ); @@ -7974,11 +7797,9 @@ def max_value(evaluator: MacroEvaluator) -> int: FROM @this_model WHERE id > @max_value; - """ - ) + """) - not_zero_audit = parse( - """ + not_zero_audit = parse(""" AUDIT ( name assert_not_zero, ); @@ -7986,11 +7807,9 @@ def max_value(evaluator: MacroEvaluator) -> int: FROM @this_model WHERE id = @zero_value; - """ - ) + """) - model_expression = d.parse( - """ + model_expression = d.parse(""" MODEL ( name db.audit_model, audits (assert_max_value, assert_positive_ids), @@ -8004,8 +7823,7 @@ def max_value(evaluator: MacroEvaluator) -> int: FROM @this_model WHERE id < @min_value; - """ - ) + """) audits = { "assert_max_value": load_audit(audit_expression, dialect="duckdb"), @@ -8042,7 +7860,9 @@ def test_python_model_dialect(): @model( name="a", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="x", format="YYMMDD")), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="x", format="YYMMDD") + ), columns={}, ) def test(context, **kwargs): @@ -8084,7 +7904,9 @@ def a_model(context): # column type not parseable by default dialect and no explicit dialect: error model._dialect = "snowflake" - with pytest.raises(ParseError, match="No expression was parsed from 'DateTime64\\(9\\)'"): + with pytest.raises( + ParseError, match="No expression was parsed from 'DateTime64\\(9\\)'" + ): @model("bad", columns={'"COL"': "DateTime64(9)"}) def a_model(context): @@ -8099,8 +7921,7 @@ def a_model(context): def test_jinja_runtime_stage(assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test.jinja ); @@ -8110,8 +7931,7 @@ def test_jinja_runtime_stage(assert_exp_eq): SELECT '{{ runtime_stage }}' as a, {{ runtime_stage == 'loading' }} as b JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) assert_exp_eq(model.render_query(), '''SELECT 'loading' as "a", TRUE as "b"''') @@ -8122,15 +7942,13 @@ def test_forward_only_on_destructive_change_config() -> None: config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind FULL, ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") @@ -8140,8 +7958,7 @@ def test_forward_only_on_destructive_change_config() -> None: config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -8150,8 +7967,7 @@ def test_forward_only_on_destructive_change_config() -> None: ), ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") @@ -8161,8 +7977,7 @@ def test_forward_only_on_destructive_change_config() -> None: config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -8172,8 +7987,7 @@ def test_forward_only_on_destructive_change_config() -> None: ), ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") @@ -8181,12 +7995,13 @@ def test_forward_only_on_destructive_change_config() -> None: # WARN specified as model default, overrides incremental model sqlmesh default ERROR config = Config( - model_defaults=ModelDefaultsConfig(dialect="duckdb", on_destructive_change="warn") + model_defaults=ModelDefaultsConfig( + dialect="duckdb", on_destructive_change="warn" + ) ) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -8195,8 +8010,7 @@ def test_forward_only_on_destructive_change_config() -> None: ), ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") @@ -8204,19 +8018,19 @@ def test_forward_only_on_destructive_change_config() -> None: # WARN specified as model default, does not override non-incremental sqlmesh default ALLOW config = Config( - model_defaults=ModelDefaultsConfig(dialect="duckdb", on_destructive_change="warn") + model_defaults=ModelDefaultsConfig( + dialect="duckdb", on_destructive_change="warn" + ) ) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind FULL, ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") @@ -8228,8 +8042,7 @@ def test_batch_concurrency_config() -> None: config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -8237,19 +8050,19 @@ def test_batch_concurrency_config() -> None: ), ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") assert context_model.batch_concurrency is None # batch_concurrency specified in model defaults applies to incremental models - config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5)) + config = Config( + model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5) + ) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -8257,19 +8070,19 @@ def test_batch_concurrency_config() -> None: ), ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") assert context_model.batch_concurrency == 5 # batch_concurrency specified in model definition overrides default - config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5)) + config = Config( + model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5) + ) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -8278,37 +8091,37 @@ def test_batch_concurrency_config() -> None: ), ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") assert context_model.batch_concurrency == 10 # batch_concurrency default does not apply to non-incremental models - config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5)) + config = Config( + model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5) + ) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind FULL, ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") assert context_model.batch_concurrency is None # batch_concurrency default does not apply to INCREMENTAL_BY_UNIQUE_KEY models - config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5)) + config = Config( + model_defaults=ModelDefaultsConfig(dialect="duckdb", batch_concurrency=5) + ) context = Context(config=config) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name memory.db.table, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -8316,8 +8129,7 @@ def test_batch_concurrency_config() -> None: ), ); SELECT a, b, c FROM source_table; - """ - ) + """) model = load_sql_based_model(expressions, defaults=config.model_defaults.dict()) context.upsert_model(model) context_model = context.get_model("memory.db.table") @@ -8326,7 +8138,8 @@ def test_batch_concurrency_config() -> None: def test_model_meta_on_additive_change_property() -> None: """Test that ModelMeta has on_additive_change property that works like on_destructive_change.""" - from sqlmesh.core.model.kind import IncrementalByTimeRangeKind, OnAdditiveChange + from sqlmesh.core.model.kind import (IncrementalByTimeRangeKind, + OnAdditiveChange) from sqlmesh.core.model.meta import ModelMeta # Test incremental model with on_additive_change=ERROR @@ -8365,8 +8178,7 @@ def test_model_meta_on_additive_change_property() -> None: def test_incremental_by_partition(sushi_context, assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_PARTITION, @@ -8374,14 +8186,12 @@ def test_incremental_by_partition(sushi_context, assert_exp_eq): ); SELECT a, b - """ - ) + """) model = load_sql_based_model(expressions) assert model.kind.is_incremental_by_partition assert not model.kind.disable_restatement - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_PARTITION ( @@ -8391,8 +8201,7 @@ def test_incremental_by_partition(sushi_context, assert_exp_eq): ); SELECT a, b - """ - ) + """) model = load_sql_based_model(expressions) assert model.kind.is_incremental_by_partition assert not model.kind.disable_restatement @@ -8401,24 +8210,21 @@ def test_incremental_by_partition(sushi_context, assert_exp_eq): ValidationError, match=r".*partitioned_by field is required for INCREMENTAL_BY_PARTITION models.*", ): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_PARTITION, ); SELECT a, b - """ - ) + """) load_sql_based_model(expressions) with pytest.raises( ConfigError, match=r".*Do not specify the `forward_only` configuration key.*", ): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_PARTITION ( @@ -8427,8 +8233,7 @@ def test_incremental_by_partition(sushi_context, assert_exp_eq): ); SELECT a, b - """ - ) + """) load_sql_based_model(expressions) @@ -8485,7 +8290,9 @@ def test_model_table_name_inference( ], ], ) -def test_python_model_name_inference(tmp_path: Path, path: str, expected_name: str) -> None: +def test_python_model_name_inference( + tmp_path: Path, path: str, expected_name: str +) -> None: init_example_project(tmp_path, engine_type="duckdb") config = Config( model_defaults=ModelDefaultsConfig(dialect="duckdb"), @@ -8527,13 +8334,16 @@ def my_model(context, **kwargs): path_b.write_text(model_payload) context = Context(paths=tmp_path, config=config) - assert context.get_model("test_schema.test_model_a").name == "test_schema.test_model_a" - assert context.get_model("test_schema.test_model_b").name == "test_schema.test_model_b" + assert ( + context.get_model("test_schema.test_model_a").name == "test_schema.test_model_a" + ) + assert ( + context.get_model("test_schema.test_model_b").name == "test_schema.test_model_b" + ) def test_custom_kind(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind CUSTOM ( @@ -8553,11 +8363,11 @@ def test_custom_kind(): ); SELECT a, b - """ - ) + """) with pytest.raises( - ConfigError, match=r"Materialization strategy with name 'MyTestStrategy' was not found.*" + ConfigError, + match=r"Materialization strategy with name 'MyTestStrategy' was not found.*", ): model = load_sql_based_model(expressions) model.validate_definition() @@ -8582,9 +8392,7 @@ class MyTestStrategy(CustomMaterialization): assert kind.batch_concurrency == 2 assert kind.lookback == 3 - assert ( - kind.to_expression().sql() - == """CUSTOM ( + assert kind.to_expression().sql() == """CUSTOM ( materialization 'MyTestStrategy', materialization_properties ('key_a' = 'value_a', key_b = 2, 'key_c' = TRUE, 'key_d' = 1.23), forward_only TRUE, @@ -8593,7 +8401,6 @@ class MyTestStrategy(CustomMaterialization): batch_concurrency 2, lookback 3 )""" - ) def test_custom_kind_lookback_property(): @@ -8607,8 +8414,7 @@ def test_custom_kind_lookback_property(): class MyTestStrategy(CustomMaterialization): pass - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.custom_table, kind CUSTOM ( @@ -8617,8 +8423,7 @@ class MyTestStrategy(CustomMaterialization): ) ); SELECT a, b FROM upstream - """ - ) + """) model = load_sql_based_model(expressions) assert model.kind.is_custom @@ -8629,11 +8434,12 @@ class MyTestStrategy(CustomMaterialization): # The bug: model.lookback should return 3, but with the old implementation # using isinstance(self.kind, _IncrementalBy), it would return 0 - assert model.lookback == 3, "CustomKind lookback not accessible via model.lookback property" + assert ( + model.lookback == 3 + ), "CustomKind lookback not accessible via model.lookback property" # Test 2: CustomKind without lookback (should default to 0) - expressions_no_lookback = d.parse( - """ + expressions_no_lookback = d.parse(""" MODEL ( name db.custom_table_no_lookback, kind CUSTOM ( @@ -8641,15 +8447,13 @@ class MyTestStrategy(CustomMaterialization): ) ); SELECT a, b FROM upstream - """ - ) + """) model_no_lookback = load_sql_based_model(expressions_no_lookback) assert model_no_lookback.lookback == 0 # Test 3: Ensure IncrementalByTimeRangeKind still works correctly - incremental_expressions = d.parse( - """ + incremental_expressions = d.parse(""" MODEL ( name db.incremental_table, kind INCREMENTAL_BY_TIME_RANGE ( @@ -8658,8 +8462,7 @@ class MyTestStrategy(CustomMaterialization): ) ); SELECT ds, a, b FROM upstream - """ - ) + """) incremental_model = load_sql_based_model(incremental_expressions) assert incremental_model.lookback == 5 @@ -8682,8 +8485,7 @@ def time_column(self): class TimeColumnMaterialization(CustomMaterialization[TimeColumnCustomKind]): NAME = "time_column_custom_strategy" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind CUSTOM ( @@ -8696,8 +8498,7 @@ class TimeColumnMaterialization(CustomMaterialization[TimeColumnCustomKind]): ); SELECT a, b, '2020-01-01' as ts - """ - ) + """) model = load_sql_based_model(expressions, time_column_format="%d-%m-%Y") assert isinstance(model.kind, TimeColumnCustomKind) @@ -8709,8 +8510,7 @@ class TimeColumnMaterialization(CustomMaterialization[TimeColumnCustomKind]): ) # dialect should not be serialized against the kind # explicit time_column format within the model - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind CUSTOM ( @@ -8722,8 +8522,7 @@ class TimeColumnMaterialization(CustomMaterialization[TimeColumnCustomKind]): ); SELECT a, b, '2020-01-01' as ts - """ - ) + """) model = load_sql_based_model(expressions, time_column_format="%d-%m-%Y") assert model.time_column.format == "%Y-%m-%d" @@ -8731,9 +8530,7 @@ class TimeColumnMaterialization(CustomMaterialization[TimeColumnCustomKind]): def test_model_kind_to_expression(): assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -8741,11 +8538,7 @@ def test_model_kind_to_expression(): ), ); SELECT a, b - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """INCREMENTAL_BY_TIME_RANGE ( time_column ("a", '%Y-%m-%d'), partition_by_time_column TRUE, @@ -8757,9 +8550,7 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -8773,11 +8564,7 @@ def test_model_kind_to_expression(): ), ); SELECT a, b - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """INCREMENTAL_BY_TIME_RANGE ( time_column ("a", '%Y-%m-%d'), partition_by_time_column TRUE, @@ -8792,9 +8579,7 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_UNIQUE_KEY( @@ -8802,11 +8587,7 @@ def test_model_kind_to_expression(): ), ); SELECT a, b - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """INCREMENTAL_BY_UNIQUE_KEY ( unique_key ("a"), batch_concurrency 1, @@ -8818,9 +8599,7 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_UNIQUE_KEY( @@ -8829,11 +8608,7 @@ def test_model_kind_to_expression(): ), ); SELECT a, b - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """INCREMENTAL_BY_UNIQUE_KEY ( unique_key ("a"), when_matched (WHEN MATCHED THEN UPDATE SET "__MERGE_TARGET__"."b" = COALESCE("__MERGE_SOURCE__"."b", "__MERGE_TARGET__"."b")), @@ -8846,9 +8621,7 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_UNIQUE_KEY( @@ -8858,11 +8631,7 @@ def test_model_kind_to_expression(): ), ); SELECT a, b - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """INCREMENTAL_BY_UNIQUE_KEY ( unique_key ("a"), when_matched (WHEN MATCHED AND "__MERGE_SOURCE__"."x" = 1 THEN UPDATE SET "__MERGE_TARGET__"."b" = COALESCE("__MERGE_SOURCE__"."b", "__MERGE_TARGET__"."b") WHEN MATCHED THEN UPDATE SET "__MERGE_TARGET__"."b" = COALESCE("__MERGE_SOURCE__"."b", "__MERGE_TARGET__"."b")), @@ -8875,20 +8644,14 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_PARTITION, partitioned_by ["a"], ); SELECT a, b - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """INCREMENTAL_BY_PARTITION ( forward_only TRUE, disable_restatement FALSE, @@ -8899,16 +8662,14 @@ def test_model_kind_to_expression(): assert ( load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL ( name db.seed, kind SEED ( path '../seeds/waiter_names.csv', ) ); - """ - ), + """), path=Path("./examples/sushi/models/test_model.sql"), ) .kind.to_expression() @@ -8920,9 +8681,7 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind SCD_TYPE_2_BY_TIME ( @@ -8930,11 +8689,7 @@ def test_model_kind_to_expression(): ) ); SELECT a, b - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """SCD_TYPE_2_BY_TIME ( updated_at_name "updated_at", updated_at_as_valid_from FALSE, @@ -8951,9 +8706,7 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind SCD_TYPE_2_BY_COLUMN ( @@ -8962,11 +8715,7 @@ def test_model_kind_to_expression(): ) ); SELECT a, b, c - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """SCD_TYPE_2_BY_COLUMN ( columns ("b"), execution_time_as_valid_from FALSE, @@ -8983,9 +8732,7 @@ def test_model_kind_to_expression(): ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind SCD_TYPE_2_BY_COLUMN ( @@ -8994,11 +8741,7 @@ def test_model_kind_to_expression(): ) ); SELECT a, b, c - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """SCD_TYPE_2_BY_COLUMN ( columns (*), execution_time_as_valid_from FALSE, @@ -9014,56 +8757,35 @@ def test_model_kind_to_expression(): )""" ) - assert ( - load_sql_based_model( - d.parse( - """ + assert load_sql_based_model(d.parse(""" MODEL ( name db.table, kind FULL ); SELECT a, b, c - """ - ) - ) - .kind.to_expression() - .sql() - == "FULL" - ) + """)).kind.to_expression().sql() == "FULL" assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind VIEW - ); - SELECT a, b, c - """ - ) - ) - .kind.to_expression() - .sql() + ); + SELECT a, b, c + """)).kind.to_expression().sql() == """VIEW ( materialized FALSE )""" ) assert ( - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind VIEW (materialized true) ); SELECT a, b, c - """ - ) - ) - .kind.to_expression() - .sql() + """)).kind.to_expression().sql() == """VIEW ( materialized TRUE )""" @@ -9075,20 +8797,17 @@ def test_bad_model_kind(): SQLMeshError, match=f"Model kind specified as 'BAD_KIND', but that is not a valid model kind.\n\nPlease specify one of {', '.join(ModelKindName)}.", ): - d.parse( - """ + d.parse(""" MODEL ( name db.table, kind BAD_KIND ); SELECT a, b - """ - ) + """) def test_merge_filter(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -9097,8 +8816,7 @@ def test_merge_filter(): ) ); SELECT 'name' AS name, 1 AS salary; - """ - ) + """) expected_incremental_predicate = f"`{MERGE_SOURCE_ALIAS}`.`salary` > 0" @@ -9109,8 +8827,7 @@ def test_merge_filter(): assert model.kind.merge_filter.sql(dialect="hive") == expected_incremental_predicate assert model.dialect == "hive" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.test, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -9132,11 +8849,14 @@ def test_merge_filter(): purchase_order_id, start_date FROM db.upstream - """ - ) + """) - model = SqlModel.parse_raw(load_sql_based_model(expressions, dialect="duckdb").json()) - assert d.format_model_expressions(model.render_definition(), dialect=model.dialect) == ( + model = SqlModel.parse_raw( + load_sql_based_model(expressions, dialect="duckdb").json() + ) + assert d.format_model_expressions( + model.render_definition(), dialect=model.dialect + ) == ( f"""MODEL ( name db.test, dialect duckdb, @@ -9171,7 +8891,9 @@ def test_merge_filter(): FROM db.upstream""" ) - rendered_merge_filters = model.render_merge_filter(start="2023-01-01", end="2023-01-02") + rendered_merge_filters = model.render_merge_filter( + start="2023-01-01", end="2023-01-02" + ) assert ( rendered_merge_filters.sql(dialect="hive") == "(`__MERGE_SOURCE__`.`ds` > (SELECT MAX(`ds`) FROM `db`.`test`) AND `__MERGE_SOURCE__`.`ds` > '2023-01-01' AND `__MERGE_SOURCE__`.`_operation` <> 1 AND `__MERGE_TARGET__`.`start_date` > CURRENT_DATE + INTERVAL '7' DAY)" @@ -9180,8 +8902,7 @@ def test_merge_filter(): def test_merge_filter_normalization(): # unquoted gets normalized and quoted - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -9190,15 +8911,15 @@ def test_merge_filter_normalization(): ) ); SELECT 'name' AS name, 1 AS salary; - """ - ) + """) model = load_sql_based_model(expressions, dialect="snowflake") - assert model.merge_filter.sql(dialect="snowflake") == '"__MERGE_SOURCE__"."SALARY" > 0' + assert ( + model.merge_filter.sql(dialect="snowflake") == '"__MERGE_SOURCE__"."SALARY" > 0' + ) # quoted gets preserved - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.employees, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -9207,11 +8928,12 @@ def test_merge_filter_normalization(): ) ); SELECT 'name' AS name, 1 AS "SaLArY"; - """ - ) + """) model = load_sql_based_model(expressions, dialect="snowflake") - assert model.merge_filter.sql(dialect="snowflake") == '"__MERGE_SOURCE__"."SaLArY" > 0' + assert ( + model.merge_filter.sql(dialect="snowflake") == '"__MERGE_SOURCE__"."SaLArY" > 0' + ) def test_merge_filter_macro(): @@ -9220,10 +8942,11 @@ def predicate( evaluator: MacroEvaluator, cluster_column: exp.Column, ) -> exp.Expr: - return parse_one(f"source.{cluster_column} > dateadd(day, -7, target.{cluster_column})") + return parse_one( + f"source.{cluster_column} > dateadd(day, -7, target.{cluster_column})" + ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.incremental_model, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -9233,8 +8956,7 @@ def predicate( clustered_by update_datetime ); SELECT id, update_datetime FROM db.test_model; - """ - ) + """) unrendered_merge_filter = f"""@predicate("UPDATE_DATETIME") AND "{MERGE_TARGET_ALIAS}"."UPDATE_DATETIME" > @start_dt""" expected_merge_filter = ( @@ -9263,19 +8985,18 @@ def test_macro_func_hash(mocker: MockerFixture, metadata_only: bool): def noop(evaluator) -> None: return None - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.model, ); SELECT 1; - """ + """) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") ) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.model, ); @@ -9283,8 +9004,7 @@ def noop(evaluator) -> None: SELECT 1; @noop(); - """ - ) + """) new_model = load_sql_based_model( expressions, path=Path("./examples/sushi/models/test_model.sql") ) @@ -9314,8 +9034,7 @@ def noop(evaluator) -> None: def test_managed_kind_sql(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind MANAGED, @@ -9327,17 +9046,16 @@ def test_managed_kind_sql(): ); SELECT a, b - """ - ) + """) model = load_sql_based_model(expressions) assert model.kind.is_managed - with pytest.raises(ConfigError, match=r".*must specify the 'target_lag' physical property.*"): - load_sql_based_model( - d.parse( - """ + with pytest.raises( + ConfigError, match=r".*must specify the 'target_lag' physical property.*" + ): + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind MANAGED, @@ -9345,9 +9063,7 @@ def test_managed_kind_sql(): ); SELECT a, b - """ - ) - ).validate_definition() + """)).validate_definition() def test_managed_kind_python(): @@ -9372,8 +9088,7 @@ def execute( def test_physical_version(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -9384,8 +9099,7 @@ def test_physical_version(): ); SELECT a, b - """ - ) + """) model = load_sql_based_model(expressions) assert model.physical_version == "1234" @@ -9394,9 +9108,7 @@ def test_physical_version(): ConfigError, match=r"Pinning a physical version is only supported for forward only models( at.*)?", ): - load_sql_based_model( - d.parse( - """ + load_sql_based_model(d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -9406,28 +9118,23 @@ def test_physical_version(): ); SELECT a, b - """ - ) - ).validate_definition() + """)).validate_definition() def test_trailing_comments(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); /* some comment A */ SELECT 1; /* some comment B */ - """ - ) + """) model = load_sql_based_model(expressions) assert not model.render_pre_statements() assert not model.render_post_statements() - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -9437,16 +9144,16 @@ def test_trailing_comments(): ); /* some comment A */ - """ + """) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") ) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) assert not model.render_pre_statements() assert not model.render_post_statements() def test_comments_in_jinja_query(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); JINJA_QUERY_BEGIN; @@ -9456,13 +9163,11 @@ def test_comments_in_jinja_query(): /* some comment B */ JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) assert model.render_query().sql() == '/* some comment A */ SELECT 1 AS "1"' - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); JINJA_QUERY_BEGIN; @@ -9473,23 +9178,20 @@ def test_comments_in_jinja_query(): /* some comment B */ JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) with pytest.raises(ConfigError, match=r"Too many statements in query.*"): model.render_query() def test_jinja_render_parse_error(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.test_model); JINJA_QUERY_BEGIN; {{ unknown_macro() }} JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) @@ -9505,15 +9207,13 @@ def test_jinja_render_debug_logging(caplog): caplog.set_level(logging.DEBUG, logger="sqlmesh.core.renderer") # Create a model with unparseable Jinja that will be rendered - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.test_model); JINJA_QUERY_BEGIN; {{ 'SELECT invalid syntax here!' }} JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) @@ -9530,27 +9230,26 @@ def test_jinja_render_debug_logging(caplog): def test_staged_file_path(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name test, dialect snowflake); SELECT * FROM @a.b/c/d.csv(FILE_FORMAT => 'b.ff') - """ - ) + """) model = load_sql_based_model(expressions) query = model.render_query() - assert query.sql(dialect="snowflake") == "SELECT * FROM @a.b/c/d.csv (FILE_FORMAT => 'b.ff')" + assert ( + query.sql(dialect="snowflake") + == "SELECT * FROM @a.b/c/d.csv (FILE_FORMAT => 'b.ff')" + ) - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name test, dialect snowflake); SELECT * FROM @variable (FILE_FORMAT => 'foo'), @non_variable (FILE_FORMAT => 'bar') LIMIT 100 - """ - ) + """) model = load_sql_based_model(expressions, variables={"variable": "some_path"}) query = model.render_query() assert ( @@ -9560,15 +9259,13 @@ def test_staged_file_path(): def test_cache(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name test); SELECT 1 x FROM y - """ - ) + """) model = load_sql_based_model(expressions) assert model.depends_on == {'"y"'} assert model.copy(update={"depends_on_": {'"z"'}}).depends_on == {'"z"', '"y"'} @@ -9627,15 +9324,13 @@ def resolve_parent(evaluator, name): "JINJA_STATEMENT_BEGIN; {{ resolve_table('parent') }}; JINJA_END;", "@resolve_parent('parent')", ): - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL (name child); SELECT c FROM parent; {post_statement} - """ - ) + """) child = load_sql_based_model(expressions) parent = load_sql_based_model(d.parse("MODEL (name parent); SELECT 1 AS c")) @@ -9643,7 +9338,9 @@ def resolve_parent(evaluator, name): parent_snapshot.categorize_as(SnapshotChangeCategory.BREAKING) version = parent_snapshot.version - post_statements = child.render_post_statements(snapshots={'"parent"': parent_snapshot}) + post_statements = child.render_post_statements( + snapshots={'"parent"': parent_snapshot} + ) assert len(post_statements) == 1 assert post_statements[0].sql() == f'"sqlmesh__default"."parent__{version}"' @@ -9653,15 +9350,13 @@ def resolve_parent(evaluator, name): "JINJA_STATEMENT_BEGIN; {{ resolve_table('schema.parent') }}; JINJA_END;", "@resolve_parent('schema.parent')", ): - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL (name schema.child); SELECT c FROM schema.parent; {post_statement} - """ - ) + """) child = load_sql_based_model(expressions, default_catalog="main") parent = load_sql_based_model( d.parse("MODEL (name schema.parent); SELECT 1 AS c"), default_catalog="main" @@ -9676,12 +9371,14 @@ def resolve_parent(evaluator, name): ) assert len(post_statements) == 1 - assert post_statements[0].sql() == f'"main"."sqlmesh__schema"."schema__parent__{version}"' + assert ( + post_statements[0].sql() + == f'"main"."sqlmesh__schema"."schema__parent__{version}"' + ) def test_cluster_with_complex_expression(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test, dialect snowflake, @@ -9692,16 +9389,16 @@ def test_cluster_with_complex_expression(): SELECT 1 AS c, CAST('2020-01-01 12:05:03' AS TIMESTAMPTZ) AS cluster_col - """ - ) + """) model = load_sql_based_model(expressions) - assert [expr.sql("snowflake") for expr in model.clustered_by] == ['(TO_DATE("CLUSTER_COL"))'] + assert [expr.sql("snowflake") for expr in model.clustered_by] == [ + '(TO_DATE("CLUSTER_COL"))' + ] def test_parametric_model_kind(): - parsed_definition = d.parse( - """ + parsed_definition = d.parse(""" MODEL ( name db.test_schema.test_model, kind @IF(@gateway = 'main', VIEW, FULL) @@ -9709,8 +9406,7 @@ def test_parametric_model_kind(): SELECT 1 AS c - """ - ) + """) model = load_sql_based_model(parsed_definition, variables={c.GATEWAY: "main"}) assert isinstance(model.kind, ViewKind) @@ -9724,8 +9420,7 @@ def test_fingerprint_signals(): def test_signal_hash(batch): return True - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, signals [ @@ -9733,8 +9428,7 @@ def test_signal_hash(batch): ], ); SELECT 1; - """ - ) + """) model = load_sql_based_model(expressions, signal_definitions=signal.get_registry()) metadata_hash = model.metadata_hash @@ -9747,7 +9441,9 @@ def assert_metadata_only(): assert model.data_hash == data_hash executable = model.python_env["test_signal_hash"] - model.python_env["test_signal_hash"].payload = executable.payload.replace("True", "False") + model.python_env["test_signal_hash"].payload = executable.payload.replace( + "True", "False" + ) assert_metadata_only() model = load_sql_based_model(expressions, signal_definitions=signal.get_registry()) @@ -9764,59 +9460,57 @@ def test_model_optimize(tmp_path: Path, assert_exp_eq): optimized_sql = 'SELECT 3 AS "new_col"' # Model flag is False, overriding defaults - disabled_opt = d.parse( - """ + disabled_opt = d.parse(""" MODEL ( name test, optimize_query False, ); SELECT 1 + 2 AS new_col - """ - ) + """) for default in defaults: model = load_sql_based_model(disabled_opt, defaults=default) assert_exp_eq(model.render_query(), non_optimized_sql) # Model flag is True, overriding defaults - enabled_opt = d.parse( - """ + enabled_opt = d.parse(""" MODEL ( name test, optimize_query True, ); SELECT 1 + 2 AS new_col - """ - ) + """) for default in defaults: model = load_sql_based_model(enabled_opt, defaults=default) assert_exp_eq(model.render_query(), optimized_sql) # Model flag is not defined, behavior is set according to the defaults - none_opt = d.parse( - """ + none_opt = d.parse(""" MODEL ( name test, ); SELECT 1 + 2 AS new_col - """ - ) + """) assert_exp_eq(load_sql_based_model(none_opt).render_query(), optimized_sql) assert_exp_eq( - load_sql_based_model(none_opt, defaults=defaults[0]).render_query(), optimized_sql + load_sql_based_model(none_opt, defaults=defaults[0]).render_query(), + optimized_sql, ) assert_exp_eq( - load_sql_based_model(none_opt, defaults=defaults[1]).render_query(), non_optimized_sql + load_sql_based_model(none_opt, defaults=defaults[1]).render_query(), + non_optimized_sql, ) # Ensure that plan works as expected (optimize_query flag affects the model's data hash) for parsed_model in [enabled_opt, disabled_opt, none_opt]: - context = Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))) + context = Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ) context.upsert_model(load_sql_based_model(parsed_model)) context.plan(auto_apply=True, no_prompts=True) @@ -9824,14 +9518,16 @@ def test_model_optimize(tmp_path: Path, assert_exp_eq): seed_path = tmp_path / "seed.csv" model_kind = SeedKind(path=str(seed_path.absolute())) with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_c 1,text_a,1.0 -2,text_b,2.0""" - ) - model = create_seed_model("test_db.test_seed_model", model_kind, optimize_query=True) - context = Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))) +2,text_b,2.0""") + model = create_seed_model( + "test_db.test_seed_model", model_kind, optimize_query=True + ) + context = Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ) with pytest.raises( ConfigError, @@ -9840,7 +9536,9 @@ def test_model_optimize(tmp_path: Path, assert_exp_eq): context.upsert_model(model) context.plan(auto_apply=True, no_prompts=True) - model = create_seed_model("test_db.test_seed_model", model_kind, optimize_query=False) + model = create_seed_model( + "test_db.test_seed_model", model_kind, optimize_query=False + ) context.upsert_model(model) context.plan(auto_apply=True, no_prompts=True) @@ -9849,8 +9547,7 @@ def test_column_description_metadata_change(): context = Context(config=Config()) model = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL ( name db.test_model, kind full @@ -9858,8 +9555,7 @@ def test_column_description_metadata_change(): SELECT 1 AS id /* description */ - """ - ), + """), default_catalog=context.default_catalog, ) @@ -9880,8 +9576,7 @@ def test_column_description_metadata_change(): def test_auto_restatement(): - parsed_definition = d.parse( - """ + parsed_definition = d.parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE( @@ -9890,13 +9585,10 @@ def test_auto_restatement(): ) ); SELECT 1 AS c - """ - ) + """) model = load_sql_based_model(parsed_definition) assert model.auto_restatement_cron == "@daily" - assert ( - model.kind.to_expression().sql(pretty=True) - == """INCREMENTAL_BY_TIME_RANGE ( + assert model.kind.to_expression().sql(pretty=True) == """INCREMENTAL_BY_TIME_RANGE ( time_column ("a", '%Y-%m-%d'), partition_by_time_column TRUE, forward_only FALSE, @@ -9905,10 +9597,8 @@ def test_auto_restatement(): on_additive_change 'ALLOW', auto_restatement_cron '@daily' )""" - ) - parsed_definition = d.parse( - """ + parsed_definition = d.parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE( @@ -9918,14 +9608,11 @@ def test_auto_restatement(): ) ); SELECT 1 AS c - """ - ) + """) model = load_sql_based_model(parsed_definition) assert model.auto_restatement_cron == "@daily" assert model.auto_restatement_intervals == 1 - assert ( - model.kind.to_expression().sql(pretty=True) - == """INCREMENTAL_BY_TIME_RANGE ( + assert model.kind.to_expression().sql(pretty=True) == """INCREMENTAL_BY_TIME_RANGE ( time_column ("a", '%Y-%m-%d'), partition_by_time_column TRUE, auto_restatement_intervals 1, @@ -9935,10 +9622,8 @@ def test_auto_restatement(): on_additive_change 'ALLOW', auto_restatement_cron '@daily' )""" - ) - parsed_definition = d.parse( - """ + parsed_definition = d.parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE( @@ -9947,8 +9632,7 @@ def test_auto_restatement(): ) ); SELECT 1 AS c - """ - ) + """) with pytest.raises(ValueError, match="Invalid cron expression '@invalid'.*"): load_sql_based_model(parsed_definition) @@ -9976,7 +9660,9 @@ def test_gateway_specific_render(assert_exp_eq) -> None: def dummy_model_entry(evaluator: MacroEvaluator) -> exp.Select: return exp.select("x").from_(exp.values([("1", 2)], "_v", ["x"])) - dummy_model = model.get_registry()["dummy_model"].model(module_path=Path("."), path=Path(".")) + dummy_model = model.get_registry()["dummy_model"].model( + module_path=Path("."), path=Path(".") + ) context.upsert_model(dummy_model) assert isinstance(dummy_model, SqlModel) assert dummy_model.gateway == "duckdb" @@ -10057,7 +9743,9 @@ def resolve_parent_name(evaluator, name): model_snapshot = make_snapshot(model) model_snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - assert model.on_virtual_update == d.parse(virtual_update_statements, default_dialect=dialect) + assert model.on_virtual_update == d.parse( + virtual_update_statements, default_dialect=dialect + ) assert parent.on_virtual_update == d.parse( "JINJA_STATEMENT_BEGIN; GRANT SELECT ON VIEW {{this_model}} TO ROLE admin; JINJA_END;", @@ -10145,7 +9833,10 @@ def model_with_virtual_statements(context, **kwargs): ) python_model = model.get_registry()["db.test_model"].model( - module_path=Path("."), path=Path("."), dialect="duckdb", jinja_macros=jinja_macros + module_path=Path("."), + path=Path("."), + dialect="duckdb", + jinja_macros=jinja_macros, ) assert len(jinja_macros.root_macros) == 1 @@ -10154,7 +9845,8 @@ def model_with_virtual_statements(context, **kwargs): assert len(python_model.on_virtual_update) == 3 rendered_statements = python_model._render_statements( - python_model.on_virtual_update, table_mapping={'"db"."test_model"': "db.test_model"} + python_model.on_virtual_update, + table_mapping={'"db"."test_model"': "db.test_model"}, ) assert ( @@ -10176,7 +9868,8 @@ def test_compile_time_checks(tmp_path: Path): config=Config( model_defaults=ModelDefaultsConfig(dialect="duckdb"), linter=LinterConfig( - enabled=True, rules=["ambiguousorinvalidcolumn", "invalidselectstarexpansion"] + enabled=True, + rules=["ambiguousorinvalidcolumn", "invalidselectstarexpansion"], ), ), paths=tmp_path, @@ -10185,30 +9878,26 @@ def test_compile_time_checks(tmp_path: Path): cfg_err = "Linter detected errors in the code. Please fix them before proceeding." # Strict SELECT * expansion - strict_query = d.parse( - """ + strict_query = d.parse(""" MODEL ( name test, ); SELECT * FROM tbl - """ - ) + """) with pytest.raises(LinterError, match=cfg_err): ctx.upsert_model(load_sql_based_model(strict_query)) ctx.plan_builder("dev") # Strict column resolution - strict_query = d.parse( - """ + strict_query = d.parse(""" MODEL ( name test, ); SELECT foo - """ - ) + """) with pytest.raises(LinterError, match=cfg_err): ctx.upsert_model(load_sql_based_model(strict_query)) @@ -10216,8 +9905,7 @@ def test_compile_time_checks(tmp_path: Path): def test_partition_interval_unit(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test, kind INCREMENTAL_BY_TIME_RANGE( @@ -10226,14 +9914,12 @@ def test_partition_interval_unit(): cron '0 0 1 * *' ); SELECT '2024-01-01' AS ds; - """ - ) + """) model = load_sql_based_model(expressions) assert model.partition_interval_unit == IntervalUnit.MONTH # Partitioning was explicitly set by the user - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test, kind INCREMENTAL_BY_TIME_RANGE( @@ -10243,8 +9929,7 @@ def test_partition_interval_unit(): partitioned_by (ds) ); SELECT '2024-01-01' AS ds; - """ - ) + """) model = load_sql_based_model(expressions) assert model.partition_interval_unit is None @@ -10266,18 +9951,15 @@ def test_model_blueprinting(tmp_path: Path) -> None: identity_macro = tmp_path / "macros" / "identity_macro.py" identity_macro.parent.mkdir(parents=True, exist_ok=True) - identity_macro.write_text( - """from sqlmesh import macro + identity_macro.write_text("""from sqlmesh import macro @macro() def identity(evaluator, value): return value -""" - ) +""") blueprint_sql = tmp_path / "models" / "blueprint.sql" blueprint_sql.parent.mkdir(parents=True, exist_ok=True) - blueprint_sql.write_text( - """ + blueprint_sql.write_text(""" MODEL ( name @{blueprint}.test_model_sql, gateway @identity(@blueprint), @@ -10287,12 +9969,10 @@ def identity(evaluator, value): SELECT @x AS x - """ - ) + """) blueprint_pydf = tmp_path / "models" / "blueprint_df.py" blueprint_pydf.parent.mkdir(parents=True, exist_ok=True) - blueprint_pydf.write_text( - """ + blueprint_pydf.write_text(""" import pandas as pd # noqa: TID253 from sqlmesh import model @@ -10307,12 +9987,10 @@ def identity(evaluator, value): def entrypoint(context, *args, **kwargs): x_var = context.var("x") assert context.blueprint_var("blueprint").startswith("gw") - return pd.DataFrame({"x": [x_var]})""" - ) + return pd.DataFrame({"x": [x_var]})""") blueprint_pysql = tmp_path / "models" / "blueprint_sql.py" blueprint_pysql.parent.mkdir(parents=True, exist_ok=True) - blueprint_pysql.write_text( - """ + blueprint_pysql.write_text(""" from sqlmesh import model @@ -10326,8 +10004,7 @@ def entrypoint(context, *args, **kwargs): def entrypoint(evaluator): x_var = evaluator.var("x") assert evaluator.blueprint_var("blueprint", default="").startswith("gw") - return f'SELECT {x_var} AS x'""" - ) + return f'SELECT {x_var} AS x'""") context = Context(paths=tmp_path, config=config) models = context.models @@ -10356,12 +10033,15 @@ def entrypoint(evaluator): {"blueprint": blueprint_value} ) - assert context.fetchdf(f"from {model.fqn}").to_dict() == {"x": {0: gateway_no}} + assert context.fetchdf(f"from {model.fqn}").to_dict() == { + "x": {0: gateway_no} + } - multi_variable_blueprint_example = tmp_path / "models" / "multi_variable_blueprint_example.sql" + multi_variable_blueprint_example = ( + tmp_path / "models" / "multi_variable_blueprint_example.sql" + ) multi_variable_blueprint_example.parent.mkdir(parents=True, exist_ok=True) - multi_variable_blueprint_example.write_text( - """ + multi_variable_blueprint_example.write_text(""" MODEL ( name @{customer}.my_table, blueprints ( @@ -10376,8 +10056,7 @@ def entrypoint(evaluator): @{customer_field} AS foo2, @BLUEPRINT_VAR('customer_field') AS foo3, FROM @{customer}.my_source - """ - ) + """) context = Context(paths=tmp_path, config=config) models = context.models @@ -10417,8 +10096,7 @@ def test_dynamic_blueprinting_using_custom_macro(tmp_path: Path) -> None: dynamic_template_sql = tmp_path / "models/dynamic_template_custom_macro.sql" dynamic_template_sql.parent.mkdir(parents=True, exist_ok=True) - dynamic_template_sql.write_text( - """ + dynamic_template_sql.write_text(""" MODEL ( name @customer.some_table, kind FULL, @@ -10430,13 +10108,11 @@ def test_dynamic_blueprinting_using_custom_macro(tmp_path: Path) -> None: @{field_b} AS field_b FROM @customer.some_source - """ - ) + """) dynamic_template_py = tmp_path / "models/dynamic_template_custom_macro.py" dynamic_template_py.parent.mkdir(parents=True, exist_ok=True) - dynamic_template_py.write_text( - """ + dynamic_template_py.write_text(""" from sqlmesh import model @model( @@ -10448,24 +10124,22 @@ def test_dynamic_blueprinting_using_custom_macro(tmp_path: Path) -> None: def entrypoint(evaluator): field_a = evaluator.blueprint_var("field_a") return f"SELECT {field_a}, @BLUEPRINT_VAR('field_b') AS field_b FROM @customer.some_source" -""" - ) +""") gen_blueprints = tmp_path / "macros/gen_blueprints.py" gen_blueprints.parent.mkdir(parents=True, exist_ok=True) - gen_blueprints.write_text( - """from sqlmesh import macro + gen_blueprints.write_text("""from sqlmesh import macro @macro() def gen_blueprints(evaluator): return ( "((customer := customer1, field_a := x, field_b := y)," " (customer := customer2, field_a := z, field_b := w))" - )""" - ) + )""") ctx = Context( - config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), paths=tmp_path + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), + paths=tmp_path, ) assert len(ctx.models) == 4 @@ -10480,8 +10154,7 @@ def test_dynamic_blueprinting_using_each(tmp_path: Path) -> None: dynamic_template_sql = tmp_path / "models/dynamic_template_each.sql" dynamic_template_sql.parent.mkdir(parents=True, exist_ok=True) - dynamic_template_sql.write_text( - """ + dynamic_template_sql.write_text(""" MODEL ( name @customer.some_table, kind FULL, @@ -10490,13 +10163,11 @@ def test_dynamic_blueprinting_using_each(tmp_path: Path) -> None: SELECT 1 AS c - """ - ) + """) dynamic_template_py = tmp_path / "models/dynamic_template_each.py" dynamic_template_py.parent.mkdir(parents=True, exist_ok=True) - dynamic_template_py.write_text( - """ + dynamic_template_py.write_text(""" from sqlmesh import model @model( @@ -10506,9 +10177,8 @@ def test_dynamic_blueprinting_using_each(tmp_path: Path) -> None: is_sql=True, ) def entrypoint(evaluator): - return "SELECT 1 AS c" -""" - ) + return "SELECT 1 AS c" +""") model_defaults = ModelDefaultsConfig(dialect="duckdb") variables = {"values": ["customer1", "customer2"]} @@ -10527,8 +10197,7 @@ def test_single_blueprint(tmp_path: Path) -> None: single_blueprint = tmp_path / "models/single_blueprint.sql" single_blueprint.parent.mkdir(parents=True, exist_ok=True) - single_blueprint.write_text( - """ + single_blueprint.write_text(""" MODEL ( name @single_blueprint.some_table, kind FULL, @@ -10536,11 +10205,11 @@ def test_single_blueprint(tmp_path: Path) -> None: ); SELECT 1 AS c - """ - ) + """) ctx = Context( - config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), paths=tmp_path + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), + paths=tmp_path, ) assert len(ctx.models) == 1 @@ -10552,8 +10221,7 @@ def test_blueprinting_with_quotes(tmp_path: Path) -> None: template_with_quoted_vars = tmp_path / "models/template_with_quoted_vars.sql" template_with_quoted_vars.parent.mkdir(parents=True, exist_ok=True) - template_with_quoted_vars.write_text( - """ + template_with_quoted_vars.write_text(""" MODEL ( name m.@{bp_var}, blueprints ( @@ -10563,11 +10231,11 @@ def test_blueprinting_with_quotes(tmp_path: Path) -> None: ); SELECT @bp_var AS c1, @{bp_var} AS c2 - """ - ) + """) ctx = Context( - config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), paths=tmp_path + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), + paths=tmp_path, ) assert len(ctx.models) == 2 @@ -10576,17 +10244,24 @@ def test_blueprinting_with_quotes(tmp_path: Path) -> None: assert m1.name == 'm."a b"' assert m2.name == 'm."c d"' - assert t.cast(exp.Query, m1.render_query()).sql() == '''SELECT "a b" AS "c1", "a b" AS "c2"''' - assert t.cast(exp.Query, m2.render_query()).sql() == '''SELECT 'c d' AS "c1", "c d" AS "c2"''' + assert ( + t.cast(exp.Query, m1.render_query()).sql() + == '''SELECT "a b" AS "c1", "a b" AS "c2"''' + ) + assert ( + t.cast(exp.Query, m2.render_query()).sql() + == '''SELECT 'c d' AS "c1", "c d" AS "c2"''' + ) -def test_blueprint_variable_precedence_sql(tmp_path: Path, assert_exp_eq: t.Callable) -> None: +def test_blueprint_variable_precedence_sql( + tmp_path: Path, assert_exp_eq: t.Callable +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) blueprint_variables = tmp_path / "models/blueprint_variables.sql" blueprint_variables.parent.mkdir(parents=True, exist_ok=True) - blueprint_variables.write_text( - """ + blueprint_variables.write_text(""" MODEL ( name s.@{bp_name}, blueprints ( @@ -10612,8 +10287,7 @@ def test_blueprint_variable_precedence_sql(tmp_path: Path, assert_exp_eq: t.Call @{bp_name} AS bp_name_identifier, @VAR('bp_name') AS bp_name_var_macro_func, @BLUEPRINT_VAR('bp_name') AS bp_name_blueprint_var_macro_func, - """ - ) + """) ctx = Context( config=Config( @@ -10670,8 +10344,7 @@ def test_blueprint_variable_jinja(tmp_path: Path, assert_exp_eq: t.Callable) -> blueprint_variables = tmp_path / "models/blueprint_variables.sql" blueprint_variables.parent.mkdir(parents=True, exist_ok=True) - blueprint_variables.write_text( - """ + blueprint_variables.write_text(""" MODEL ( name s.@{bp_name}, blueprints ( @@ -10690,8 +10363,7 @@ def test_blueprint_variable_jinja(tmp_path: Path, assert_exp_eq: t.Callable) -> '{{ blueprint_var('bp_name') }}' AS bp_name FROM s.{{ blueprint_var('bp_name') }}_source; JINJA_END; - """ - ) + """) ctx = Context( config=Config( model_defaults=ModelDefaultsConfig(dialect="duckdb"), @@ -10714,13 +10386,14 @@ def test_blueprint_variable_jinja(tmp_path: Path, assert_exp_eq: t.Callable) -> ) -def test_blueprint_variable_precedence_python(tmp_path: Path, mocker: MockerFixture) -> None: +def test_blueprint_variable_precedence_python( + tmp_path: Path, mocker: MockerFixture +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) blueprint_variables = tmp_path / "models/blueprint_variables.py" blueprint_variables.parent.mkdir(parents=True, exist_ok=True) - blueprint_variables.write_text( - """ + blueprint_variables.write_text(""" import pandas as pd # noqa: TID253 from sqlglot import exp from sqlmesh import model @@ -10746,8 +10419,7 @@ def entrypoint(context, *args, **kwargs): assert context.blueprint_var("var2") == 1 return pd.DataFrame({"x": [1]}) - """ - ) + """) ctx = Context( config=Config( @@ -10761,14 +10433,15 @@ def entrypoint(context, *args, **kwargs): m = ctx.get_model("s.m", raise_if_missing=True) context = ExecutionContext(mocker.Mock(), {}, None, None) - assert t.cast(pd.DataFrame, list(m.render(context=context))[0]).to_dict() == {"x": {0: 1}} + assert t.cast(pd.DataFrame, list(m.render(context=context))[0]).to_dict() == { + "x": {0: 1} + } def test_python_model_depends_on_blueprints(tmp_path: Path) -> None: sql_model = tmp_path / "models" / "base_blueprints.sql" sql_model.parent.mkdir(parents=True, exist_ok=True) - sql_model.write_text( - """ + sql_model.write_text(""" MODEL ( name test_schema1.@{model_name}, blueprints ((model_name := foo), (model_name := bar)), @@ -10776,13 +10449,11 @@ def test_python_model_depends_on_blueprints(tmp_path: Path) -> None: ); SELECT 1 AS id - """ - ) + """) py_model = tmp_path / "models" / "depends_on_with_blueprint_vars.py" py_model.parent.mkdir(parents=True, exist_ok=True) - py_model.write_text( - """ + py_model.write_text(""" import pandas as pd # noqa: TID253 from sqlmesh import model @@ -10799,8 +10470,7 @@ def test_python_model_depends_on_blueprints(tmp_path: Path) -> None: ) def entrypoint(context, *args, **kwargs): table = context.resolve_table(f"test_schema1.{context.blueprint_var('model_name')}") - return context.fetchdf(f"SELECT * FROM {table}")""" - ) + return context.fetchdf(f"SELECT * FROM {table}")""") ctx = Context( config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), @@ -10816,8 +10486,7 @@ def test_python_model_blueprint_column_names(tmp_path: Path) -> None: """Blueprint variables can be used as column names and types in Python model definitions.""" py_model = tmp_path / "models" / "blueprint_col_names.py" py_model.parent.mkdir(parents=True, exist_ok=True) - py_model.write_text( - """ + py_model.write_text(""" import pandas as pd # noqa: TID253 from sqlmesh import model @@ -10838,8 +10507,7 @@ def entrypoint(context, *args, **kwargs): context.blueprint_var("col_a"): [1], context.blueprint_var("col_b"): [1.5], }) - """ - ) + """) ctx = Context( config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), @@ -10865,8 +10533,7 @@ def test_python_model_variable_column_names(tmp_path: Path) -> None: """Global variables can be used as column names in Python model definitions.""" py_model = tmp_path / "models" / "var_col_names.py" py_model.parent.mkdir(parents=True, exist_ok=True) - py_model.write_text( - """ + py_model.write_text(""" import pandas as pd # noqa: TID253 from sqlmesh import model @@ -10880,8 +10547,7 @@ def test_python_model_variable_column_names(tmp_path: Path) -> None: ) def entrypoint(context, *args, **kwargs): return pd.DataFrame({"revenue": [1], "static_col": ["x"]}) - """ - ) + """) ctx = Context( config=Config( @@ -10908,8 +10574,7 @@ def get_current_date(evaluator): return f"'{now().date()}'" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name test_model, dialect duckdb); @DEF(curr_date, @get_current_date()); @@ -10919,8 +10584,7 @@ def get_current_date(evaluator): ) SELECT * FROM discount_promotion_dates - """ - ) + """) model = load_sql_based_model(expressions) assert_exp_eq( model.render_query(), @@ -10943,8 +10607,7 @@ def test_seed_dont_coerce_na_into_null(tmp_path): with open(model_csv_path, "w", encoding="utf-8") as fd: fd.write("code\nNA") - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.seed, kind SEED ( @@ -10956,10 +10619,11 @@ def test_seed_dont_coerce_na_into_null(tmp_path): ), ), ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) assert isinstance(model.kind, SeedKind) assert model.seed is not None @@ -10973,8 +10637,7 @@ def test_seed_coerce_datetime(tmp_path): with open(model_csv_path, "w", encoding="utf-8") as fd: fd.write("bad_datetime\n9999-12-31 23:59:59") - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.seed, kind SEED ( @@ -10984,10 +10647,11 @@ def test_seed_coerce_datetime(tmp_path): bad_datetime datetime, ), ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) df = next(model.render(context=None)) assert df["bad_datetime"].iloc[0] == "9999-12-31 23:59:59" @@ -10998,8 +10662,7 @@ def test_seed_invalid_date_column(tmp_path): with open(model_csv_path, "w", encoding="utf-8") as fd: fd.write("bad_date\n9999-12-31\n2025-01-01\n1000-01-01") - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.seed, kind SEED ( @@ -11009,10 +10672,11 @@ def test_seed_invalid_date_column(tmp_path): bad_date date, ), ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) df = next(model.render(context=None)) # The conversion to date should not raise an error assert df["bad_date"].to_list() == ["9999-12-31", "2025-01-01", "1000-01-01"] @@ -11024,8 +10688,7 @@ def test_seed_missing_columns(tmp_path): with open(model_csv_path, "w", encoding="utf-8") as fd: fd.write("key,value\n1,2\n3,4") - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL ( name db.seed, kind SEED ( @@ -11037,19 +10700,20 @@ def test_seed_missing_columns(tmp_path): missing_column int, ), ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) with pytest.raises( - ConfigError, match="Seed model 'db.seed' has missing columns: {'missing_column'}.*" + ConfigError, + match="Seed model 'db.seed' has missing columns: {'missing_column'}.*", ): next(model.render(context=None)) def test_missing_column_data_in_columns_key(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -11059,10 +10723,11 @@ def test_missing_column_data_in_columns_key(): culprit, other_column double, ) ); - """ - ) + """) with pytest.raises(ConfigError, match="Missing data type for column 'culprit'."): - load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) def test_ignored_rules_serialization(): @@ -11107,12 +10772,16 @@ def test_data_hash_unchanged_when_column_type_uses_default_dialect(): assert model.data_hash == deserialized_model.data_hash -def test_transitive_dependency_of_metadata_only_object_is_metadata_only(tmp_path: Path) -> None: +def test_transitive_dependency_of_metadata_only_object_is_metadata_only( + tmp_path: Path, +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) test_model = tmp_path / "models/test_model.sql" test_model.parent.mkdir(parents=True, exist_ok=True) - test_model.write_text("MODEL (name test_model, kind FULL); @metadata_macro(); SELECT 1 AS c") + test_model.write_text( + "MODEL (name test_model, kind FULL); @metadata_macro(); SELECT 1 AS c" + ) metadata_macro_code = """ from sqlglot import parse_one @@ -11186,7 +10855,9 @@ def metadata_macro(evaluator): assert new_snapshot.change_category == SnapshotChangeCategory.METADATA -def test_vars_are_taken_into_account_when_propagating_metadata_status(tmp_path: Path) -> None: +def test_vars_are_taken_into_account_when_propagating_metadata_status( + tmp_path: Path, +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) test_model = tmp_path / "models/test_model.sql" @@ -11279,17 +10950,24 @@ def m4_non_metadata_references_v6(evaluator): assert macro_evaluator.blueprint_var("v5") == exp.Literal.number("5") query_with_vars = macro_evaluator.transform( - parse_one("SELECT " + ", ".join(f"@v{var}, @VAR('v{var}')" for var in [1, 2, 3, 6])) + parse_one( + "SELECT " + ", ".join(f"@v{var}, @VAR('v{var}')" for var in [1, 2, 3, 6]) + ) ) assert t.cast(exp.Expr, query_with_vars).sql() == "SELECT 1, 1, 2, 2, 3, 3, 6, 6" query_with_blueprint_vars = macro_evaluator.transform( - parse_one("SELECT " + ", ".join(f"@v{var}, @BLUEPRINT_VAR('v{var}')" for var in [4, 5])) + parse_one( + "SELECT " + + ", ".join(f"@v{var}, @BLUEPRINT_VAR('v{var}')" for var in [4, 5]) + ) ) assert t.cast(exp.Expr, query_with_blueprint_vars).sql() == "SELECT 4, 4, 5, 5" -def test_variable_mentioned_in_both_metadata_and_non_metadata_macro(tmp_path: Path) -> None: +def test_variable_mentioned_in_both_metadata_and_non_metadata_macro( + tmp_path: Path, +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) test_model = tmp_path / "models/test_model.sql" @@ -11316,7 +10994,9 @@ def m2_references_v_non_metadata(evaluator): test_macros.write_text(macro_code) ctx = Context( - config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"), variables={"v": 1}), + config=Config( + model_defaults=ModelDefaultsConfig(dialect="duckdb"), variables={"v": 1} + ), paths=tmp_path, ) model = ctx.get_model("test_model") @@ -11324,11 +11004,16 @@ def m2_references_v_non_metadata(evaluator): python_env = model.python_env assert len(python_env) == 3 - assert set(python_env) > {"m1_references_v_metadata", "m2_references_v_non_metadata"} + assert set(python_env) > { + "m1_references_v_metadata", + "m2_references_v_non_metadata", + } assert python_env.get(c.SQLMESH_VARS) == Executable.value({"v": 1}) -def test_only_top_level_macro_func_impacts_var_descendant_metadata_status(tmp_path: Path) -> None: +def test_only_top_level_macro_func_impacts_var_descendant_metadata_status( + tmp_path: Path, +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) test_model = tmp_path / "models/test_model.sql" @@ -11353,7 +11038,9 @@ def m2_non_metadata(evaluator, *args): test_macros.write_text(macro_code) ctx = Context( - config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"), variables={"v": 1}), + config=Config( + model_defaults=ModelDefaultsConfig(dialect="duckdb"), variables={"v": 1} + ), paths=tmp_path, ) model = ctx.get_model("test_model") @@ -11362,15 +11049,21 @@ def m2_non_metadata(evaluator, *args): assert len(python_env) == 3 assert set(python_env) > {"m1_metadata", "m2_non_metadata"} - assert python_env.get(c.SQLMESH_VARS_METADATA) == Executable.value({"v": 1}, is_metadata=True) + assert python_env.get(c.SQLMESH_VARS_METADATA) == Executable.value( + {"v": 1}, is_metadata=True + ) -def test_non_metadata_object_takes_precedence_over_metadata_only_object(tmp_path: Path) -> None: +def test_non_metadata_object_takes_precedence_over_metadata_only_object( + tmp_path: Path, +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) test_model = tmp_path / "models/test_model.sql" test_model.parent.mkdir(parents=True, exist_ok=True) - test_model.write_text("MODEL (name test_model, kind FULL); @m1(); @m2(); SELECT 1 AS c") + test_model.write_text( + "MODEL (name test_model, kind FULL); @m1(); @m2(); SELECT 1 AS c" + ) macro_code = """ from sqlglot import parse_one @@ -11426,8 +11119,7 @@ def test_macros_referenced_in_metadata_statements_and_properties_are_metadata_on test_model = tmp_path / "models/test_model.sql" test_model.parent.mkdir(parents=True, exist_ok=True) - test_model.write_text( - """ + test_model.write_text(""" MODEL ( name test_model, kind FULL, @@ -11457,8 +11149,7 @@ def test_macros_referenced_in_metadata_statements_and_properties_are_metadata_on ON_VIRTUAL_UPDATE_END; - """ - ) + """) macro_code = """ from sqlglot import exp @@ -11574,7 +11265,9 @@ def test_scd_type_2_full_history_restatement(): assert ModelKindName.SCD_TYPE_2.full_history_restatement_only is True assert ModelKindName.SCD_TYPE_2_BY_TIME.full_history_restatement_only is True assert ModelKindName.SCD_TYPE_2_BY_COLUMN.full_history_restatement_only is True - assert ModelKindName.INCREMENTAL_BY_TIME_RANGE.full_history_restatement_only is False + assert ( + ModelKindName.INCREMENTAL_BY_TIME_RANGE.full_history_restatement_only is False + ) def test_python_model_boolean_values(): @@ -11600,8 +11293,7 @@ def test_model(context, **kwargs): def test_var_in_def(assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, kind INCREMENTAL_BY_TIME_RANGE( @@ -11612,8 +11304,7 @@ def test_var_in_def(assert_exp_eq): @DEF(var, @start_ds); SELECT @var AS ds - """ - ) + """) model = load_sql_based_model(expressions) @@ -11638,7 +11329,10 @@ def test_formatting_flag_serde(): ) model = load_sql_based_model(expressions) - assert model.render_definition()[0].sql() == "MODEL (\nname test_model,\nformatting False\n)" + assert ( + model.render_definition()[0].sql() + == "MODEL (\nname test_model,\nformatting False\n)" + ) model_json = model.json() assert "formatting" not in json.loads(model_json) @@ -11656,8 +11350,7 @@ def test_runtime_stage(evaluator): noop() return evaluator.runtime_stage - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect spark, @@ -11667,8 +11360,7 @@ def test_runtime_stage(evaluator): JINJA_QUERY_BEGIN; SELECT '{{ test_runtime_stage() }}' AS a, '{{ test_runtime_stage_jinja('bla') }}' AS b; JINJA_END; - """ - ) + """) jinja_macros = JinjaMacroRegistry( root_macros={ @@ -11684,7 +11376,9 @@ def test_runtime_stage(evaluator): assert set(model.python_env) == {"noop", "test_runtime_stage"} -def test_python_env_references_are_unequal_but_point_to_same_definition(tmp_path: Path) -> None: +def test_python_env_references_are_unequal_but_point_to_same_definition( + tmp_path: Path, +) -> None: # This tests for regressions against an edge case bug which was due to reloading modules # in sqlmesh.utils.metaprogramming.import_python_file. Depending on the module loading # order, we could get a "duplicate symbol in python env" error, even though the references @@ -11703,15 +11397,12 @@ def test_python_env_references_are_unequal_but_point_to_same_definition(tmp_path file_b = tmp_path / "macros" / "b.py" file_c = tmp_path / "macros" / "c.py" - file_a.write_text( - """from macros.c import target + file_a.write_text("""from macros.c import target def f1(): target() -""" - ) - file_b.write_text( - """from sqlmesh import macro +""") + file_b.write_text("""from sqlmesh import macro from macros.a import f1 from macros.c import target @@ -11723,16 +11414,15 @@ def first_macro(evaluator): @macro() def second_macro(evaluator): target() -""" - ) - file_c.write_text( - """def target(): +""") + file_c.write_text("""def target(): pass -""" - ) +""") model_file = tmp_path / "models" / "model.sql" - model_file.write_text("MODEL (name a); @first_macro(); @second_macro(); SELECT 1 AS c") + model_file.write_text( + "MODEL (name a); @first_macro(); @second_macro(); SELECT 1 AS c" + ) ctx = Context(paths=tmp_path, config=config, load=False) loader = ctx._loaders[0] @@ -11783,8 +11473,7 @@ def test_unequal_duplicate_python_env_references_are_prohibited(tmp_path: Path) file_a = tmp_path / "macros" / "unimportant_macro.py" file_b = tmp_path / "macros" / "just_f.py" - file_a.write_text( - """from sqlmesh import macro + file_a.write_text("""from sqlmesh import macro from macros.just_f import f a = False @@ -11794,18 +11483,17 @@ def unimportant_macro(evaluator): print(a) f() return 1 -""" - ) - file_b.write_text( - """a = 0 +""") + file_b.write_text("""a = 0 def f(): print(a) -""" - ) +""") model_file = tmp_path / "models" / "model.sql" - model_file.write_text("MODEL (name m); SELECT @unimportant_macro() AS unimportant_macro") + model_file.write_text( + "MODEL (name m); SELECT @unimportant_macro() AS unimportant_macro" + ) with pytest.raises(SQLMeshError, match=r"duplicate definitions found"): Context(paths=tmp_path, config=config) @@ -11821,8 +11509,7 @@ def test_semicolon_is_metadata_only_change(tmp_path, assert_exp_eq): ) model_file = tmp_path / "models" / "model_with_semicolon.sql" - model_file.write_text( - """ + model_file.write_text(""" MODEL ( name sqlmesh_example.incremental_model_with_semicolon, kind INCREMENTAL_BY_TIME_RANGE ( @@ -11840,8 +11527,7 @@ def test_semicolon_is_metadata_only_change(tmp_path, assert_exp_eq): ; --Just a comment - """ - ) + """) ctx = Context(paths=tmp_path, config=config) model = ctx.get_model("sqlmesh_example.incremental_model_with_semicolon") @@ -11855,9 +11541,7 @@ def test_semicolon_is_metadata_only_change(tmp_path, assert_exp_eq): ) ctx.format() - assert ( - model_file.read_text() - == """MODEL ( + assert model_file.read_text() == """MODEL ( name sqlmesh_example.incremental_model_with_semicolon, kind INCREMENTAL_BY_TIME_RANGE ( time_column event_date @@ -11873,13 +11557,11 @@ def test_semicolon_is_metadata_only_change(tmp_path, assert_exp_eq): '2020-01-01'::DATE AS event_date; /* Just a comment */""" - ) ctx.plan(no_prompts=True, auto_apply=True) model_file = tmp_path / "models" / "model_with_semicolon.sql" - model_file.write_text( - """ + model_file.write_text(""" MODEL ( name sqlmesh_example.incremental_model_with_semicolon, kind INCREMENTAL_BY_TIME_RANGE ( @@ -11894,8 +11576,7 @@ def test_semicolon_is_metadata_only_change(tmp_path, assert_exp_eq): 1 AS id, 1 AS item_id, CAST('2020-01-01' AS DATE) AS event_date - """ - ) + """) ctx.load() plan = ctx.plan(no_prompts=True, auto_apply=True) @@ -12008,7 +11689,9 @@ def test_resolve_interpolated_variables_when_parsing_python_deps(): ) def unimportant_testing_model(context, **kwargs): table1 = context.resolve_table(f"{context.var('schema_name')}.table_name") - table2 = context.resolve_table(f"{context.blueprint_var('schema_name')}.table_name") + table2 = context.resolve_table( + f"{context.blueprint_var('schema_name')}.table_name" + ) return context.fetchdf(exp.select("*").from_(table)) @@ -12021,7 +11704,9 @@ def unimportant_testing_model(context, **kwargs): assert m.depends_on == {'"foo"."table_name"', '"baz"."table_name"'} assert m.python_env.get(c.SQLMESH_VARS) == Executable.value({"schema_name": "foo"}) - assert m.python_env.get(c.SQLMESH_BLUEPRINT_VARS) == Executable.value({"schema_name": "baz"}) + assert m.python_env.get(c.SQLMESH_BLUEPRINT_VARS) == Executable.value( + {"schema_name": "baz"} + ) @macro() def unimportant_testing_macro(evaluator, *projections): @@ -12067,16 +11752,14 @@ def test_extract_schema_in_post_statement(tmp_path: Path) -> None: model2 = tmp_path / "models" / "child_model.sql" model2.parent.mkdir(parents=True, exist_ok=True) - model2.write_text( - """ + model2.write_text(""" MODEL (name y); SELECT c FROM x; ON_VIRTUAL_UPDATE_BEGIN; @check_schema('y'); @check_self_schema(); ON_VIRTUAL_UPDATE_END; - """ - ) + """) check_schema = tmp_path / "macros/check_schema.py" check_schema.parent.mkdir(parents=True, exist_ok=True) @@ -12099,19 +11782,19 @@ def check_self_schema(evaluator): context.plan(no_prompts=True, auto_apply=True) -def test_model_relies_on_os_getenv(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: +def test_model_relies_on_os_getenv( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: init_example_project(tmp_path, engine_type="duckdb", template=ProjectTemplate.EMPTY) - (tmp_path / "macros" / "getenv_macro.py").write_text( - """ + (tmp_path / "macros" / "getenv_macro.py").write_text(""" from os import getenv from sqlmesh import macro @macro() def getenv_macro(evaluator): getenv("foo", None) - return 1""" - ) + return 1""") (tmp_path / "models" / "model.sql").write_text( "MODEL (name test); SELECT @getenv_macro() AS foo" ) @@ -12122,15 +11805,13 @@ def getenv_macro(evaluator): def test_invalid_sql_model_query() -> None: for kind in ("", ", KIND FULL"): - expressions = d.parse( - f""" + expressions = d.parse(f""" MODEL (name db.table{kind}); JINJA_STATEMENT_BEGIN; SELECT 1 AS c; JINJA_END; - """ - ) + """) with pytest.raises( ConfigError, @@ -12148,8 +11829,7 @@ def test_query_label_macro(evaluator): def test_authorization_macro(evaluator): return exp.Literal.string("test_authorization") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, session_properties ( @@ -12159,8 +11839,7 @@ def test_authorization_macro(evaluator): ); SELECT 1 AS c; - """ - ) + """) model = load_sql_based_model(expressions) assert model.session_properties == { @@ -12179,8 +11858,7 @@ def test_query_tags_macro() -> None: def test_query_tags_macro(evaluator): return "MAP('team', 'data-eng')" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, dialect databricks, @@ -12190,8 +11868,7 @@ def test_query_tags_macro(evaluator): ); SELECT 1 AS c; - """ - ) + """) model = load_sql_based_model(expressions) assert model.session_properties == { @@ -12204,8 +11881,7 @@ def test_query_tags_macro(evaluator): def test_boolean_property_validation() -> None: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.table, enabled @IF(TRUE, TRUE, FALSE), @@ -12213,15 +11889,13 @@ def test_boolean_property_validation() -> None: ); SELECT 1 AS c; - """ - ) + """) model = load_sql_based_model(expressions, dialect="tsql") assert model.enabled def test_datetime_without_timezone_variable_redshift() -> None: - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name test, kind INCREMENTAL_BY_TIME_RANGE ( @@ -12234,8 +11908,7 @@ def test_datetime_without_timezone_variable_redshift() -> None: ); SELECT @start_dtntz AS test_time_col - """ - ) + """) model = load_sql_based_model(expressions, dialect="redshift") assert ( @@ -12249,8 +11922,7 @@ def test_python_model_cron_with_blueprints(tmp_path: Path) -> None: cron_blueprint_model = tmp_path / "models" / "cron_blueprint.py" cron_blueprint_model.parent.mkdir(parents=True, exist_ok=True) - cron_blueprint_model.write_text( - """ + cron_blueprint_model.write_text(""" import typing as t from datetime import datetime @@ -12286,11 +11958,11 @@ def entrypoint( "customer": [context.blueprint_var("customer")], } ) -""" - ) +""") context = Context( - paths=tmp_path, config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + paths=tmp_path, + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")), ) models = context.models @@ -12333,8 +12005,7 @@ def test_python_model_cron_macro_rendering(tmp_path: Path) -> None: cron_macro_model = tmp_path / "models" / "cron_macro.py" cron_macro_model.parent.mkdir(parents=True, exist_ok=True) - cron_macro_model.write_text( - """ + cron_macro_model.write_text(""" import pandas as pd from sqlmesh import model @@ -12346,8 +12017,7 @@ def test_python_model_cron_macro_rendering(tmp_path: Path) -> None: ) def entrypoint(context, **kwargs): return pd.DataFrame([{"a": 1}]) -""" - ) +""") # Test with cron alias context_daily = Context( @@ -12380,8 +12050,7 @@ def test_python_model_normal_cron(tmp_path: Path) -> None: cron_macro_model = tmp_path / "models" / "cron_macro.py" cron_macro_model.parent.mkdir(parents=True, exist_ok=True) - cron_macro_model.write_text( - """ + cron_macro_model.write_text(""" import pandas as pd from sqlmesh import model @@ -12393,8 +12062,7 @@ def test_python_model_normal_cron(tmp_path: Path) -> None: ) def entrypoint(context, **kwargs): return pd.DataFrame([{"a": 1}]) -""" - ) +""") # Test with cron alias context_daily = Context( @@ -12416,7 +12084,9 @@ def test_render_query_optimize_query_false(assert_exp_eq, sushi_context): model = sushi_context.get_model("sushi.top_waiters") model = model.copy(update={"optimize_query": False}) - upstream_model_version = sushi_context.get_snapshot("sushi.waiter_revenue_by_day").version + upstream_model_version = sushi_context.get_snapshot( + "sushi.waiter_revenue_by_day" + ).version assert_exp_eq( model.render_query(snapshots=snapshots).sql(), @@ -12446,8 +12116,7 @@ def test_render_query_optimize_query_false(assert_exp_eq, sushi_context): def test_each_macro_with_paren_expression_arg(assert_exp_eq): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name dataset.@table_name, kind VIEW, @@ -12469,8 +12138,7 @@ def test_each_macro_with_paren_expression_arg(assert_exp_eq): ); SELECT @EACH(@event_columns, x -> x) - """ - ) + """) models = load_sql_based_models(expressions, lambda _: {}) @@ -12515,8 +12183,11 @@ def test_each_macro_with_paren_expression_arg(assert_exp_eq): ("@M1(@BLUEPRINT_VAR(@VAR('v1')))", {"v1"}), ], ) -def test_extract_macro_func_variable_references(macro_func: str, variables: t.Set[str]) -> None: - from sqlmesh.core.model.common import _extract_macro_func_variable_references +def test_extract_macro_func_variable_references( + macro_func: str, variables: t.Set[str] +) -> None: + from sqlmesh.core.model.common import \ + _extract_macro_func_variable_references macro_func_ast = parse_one(macro_func) assert _extract_macro_func_variable_references(macro_func_ast, True)[0] == variables @@ -12572,23 +12243,20 @@ def test_text_diff_optimize_query(): def test_raw_jinja_raw_tag(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name test); JINJA_QUERY_BEGIN; SELECT {% raw %} '{{ foo }}' {% endraw %} AS col; JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) assert model.render_query().sql() == "SELECT '{{ foo }}' AS \"col\"" def test_use_original_sql(): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name test); CREATE TABLE pre ( @@ -12602,13 +12270,17 @@ def test_use_original_sql(): CREATE TABLE post ( b INT ); - """ - ) + """) model = load_sql_based_model(expressions) assert model.query_.sql == "SELECT\n 1,\n 2" - assert model.pre_statements_[0].sql == "CREATE TABLE pre (\n a INT\n )" - assert model.post_statements_[0].sql == "CREATE TABLE post (\n b INT\n );" + assert ( + model.pre_statements_[0].sql == "CREATE TABLE pre (\n a INT\n )" + ) + assert ( + model.post_statements_[0].sql + == "CREATE TABLE post (\n b INT\n );" + ) # Now manually create the model and make sure that the original SQL is not used model_query = d.parse_one("SELECT 1 AS one") @@ -12642,8 +12314,7 @@ def test_case_sensitive_macro_locals(tmp_path: Path) -> None: macro_file = tmp_path / "macros" / "some_macro_with_globals.py" macro_file.parent.mkdir(parents=True, exist_ok=True) - macro_file.write_text( - """from sqlmesh import macro + macro_file.write_text("""from sqlmesh import macro x = 1 X = 2 @@ -12654,8 +12325,7 @@ def my_macro(evaluator): assert evaluator.locals.get("X") == 2 return x + X -""" - ) +""") test_model = tmp_path / "models" / "test_model.sql" test_model.parent.mkdir(parents=True, exist_ok=True) test_model.write_text("MODEL (name test_model, kind FULL); SELECT @my_macro() AS c") @@ -12684,7 +12354,10 @@ def test_grants(): assert model.grants == { "select": ["user1", "123", "admin_role", "user2"], "insert": ["admin"], - "roles/bigquery.dataViewer": ["group:data_eng@company.com", "user:someone@company.com"], + "roles/bigquery.dataViewer": [ + "group:data_eng@company.com", + "user:someone@company.com", + ], "update": ["admin"], } @@ -12845,7 +12518,9 @@ def test_grants_macro_var_in_array_flattening(): SELECT 1 as id """) - model = load_sql_based_model(expressions, variables={"admins": ["admin1", "admin2"]}) + model = load_sql_based_model( + expressions, variables={"admins": ["admin1", "admin2"]} + ) assert model.grants == {"select": ["user1", "admin1", "admin2", "user3"]} model2 = load_sql_based_model(expressions, variables={"admins": "super_admin"}) @@ -12875,7 +12550,9 @@ def test_grants_unresolved_macro_errors(): MODEL (name test.bad1, kind FULL, grants ('select' = @VAR('undefined'))); SELECT 1 as id """) - with pytest.raises(ConfigError, match=r"Invalid grants configuration for 'select': NULL value"): + with pytest.raises( + ConfigError, match=r"Invalid grants configuration for 'select': NULL value" + ): load_sql_based_model(expressions1) expressions2 = d.parse(""" @@ -12889,7 +12566,9 @@ def test_grants_unresolved_macro_errors(): MODEL (name test.bad3, kind FULL, grants ('select' = ['user', @VAR('undefined')])); SELECT 1 as id """) - with pytest.raises(ConfigError, match=r"Invalid grants configuration for 'select': NULL value"): + with pytest.raises( + ConfigError, match=r"Invalid grants configuration for 'select': NULL value" + ): load_sql_based_model(expressions3) @@ -12922,22 +12601,19 @@ def test_model_macro_using_locals_called_from_jinja(assert_exp_eq) -> None: def execution_date(evaluator): return f"""'{evaluator.locals.get("execution_date")}'""" - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); JINJA_QUERY_BEGIN; SELECT {{ execution_date() }} AS col; JINJA_END; - """ - ) + """) model = load_sql_based_model(expressions) assert_exp_eq(model.render_query(), '''SELECT '1970-01-01' AS "col"''') def test_audits_in_embedded_model(): - expression = d.parse( - """ + expression = d.parse(""" MODEL ( name test.embedded_with_audits, kind EMBEDDED, @@ -12945,9 +12621,10 @@ def test_audits_in_embedded_model(): ); SELECT 1 AS id, 'A' as value - """ - ) - with pytest.raises(ConfigError, match="Audits are not supported for embedded models"): + """) + with pytest.raises( + ConfigError, match="Audits are not supported for embedded models" + ): load_sql_based_model(expression).validate_definition() @@ -12994,9 +12671,9 @@ def test_default_catalog_not_leaked_to_unsupported_gateway(): f"Default gateway catalog leaked into catalog-unsupported gateway model. " f"Expected no catalog, got: {model.catalog}" ) - assert "example_catalog" not in model.fqn, ( - f"Default gateway catalog found in model FQN: {model.fqn}" - ) + assert ( + "example_catalog" not in model.fqn + ), f"Default gateway catalog found in model FQN: {model.fqn}" def test_default_catalog_still_applied_to_supported_gateway(): @@ -13035,7 +12712,9 @@ def test_default_catalog_still_applied_to_supported_gateway(): assert len(models) == 1 model = models[0] - assert model.catalog == "other_db", f"Expected catalog 'other_db', got: {model.catalog}" + assert ( + model.catalog == "other_db" + ), f"Expected catalog 'other_db', got: {model.catalog}" @pytest.mark.parametrize( @@ -13202,10 +12881,10 @@ def test_blueprint_catalog_not_cross_contaminated(): ch_model = next(m for m in models if "ch_schema" in m.fqn) db_model = next(m for m in models if "db_schema" in m.fqn) - assert not ch_model.catalog, ( - f"Catalog leaked into ClickHouse blueprint. Got: {ch_model.catalog}" - ) + assert ( + not ch_model.catalog + ), f"Catalog leaked into ClickHouse blueprint. Got: {ch_model.catalog}" - assert db_model.catalog == "example_catalog", ( - f"Catalog lost for DuckDB blueprint after ClickHouse iteration. Got: {db_model.catalog}" - ) + assert ( + db_model.catalog == "example_catalog" + ), f"Catalog lost for DuckDB blueprint after ClickHouse iteration. Got: {db_model.catalog}" diff --git a/tests/core/test_notification_target.py b/tests/core/test_notification_target.py index 57b21f2e0f..1221f5e53e 100644 --- a/tests/core/test_notification_target.py +++ b/tests/core/test_notification_target.py @@ -5,16 +5,16 @@ import pytest -from sqlmesh.core.notification_target import ( - ConsoleNotificationTarget, - NotificationEvent, - NotificationStatus, - NotificationTargetManager, -) +from sqlmesh.core.notification_target import (ConsoleNotificationTarget, + NotificationEvent, + NotificationStatus, + NotificationTargetManager) @pytest.fixture -def notification_target_manager_with_spy(mocker) -> tuple[NotificationTargetManager, t.Callable]: +def notification_target_manager_with_spy( + mocker, +) -> tuple[NotificationTargetManager, t.Callable]: console_notification_target_send_spy = mocker.spy(ConsoleNotificationTarget, "send") console_notification_target = ConsoleNotificationTarget() test_user_console_notification_target = ConsoleNotificationTarget( @@ -34,7 +34,9 @@ def notification_target_manager_with_spy(mocker) -> tuple[NotificationTargetMana def test_notify(notification_target_manager_with_spy): notification_target_manager, spy = notification_target_manager_with_spy - notification_target_manager.notify(NotificationEvent.APPLY_START, "prod", "a-plan-id") + notification_target_manager.notify( + NotificationEvent.APPLY_START, "prod", "a-plan-id" + ) spy.assert_called_once_with( mock.ANY, NotificationStatus.INFO, diff --git a/tests/core/test_plan.py b/tests/core/test_plan.py index 79313c36bd..3ec0be1ce9 100644 --- a/tests/core/test_plan.py +++ b/tests/core/test_plan.py @@ -4,48 +4,34 @@ from unittest.mock import patch import pytest - -from sqlmesh.core.console import TerminalConsole -from sqlmesh.utils.metaprogramming import Executable -from tests.core.test_table_diff import create_test_console import time_machine from pytest_mock.plugin import MockerFixture -from sqlglot import parse_one, exp +from sqlglot import exp, parse_one from sqlmesh.core import dialect as d +from sqlmesh.core.console import TerminalConsole from sqlmesh.core.context import Context from sqlmesh.core.context_diff import ContextDiff -from sqlmesh.core.environment import EnvironmentNamingInfo, EnvironmentStatements -from sqlmesh.core.model import ( - ExternalModel, - FullKind, - IncrementalByTimeRangeKind, - IncrementalUnmanagedKind, - SeedKind, - SeedModel, - SqlModel, - ModelKindName, -) -from sqlmesh.core.model.kind import OnDestructiveChange, OnAdditiveChange, ViewKind +from sqlmesh.core.environment import (EnvironmentNamingInfo, + EnvironmentStatements) +from sqlmesh.core.model import (ExternalModel, FullKind, + IncrementalByTimeRangeKind, + IncrementalUnmanagedKind, ModelKindName, + SeedKind, SeedModel, SqlModel) +from sqlmesh.core.model.kind import (OnAdditiveChange, OnDestructiveChange, + ViewKind) from sqlmesh.core.model.seed import Seed from sqlmesh.core.plan import Plan, PlanBuilder, SnapshotIntervals -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Snapshot, - SnapshotChangeCategory, - SnapshotDataVersion, - SnapshotFingerprint, -) +from sqlmesh.core.snapshot import (DeployabilityIndex, Snapshot, + SnapshotChangeCategory, SnapshotDataVersion, + SnapshotFingerprint) from sqlmesh.utils.dag import DAG -from sqlmesh.utils.date import ( - now, - to_date, - to_datetime, - to_timestamp, - yesterday_ds, -) -from sqlmesh.utils.errors import PlanError, NoChangesPlanError +from sqlmesh.utils.date import (now, to_date, to_datetime, to_timestamp, + yesterday_ds) +from sqlmesh.utils.errors import NoChangesPlanError, PlanError +from sqlmesh.utils.metaprogramming import Executable from sqlmesh.utils.rich import strip_ansi_codes +from tests.core.test_table_diff import create_test_console def test_forward_only_plan_sets_version(make_snapshot, mocker: MockerFixture): @@ -155,7 +141,10 @@ def test_forward_only_dev(make_snapshot, mocker: MockerFixture): plan = PlanBuilder(context_diff, forward_only=True, is_dev=True).build() assert plan.restatements == { - updated_snapshot.snapshot_id: (to_timestamp(expected_start), expected_interval_end) + updated_snapshot.snapshot_id: ( + to_timestamp(expected_start), + expected_interval_end, + ) } assert plan.start == to_date(expected_start) assert plan.end == expected_end @@ -324,10 +313,14 @@ def test_paused_forward_only_parent(make_snapshot, mocker: MockerFixture): ) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) - snapshot_b_old = make_snapshot(SqlModel(name="b", query=parse_one("select 2, ds from a"))) + snapshot_b_old = make_snapshot( + SqlModel(name="b", query=parse_one("select 2, ds from a")) + ) snapshot_b_old.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=False) - snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select 3, ds from a"))) + snapshot_b = make_snapshot( + SqlModel(name="b", query=parse_one("select 3, ds from a")) + ) assert not snapshot_b.version context_diff = ContextDiff( @@ -373,7 +366,10 @@ def test_forward_only_plan_allow_destructive_models( added=set(), removed_snapshots={}, modified_snapshots={snapshot_a.name: (snapshot_a, snapshot_a_old)}, - snapshots={snapshot_a.snapshot_id: snapshot_a, snapshot_a_old.snapshot_id: snapshot_a_old}, + snapshots={ + snapshot_a.snapshot_id: snapshot_a, + snapshot_a_old.snapshot_id: snapshot_a_old, + }, new_snapshots={snapshot_a.snapshot_id: snapshot_a}, previous_plan_id=None, previously_promoted_snapshot_ids=set(), @@ -451,7 +447,10 @@ def test_forward_only_plan_allow_destructive_models( snapshot_c.snapshot_id: snapshot_c, snapshot_c_old.snapshot_id: snapshot_c_old, }, - new_snapshots={snapshot_b.snapshot_id: snapshot_b, snapshot_c.snapshot_id: snapshot_c}, + new_snapshots={ + snapshot_b.snapshot_id: snapshot_b, + snapshot_c.snapshot_id: snapshot_c, + }, previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, @@ -470,7 +469,9 @@ def test_forward_only_plan_allow_destructive_models( PlanError, match="""Plan requires a destructive change to a forward-only model.""", ): - PlanBuilder(context_diff_b, forward_only=True, allow_destructive_models=['"b"']).build() + PlanBuilder( + context_diff_b, forward_only=True, allow_destructive_models=['"b"'] + ).build() logger = logging.getLogger("sqlmesh.core.plan.builder") with patch.object(logger, "warning") as mock_logger: @@ -498,7 +499,10 @@ def test_forward_only_plan_allow_additive_models( added=set(), removed_snapshots={}, modified_snapshots={snapshot_a.name: (snapshot_a, snapshot_a_old)}, - snapshots={snapshot_a.snapshot_id: snapshot_a, snapshot_a_old.snapshot_id: snapshot_a_old}, + snapshots={ + snapshot_a.snapshot_id: snapshot_a, + snapshot_a_old.snapshot_id: snapshot_a_old, + }, new_snapshots={snapshot_a.snapshot_id: snapshot_a}, previous_plan_id=None, previously_promoted_snapshot_ids=set(), @@ -508,13 +512,18 @@ def test_forward_only_plan_allow_additive_models( environment_statements=[], ) - with pytest.raises(PlanError, match="Plan requires an additive change to a forward-only model"): + with pytest.raises( + PlanError, match="Plan requires an additive change to a forward-only model" + ): PlanBuilder(context_diff_a, forward_only=False).build() console = TerminalConsole() log_warning_spy = mocker.spy(console, "log_warning") assert PlanBuilder( - context_diff_a, forward_only=False, allow_additive_models=['"a"'], console=console + context_diff_a, + forward_only=False, + allow_additive_models=['"a"'], + console=console, ).build() assert log_warning_spy.call_count == 0 @@ -532,7 +541,10 @@ def test_forward_only_plan_allow_additive_models( added=set(), removed_snapshots={}, modified_snapshots={snapshot_a.name: (snapshot_a, snapshot_a_old)}, - snapshots={snapshot_a.snapshot_id: snapshot_a, snapshot_a_old.snapshot_id: snapshot_a_old}, + snapshots={ + snapshot_a.snapshot_id: snapshot_a, + snapshot_a_old.snapshot_id: snapshot_a_old, + }, new_snapshots={snapshot_a.snapshot_id: snapshot_a}, previous_plan_id=None, previously_promoted_snapshot_ids=set(), @@ -686,7 +698,8 @@ def test_forward_only_model_on_destructive_change( nodes={'"a"': snapshot_a_old3.model, '"b"': snapshot_b_old3.model}, ) snapshot_c3 = make_snapshot( - snapshot_c_old3.model, nodes={'"a"': snapshot_a3.model, '"b"': snapshot_b3.model} + snapshot_c_old3.model, + nodes={'"a"': snapshot_a3.model, '"b"': snapshot_b3.model}, ) snapshot_c3.previous_versions = ( SnapshotDataVersion( @@ -738,7 +751,8 @@ def test_forward_only_model_on_destructive_change_no_column_types( make_snapshot_on_destructive_change, ): snapshot_a_old, snapshot_a = make_snapshot_on_destructive_change( - old_query="select 1 as one, '2022-01-01' ds", new_query="select one, '2022-01-01' ds" + old_query="select 1 as one, '2022-01-01' ds", + new_query="select one, '2022-01-01' ds", ) context_diff_1 = ContextDiff( @@ -888,15 +902,21 @@ def test_restate_models(sushi_context_pre_scheduling: Context): PlanError, match="Selector did not return any models. Please check your model selection and try again.", ): - sushi_context_pre_scheduling.plan(restate_models=["unknown_model"], no_prompts=True) + sushi_context_pre_scheduling.plan( + restate_models=["unknown_model"], no_prompts=True + ) with pytest.raises( PlanError, match="Selector did not return any models. Please check your model selection and try again.", ): - sushi_context_pre_scheduling.plan(restate_models=["tag:unknown_tag"], no_prompts=True) + sushi_context_pre_scheduling.plan( + restate_models=["tag:unknown_tag"], no_prompts=True + ) - plan = sushi_context_pre_scheduling.plan(restate_models=["raw.demographics"], no_prompts=True) + plan = sushi_context_pre_scheduling.plan( + restate_models=["raw.demographics"], no_prompts=True + ) assert not plan.has_changes assert plan.restatements assert plan.models_to_backfill == { @@ -916,20 +936,26 @@ def test_restate_models(sushi_context_pre_scheduling: Context): @pytest.mark.slow @time_machine.travel(now(minute_floor=False), tick=False) -def test_restate_models_with_existing_missing_intervals(init_and_plan_context: t.Callable): +def test_restate_models_with_existing_missing_intervals( + init_and_plan_context: t.Callable, +): sushi_context, plan = init_and_plan_context("examples/sushi") sushi_context.apply(plan) yesterday_ts = to_timestamp(yesterday_ds()) assert not sushi_context.plan(no_prompts=True).requires_backfill - waiter_revenue_by_day = sushi_context.snapshots['"memory"."sushi"."waiter_revenue_by_day"'] + waiter_revenue_by_day = sushi_context.snapshots[ + '"memory"."sushi"."waiter_revenue_by_day"' + ] sushi_context.state_sync.remove_intervals( [(waiter_revenue_by_day, (yesterday_ts, waiter_revenue_by_day.intervals[0][1]))] ) assert sushi_context.plan(no_prompts=True).requires_backfill - plan = sushi_context.plan(restate_models=["sushi.waiter_revenue_by_day"], no_prompts=True) + plan = sushi_context.plan( + restate_models=["sushi.waiter_revenue_by_day"], no_prompts=True + ) one_day_ms = 24 * 60 * 60 * 1000 @@ -1109,9 +1135,7 @@ def test_end_validation(make_snapshot, mocker: MockerFixture): dev_plan_builder.set_end("2022-01-04") assert dev_plan_builder.build().end == "2022-01-04" - start_end_not_allowed_message = ( - "The start and end dates can't be set for a production plan without restatements." - ) + start_end_not_allowed_message = "The start and end dates can't be set for a production plan without restatements." with pytest.raises(PlanError, match=start_end_not_allowed_message): PlanBuilder(context_diff, end="2022-01-03").build() @@ -1230,7 +1254,9 @@ def test_seed_model_metadata_change_no_missing_intervals( modified_snapshots={ snapshot_a_metadata_updated.name: (snapshot_a_metadata_updated, snapshot_a) }, - snapshots={snapshot_a_metadata_updated.snapshot_id: snapshot_a_metadata_updated}, + snapshots={ + snapshot_a_metadata_updated.snapshot_id: snapshot_a_metadata_updated + }, new_snapshots={snapshot_a_metadata_updated.snapshot_id: snapshot_a}, previous_plan_id=None, previously_promoted_snapshot_ids=set(), @@ -1241,7 +1267,9 @@ def test_seed_model_metadata_change_no_missing_intervals( ) plan = PlanBuilder(context_diff).build() - assert snapshot_a_metadata_updated.change_category == SnapshotChangeCategory.METADATA + assert ( + snapshot_a_metadata_updated.change_category == SnapshotChangeCategory.METADATA + ) assert not snapshot_a_metadata_updated.is_forward_only assert not plan.missing_intervals # plan should have no missing intervals assert ( @@ -1299,7 +1327,9 @@ def test_auto_categorization(make_snapshot, mocker: MockerFixture): snapshot = make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds"))) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - updated_snapshot = make_snapshot(SqlModel(name="a", query=parse_one("select 2, ds"))) + updated_snapshot = make_snapshot( + SqlModel(name="a", query=parse_one("select 2, ds")) + ) context_diff = ContextDiff( environment="test_environment", @@ -1327,10 +1357,14 @@ def test_auto_categorization(make_snapshot, mocker: MockerFixture): assert updated_snapshot.change_category == SnapshotChangeCategory.BREAKING -def test_auto_categorization_missing_schema_downstream(make_snapshot, mocker: MockerFixture): +def test_auto_categorization_missing_schema_downstream( + make_snapshot, mocker: MockerFixture +): snapshot = make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds"))) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - updated_snapshot = make_snapshot(SqlModel(name="a", query=parse_one("select 1, 2, ds"))) + updated_snapshot = make_snapshot( + SqlModel(name="a", query=parse_one("select 1, 2, ds")) + ) # selects * from `tbl` which is not defined and has an unknown schema # therefore we can't be sure what is included in the star select @@ -1355,7 +1389,10 @@ def test_auto_categorization_missing_schema_downstream(make_snapshot, mocker: Mo removed_snapshots={}, modified_snapshots={ updated_snapshot.name: (updated_snapshot, snapshot), - updated_downstream_snapshot.name: (updated_downstream_snapshot, downstream_snapshot), + updated_downstream_snapshot.name: ( + updated_downstream_snapshot, + downstream_snapshot, + ), }, snapshots={ updated_snapshot.snapshot_id: updated_snapshot, @@ -1380,7 +1417,9 @@ def test_broken_references(make_snapshot, mocker: MockerFixture): snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds"))) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) - snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select 2, ds FROM a"))) + snapshot_b = make_snapshot( + SqlModel(name="b", query=parse_one("select 2, ds FROM a")) + ) snapshot_b.categorize_as(SnapshotChangeCategory.BREAKING) context_diff = ContextDiff( @@ -1415,10 +1454,14 @@ def test_broken_references(make_snapshot, mocker: MockerFixture): def test_broken_references_external_model(make_snapshot, mocker: MockerFixture): - snapshot_a = make_snapshot(ExternalModel(name="a", kind=dict(name=ModelKindName.EXTERNAL))) + snapshot_a = make_snapshot( + ExternalModel(name="a", kind=dict(name=ModelKindName.EXTERNAL)) + ) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) - snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select 2, ds FROM a"))) + snapshot_b = make_snapshot( + SqlModel(name="b", query=parse_one("select 2, ds FROM a")) + ) snapshot_b.categorize_as(SnapshotChangeCategory.BREAKING) context_diff = ContextDiff( @@ -1452,7 +1495,10 @@ def test_broken_references_external_model(make_snapshot, mocker: MockerFixture): def test_effective_from(make_snapshot, mocker: MockerFixture): snapshot = make_snapshot( SqlModel( - name="a", query=parse_one("select 1, ds FROM b"), start="2023-01-01", dialect="duckdb" + name="a", + query=parse_one("select 1, ds FROM b"), + start="2023-01-01", + dialect="duckdb", ) ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -1460,7 +1506,10 @@ def test_effective_from(make_snapshot, mocker: MockerFixture): updated_snapshot = make_snapshot( SqlModel( - name="a", query=parse_one("select 2, ds FROM b"), start="2023-01-01", dialect="duckdb" + name="a", + query=parse_one("select 2, ds FROM b"), + start="2023-01-01", + dialect="duckdb", ) ) updated_snapshot.previous_versions = snapshot.all_versions @@ -1609,7 +1658,8 @@ def test_new_environment_no_changes(make_snapshot, mocker: MockerFixture): ) with pytest.raises( - PlanError, match="Creating a new environment requires a change, but project files match.*" + PlanError, + match="Creating a new environment requires a change, but project files match.*", ): PlanBuilder(context_diff, is_dev=True).build() @@ -1625,7 +1675,9 @@ def test_new_environment_no_changes(make_snapshot, mocker: MockerFixture): def test_new_environment_with_changes(make_snapshot, mocker: MockerFixture): snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds"))) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) - updated_snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 3, ds"))) + updated_snapshot_a = make_snapshot( + SqlModel(name="a", query=parse_one("select 3, ds")) + ) snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select 2, ds"))) snapshot_b.categorize_as(SnapshotChangeCategory.BREAKING) @@ -1655,9 +1707,9 @@ def test_new_environment_with_changes(make_snapshot, mocker: MockerFixture): # Modified the existing model. - assert PlanBuilder(context_diff, is_dev=True).build().environment.promoted_snapshot_ids == [ - updated_snapshot_a.snapshot_id - ] + assert PlanBuilder( + context_diff, is_dev=True + ).build().environment.promoted_snapshot_ids == [updated_snapshot_a.snapshot_id] # Updating the existing environment with a previously promoted snapshot. context_diff.previously_promoted_snapshot_ids = { @@ -1666,7 +1718,8 @@ def test_new_environment_with_changes(make_snapshot, mocker: MockerFixture): } context_diff.is_new_environment = False assert set( - PlanBuilder(context_diff, is_dev=True).build().environment.promoted_snapshot_ids or [] + PlanBuilder(context_diff, is_dev=True).build().environment.promoted_snapshot_ids + or [] ) == { updated_snapshot_a.snapshot_id, snapshot_b.snapshot_id, @@ -1685,7 +1738,8 @@ def test_new_environment_with_changes(make_snapshot, mocker: MockerFixture): context_diff.new_snapshots = {snapshot_c.snapshot_id: snapshot_c} assert set( - PlanBuilder(context_diff, is_dev=True).build().environment.promoted_snapshot_ids or [] + PlanBuilder(context_diff, is_dev=True).build().environment.promoted_snapshot_ids + or [] ) == { updated_snapshot_a.snapshot_id, snapshot_b.snapshot_id, @@ -1751,7 +1805,9 @@ def test_forward_only_models(make_snapshot, mocker: MockerFixture): def test_forward_only_models_model_kind_changed(make_snapshot, mocker: MockerFixture): - snapshot = make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds"), kind=FullKind())) + snapshot = make_snapshot( + SqlModel(name="a", query=parse_one("select 1, ds"), kind=FullKind()) + ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) updated_snapshot = make_snapshot( SqlModel( @@ -1844,7 +1900,9 @@ def test_forward_only_models_model_kind_changed_to_incremental_by_time_range( def test_indirectly_modified_forward_only_model(make_snapshot, mocker: MockerFixture): snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 1 as a, ds"))) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) - updated_snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 2 as a, ds"))) + updated_snapshot_a = make_snapshot( + SqlModel(name="a", query=parse_one("select 2 as a, ds")) + ) updated_snapshot_a.previous_versions = snapshot_a.all_versions snapshot_b = make_snapshot( @@ -1857,15 +1915,19 @@ def test_indirectly_modified_forward_only_model(make_snapshot, mocker: MockerFix nodes={'"a"': snapshot_a.model}, ) snapshot_b.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) - updated_snapshot_b = make_snapshot(snapshot_b.model, nodes={'"a"': updated_snapshot_a.model}) + updated_snapshot_b = make_snapshot( + snapshot_b.model, nodes={'"a"': updated_snapshot_a.model} + ) updated_snapshot_b.previous_versions = snapshot_b.all_versions snapshot_c = make_snapshot( - SqlModel(name="c", query=parse_one("select a, ds from b")), nodes={'"b"': snapshot_b.model} + SqlModel(name="c", query=parse_one("select a, ds from b")), + nodes={'"b"': snapshot_b.model}, ) snapshot_c.categorize_as(SnapshotChangeCategory.BREAKING) updated_snapshot_c = make_snapshot( - snapshot_c.model, nodes={'"b"': updated_snapshot_b.model, '"a"': updated_snapshot_a.model} + snapshot_c.model, + nodes={'"b"': updated_snapshot_b.model, '"a"': updated_snapshot_a.model}, ) updated_snapshot_c.previous_versions = snapshot_c.all_versions @@ -1878,7 +1940,8 @@ def test_indirectly_modified_forward_only_model(make_snapshot, mocker: MockerFix ) snapshot_d.categorize_as(SnapshotChangeCategory.BREAKING) updated_snapshot_d = make_snapshot( - snapshot_d.model, nodes={'"b"': updated_snapshot_b.model, '"a"': updated_snapshot_a.model} + snapshot_d.model, + nodes={'"b"': updated_snapshot_b.model, '"a"': updated_snapshot_a.model}, ) updated_snapshot_d.previous_versions = snapshot_d.all_versions @@ -1929,9 +1992,15 @@ def test_indirectly_modified_forward_only_model(make_snapshot, mocker: MockerFix assert plan.directly_modified == {updated_snapshot_a.snapshot_id} assert updated_snapshot_a.change_category == SnapshotChangeCategory.BREAKING - assert updated_snapshot_b.change_category == SnapshotChangeCategory.INDIRECT_BREAKING - assert updated_snapshot_c.change_category == SnapshotChangeCategory.INDIRECT_BREAKING - assert updated_snapshot_d.change_category == SnapshotChangeCategory.INDIRECT_BREAKING + assert ( + updated_snapshot_b.change_category == SnapshotChangeCategory.INDIRECT_BREAKING + ) + assert ( + updated_snapshot_c.change_category == SnapshotChangeCategory.INDIRECT_BREAKING + ) + assert ( + updated_snapshot_d.change_category == SnapshotChangeCategory.INDIRECT_BREAKING + ) assert not updated_snapshot_a.is_forward_only assert updated_snapshot_b.is_forward_only @@ -1954,7 +2023,9 @@ def test_added_model_with_forward_only_parent(make_snapshot, mocker: MockerFixtu snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 1 as a, ds"))) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) - snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select a, ds from a"))) + snapshot_b = make_snapshot( + SqlModel(name="b", query=parse_one("select a, ds from a")) + ) context_diff = ContextDiff( environment="test_environment", @@ -1993,7 +2064,9 @@ def test_added_forward_only_model(make_snapshot, mocker: MockerFixture): ) ) - snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select a, ds from a"))) + snapshot_b = make_snapshot( + SqlModel(name="b", query=parse_one("select a, ds from a")) + ) context_diff = ContextDiff( environment="test_environment", @@ -2060,14 +2133,19 @@ def test_disable_restatement(make_snapshot, mocker: MockerFixture): assert not plan.restatements # Effective from doesn't apply to snapshots for which restatements are disabled. - plan = PlanBuilder(context_diff, forward_only=True, effective_from="2023-01-01").build() + plan = PlanBuilder( + context_diff, forward_only=True, effective_from="2023-01-01" + ).build() assert plan.effective_from == "2023-01-01" assert snapshot.effective_from is None # Restatements should still be supported when in dev. plan = PlanBuilder(context_diff, is_dev=True, restate_models=['"a"']).build() assert plan.restatements == { - snapshot.snapshot_id: (to_timestamp(plan.start), to_timestamp(to_date("tomorrow"))) + snapshot.snapshot_id: ( + to_timestamp(plan.start), + to_timestamp(to_date("tomorrow")), + ) } # We don't want to restate a disable_restatement model if it is unpaused since that would be mean we are violating @@ -2397,7 +2475,11 @@ def test_dev_plan_depends_past(make_snapshot, mocker: MockerFixture): normalize_environment_name=True, create_from="prod", create_from_env_exists=True, - added={snapshot.snapshot_id, snapshot_child.snapshot_id, unrelated_snapshot.snapshot_id}, + added={ + snapshot.snapshot_id, + snapshot_child.snapshot_id, + unrelated_snapshot.snapshot_id, + }, removed_snapshots={}, modified_snapshots={}, snapshots={ @@ -2439,8 +2521,12 @@ def test_dev_plan_depends_past(make_snapshot, mocker: MockerFixture): ).build() assert len(dev_plan_start_ahead_of_model.new_snapshots) == 3 assert not dev_plan_start_ahead_of_model.deployability_index.is_deployable(snapshot) - assert not dev_plan_start_ahead_of_model.deployability_index.is_deployable(snapshot_child) - assert dev_plan_start_ahead_of_model.deployability_index.is_deployable(unrelated_snapshot) + assert not dev_plan_start_ahead_of_model.deployability_index.is_deployable( + snapshot_child + ) + assert dev_plan_start_ahead_of_model.deployability_index.is_deployable( + unrelated_snapshot + ) assert dev_plan_start_ahead_of_model.directly_modified == { snapshot.snapshot_id, snapshot_child.snapshot_id, @@ -2539,7 +2625,9 @@ def new_builder(start, end): def test_restatement_intervals_after_updating_start(sushi_context: Context): - plan = sushi_context.plan(no_prompts=True, restate_models=["sushi.waiter_revenue_by_day"]) + plan = sushi_context.plan( + no_prompts=True, restate_models=["sushi.waiter_revenue_by_day"] + ) snapshot_id = [ snapshot.snapshot_id for snapshot in plan.snapshots @@ -2558,7 +2646,9 @@ def test_restatement_intervals_after_updating_start(sushi_context: Context): def test_models_selected_for_backfill(make_snapshot, mocker: MockerFixture): - snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 1 as one, ds"))) + snapshot_a = make_snapshot( + SqlModel(name="a", query=parse_one("select 1 as one, ds")) + ) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) snapshot_b = make_snapshot( @@ -2657,19 +2747,29 @@ def test_categorized_uncategorized(make_snapshot, mocker: MockerFixture): def test_environment_previous_finalized_snapshots(make_snapshot, mocker: MockerFixture): - snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 1 as one, ds"))) + snapshot_a = make_snapshot( + SqlModel(name="a", query=parse_one("select 1 as one, ds")) + ) snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) - updated_snapshot_a = make_snapshot(SqlModel(name="a", query=parse_one("select 4 as four, ds"))) + updated_snapshot_a = make_snapshot( + SqlModel(name="a", query=parse_one("select 4 as four, ds")) + ) updated_snapshot_a.categorize_as(SnapshotChangeCategory.BREAKING) - snapshot_b = make_snapshot(SqlModel(name="b", query=parse_one("select 2 as two, ds"))) + snapshot_b = make_snapshot( + SqlModel(name="b", query=parse_one("select 2 as two, ds")) + ) snapshot_b.categorize_as(SnapshotChangeCategory.BREAKING) - snapshot_c = make_snapshot(SqlModel(name="c", query=parse_one("select 3 as three, ds"))) + snapshot_c = make_snapshot( + SqlModel(name="c", query=parse_one("select 3 as three, ds")) + ) snapshot_c.categorize_as(SnapshotChangeCategory.BREAKING) - snapshot_d = make_snapshot(SqlModel(name="d", query=parse_one("select 5 as five, ds"))) + snapshot_d = make_snapshot( + SqlModel(name="d", query=parse_one("select 5 as five, ds")) + ) snapshot_d.categorize_as(SnapshotChangeCategory.BREAKING) context_diff = ContextDiff( @@ -2770,7 +2870,9 @@ def test_plan_start_when_preview_enabled(make_snapshot, mocker: MockerFixture): query=parse_one("select 1, ds"), dialect="duckdb", kind=dict( - name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, time_column="ds", forward_only=True + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + time_column="ds", + forward_only=True, ), start=model_start, ) @@ -2852,7 +2954,9 @@ def test_end_override_per_model(make_snapshot): context_diff, end_override_per_model={snapshot.name: to_datetime("2023-01-09")}, ) - assert plan_builder.build().end_override_per_model == {snapshot.name: to_datetime("2023-01-09")} + assert plan_builder.build().end_override_per_model == { + snapshot.name: to_datetime("2023-01-09") + } # User-provided end should take precedence. plan_builder = PlanBuilder( @@ -2870,7 +2974,9 @@ def test_unaligned_start_model_with_forward_only_preview(make_snapshot): name="a", query=parse_one("select 1, ds"), kind=dict( - name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True, time_column="ds" + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, + time_column="ds", ), ) ) @@ -2881,7 +2987,9 @@ def test_unaligned_start_model_with_forward_only_preview(make_snapshot): name="a", query=parse_one("select 2, ds"), kind=dict( - name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True, time_column="ds" + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, + time_column="ds", ), ) ) @@ -2908,7 +3016,10 @@ def test_unaligned_start_model_with_forward_only_preview(make_snapshot): create_from_env_exists=True, added={snapshot_b.snapshot_id}, removed_snapshots={}, - snapshots={new_snapshot_a.snapshot_id: new_snapshot_a, snapshot_b.snapshot_id: snapshot_b}, + snapshots={ + new_snapshot_a.snapshot_id: new_snapshot_a, + snapshot_b.snapshot_id: snapshot_b, + }, new_snapshots={ new_snapshot_a.snapshot_id: new_snapshot_a, snapshot_b.snapshot_id: snapshot_b, @@ -2929,7 +3040,10 @@ def test_unaligned_start_model_with_forward_only_preview(make_snapshot): ) plan = plan_builder.build() - assert set(plan.restatements) == {new_snapshot_a.snapshot_id, snapshot_b.snapshot_id} + assert set(plan.restatements) == { + new_snapshot_a.snapshot_id, + snapshot_b.snapshot_id, + } assert not plan.deployability_index.is_deployable(new_snapshot_a) assert not plan.deployability_index.is_deployable(snapshot_b) @@ -2986,7 +3100,9 @@ def _make_forward_only_preview_context_diff(make_snapshot): return context_diff, new_snapshot -def test_forward_only_preview_start_does_not_override_implicit_backfill_start(make_snapshot): +def test_forward_only_preview_start_does_not_override_implicit_backfill_start( + make_snapshot, +): context_diff, new_snapshot = _make_forward_only_preview_context_diff(make_snapshot) normal_old_snapshot = make_snapshot( @@ -3171,7 +3287,10 @@ def test_restate_production_model_in_dev(make_snapshot, mocker: MockerFixture): added=set(), removed_snapshots={}, modified_snapshots={}, - snapshots={snapshot.snapshot_id: snapshot, prod_snapshot.snapshot_id: prod_snapshot}, + snapshots={ + snapshot.snapshot_id: snapshot, + prod_snapshot.snapshot_id: prod_snapshot, + }, new_snapshots={}, previous_plan_id=None, previously_promoted_snapshot_ids=set(), @@ -3327,7 +3446,10 @@ def test_plan_environment_statements_diff(make_snapshot): previous_finalized_snapshots=None, environment_statements=[ EnvironmentStatements( - before_all=["CREATE OR REPLACE TABLE table_1 AS SELECT 1", "@test_macro()"], + before_all=[ + "CREATE OR REPLACE TABLE table_1 AS SELECT 1", + "@test_macro()", + ], after_all=["CREATE OR REPLACE TABLE table_2 AS SELECT 2"], python_env={ "test_macro": Executable( @@ -4206,7 +4328,9 @@ def test_plan_builder_allow_additive_models_pattern_matching(make_snapshot): builder_with_pattern = PlanBuilder( context_diff, forward_only=True, - allow_additive_models={'"test"."model_1"'}, # Only allow test.model_1, not other.model_2 + allow_additive_models={ + '"test"."model_1"' + }, # Only allow test.model_1, not other.model_2 ) # Should still fail because other.model_2 is not allowed @@ -4217,7 +4341,10 @@ def test_plan_builder_allow_additive_models_pattern_matching(make_snapshot): builder_with_both = PlanBuilder( context_diff, forward_only=True, - allow_additive_models={'"test"."model_1"', '"other"."model_2"'}, # Allow both models + allow_additive_models={ + '"test"."model_1"', + '"other"."model_2"', + }, # Allow both models ) # Should succeed @@ -4263,7 +4390,9 @@ def test_environment_statements_change_allows_dev_environment_creation(make_snap is_dev=True, ) - with pytest.raises(NoChangesPlanError, match="Creating a new environment requires a change"): + with pytest.raises( + NoChangesPlanError, match="Creating a new environment requires a change" + ): plan_builder.build() # Now create context diff with environment statements @@ -4508,4 +4637,6 @@ def test_forward_only_indirect_change_to_materialized_view(make_snapshot): # Forward-only indirect changes to MVs should not always be classified as indirect breaking. # Instead, we want to preserve the standard categorization. - assert snapshot_b_new.change_category == SnapshotChangeCategory.INDIRECT_NON_BREAKING + assert ( + snapshot_b_new.change_category == SnapshotChangeCategory.INDIRECT_NON_BREAKING + ) diff --git a/tests/core/test_plan_evaluator.py b/tests/core/test_plan_evaluator.py index 575f5ae742..6c8da5e22b 100644 --- a/tests/core/test_plan_evaluator.py +++ b/tests/core/test_plan_evaluator.py @@ -4,12 +4,8 @@ from sqlmesh.core.context import Context from sqlmesh.core.model import FullKind, SqlModel, ViewKind -from sqlmesh.core.plan import ( - BuiltInPlanEvaluator, - Plan, - PlanBuilder, - stages as plan_stages, -) +from sqlmesh.core.plan import BuiltInPlanEvaluator, Plan, PlanBuilder +from sqlmesh.core.plan import stages as plan_stages from sqlmesh.core.snapshot import SnapshotChangeCategory @@ -42,7 +38,9 @@ def test_builtin_evaluator_push(sushi_context: Context, make_snapshot): kind=ViewKind(), owner="jen", start="2020-01-01", - query=parse_one("SELECT 1::INT AS one FROM sushi.new_test_model, sushi.waiters"), + query=parse_one( + "SELECT 1::INT AS one FROM sushi.new_test_model, sushi.waiters" + ), default_catalog="memory", ) @@ -50,7 +48,9 @@ def test_builtin_evaluator_push(sushi_context: Context, make_snapshot): sushi_context.upsert_model(new_view_model) new_model_snapshot = sushi_context.get_snapshot(new_model, raise_if_missing=True) - new_view_model_snapshot = sushi_context.get_snapshot(new_view_model, raise_if_missing=True) + new_view_model_snapshot = sushi_context.get_snapshot( + new_view_model, raise_if_missing=True + ) new_model_snapshot.categorize_as(SnapshotChangeCategory.BREAKING) new_view_model_snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -77,8 +77,14 @@ def test_builtin_evaluator_push(sushi_context: Context, make_snapshot): evaluator.visit_backfill_stage(stages[3], evaluatable_plan) assert ( - len(sushi_context.state_sync.get_snapshots([new_model_snapshot, new_view_model_snapshot])) + len( + sushi_context.state_sync.get_snapshots( + [new_model_snapshot, new_view_model_snapshot] + ) + ) == 2 ) assert sushi_context.engine_adapter.table_exists(new_model_snapshot.table_name()) - assert sushi_context.engine_adapter.table_exists(new_view_model_snapshot.table_name()) + assert sushi_context.engine_adapter.table_exists( + new_view_model_snapshot.table_name() + ) diff --git a/tests/core/test_plan_stages.py b/tests/core/test_plan_stages.py index eb3f965761..edca010dec 100644 --- a/tests/core/test_plan_stages.py +++ b/tests/core/test_plan_stages.py @@ -1,39 +1,31 @@ -import pytest import typing as t -from sqlglot import parse_one + +import pytest from pytest_mock.plugin import MockerFixture +from sqlglot import parse_one from sqlmesh.core.config import EnvironmentSuffixTarget from sqlmesh.core.config.common import VirtualEnvironmentMode -from sqlmesh.core.model import SqlModel, ModelKindName +from sqlmesh.core.environment import Environment, EnvironmentStatements +from sqlmesh.core.model import ModelKindName, SqlModel from sqlmesh.core.plan.common import SnapshotIntervalClearRequest from sqlmesh.core.plan.definition import EvaluatablePlan -from sqlmesh.core.plan.stages import ( - build_plan_stages, - AfterAllStage, - AuditOnlyRunStage, - PhysicalLayerUpdateStage, - PhysicalLayerSchemaCreationStage, - CreateSnapshotRecordsStage, - BeforeAllStage, - BackfillStage, - EnvironmentRecordUpdateStage, - VirtualLayerUpdateStage, - RestatementStage, - MigrateSchemasStage, - FinalizeEnvironmentStage, - UnpauseStage, -) from sqlmesh.core.plan.explainer import ExplainableRestatementStage -from sqlmesh.core.snapshot.definition import ( - SnapshotChangeCategory, - DeployabilityIndex, - Snapshot, - SnapshotId, - SnapshotIdLike, -) +from sqlmesh.core.plan.stages import (AfterAllStage, AuditOnlyRunStage, + BackfillStage, BeforeAllStage, + CreateSnapshotRecordsStage, + EnvironmentRecordUpdateStage, + FinalizeEnvironmentStage, + MigrateSchemasStage, + PhysicalLayerSchemaCreationStage, + PhysicalLayerUpdateStage, + RestatementStage, UnpauseStage, + VirtualLayerUpdateStage, + build_plan_stages) +from sqlmesh.core.snapshot.definition import (DeployabilityIndex, Snapshot, + SnapshotChangeCategory, + SnapshotId, SnapshotIdLike) from sqlmesh.core.state_sync import StateReader -from sqlmesh.core.environment import Environment, EnvironmentStatements from sqlmesh.utils.date import to_timestamp @@ -165,14 +157,20 @@ def test_build_plan_stages_basic( # Verify UnpauseStage assert isinstance(stages[4], UnpauseStage) - assert {s.name for s in stages[4].promoted_snapshots} == {snapshot_a.name, snapshot_b.name} + assert {s.name for s in stages[4].promoted_snapshots} == { + snapshot_a.name, + snapshot_b.name, + } # Verify VirtualLayerUpdateStage virtual_stage = stages[5] assert isinstance(virtual_stage, VirtualLayerUpdateStage) assert len(virtual_stage.promoted_snapshots) == 2 assert len(virtual_stage.demoted_snapshots) == 0 - assert {s.name for s in virtual_stage.promoted_snapshots} == {snapshot_a.name, snapshot_b.name} + assert {s.name for s in virtual_stage.promoted_snapshots} == { + snapshot_a.name, + snapshot_b.name, + } state_reader.refresh_snapshot_intervals.assert_called_once() @@ -281,7 +279,10 @@ def test_build_plan_stages_with_before_all_and_after_all( # Verify UnpauseStage assert isinstance(stages[5], UnpauseStage) - assert {s.name for s in stages[5].promoted_snapshots} == {snapshot_a.name, snapshot_b.name} + assert {s.name for s in stages[5].promoted_snapshots} == { + snapshot_a.name, + snapshot_b.name, + } # Verify VirtualLayerUpdateStage virtual_stage = stages[6] @@ -493,7 +494,10 @@ def test_build_plan_stages_basic_no_backfill( # Verify UnpauseStage assert isinstance(stages[5], UnpauseStage) - assert {s.name for s in stages[5].promoted_snapshots} == {snapshot_a.name, snapshot_b.name} + assert {s.name for s in stages[5].promoted_snapshots} == { + snapshot_a.name, + snapshot_b.name, + } # Verify VirtualLayerUpdateStage virtual_stage = stages[6] @@ -606,7 +610,9 @@ def test_build_plan_stages_restatement_prod_only( assert isinstance(backfill_stage, BackfillStage) assert len(backfill_stage.snapshot_to_intervals) == 2 assert backfill_stage.deployability_index == DeployabilityIndex.all_deployable() - expected_backfill_interval = [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))] + expected_backfill_interval = [ + (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) + ] for intervals in backfill_stage.snapshot_to_intervals.values(): assert intervals == expected_backfill_interval @@ -763,7 +769,9 @@ def _get_snapshots(snapshot_ids: t.Iterable[SnapshotIdLike]): assert isinstance(backfill_stage, BackfillStage) assert len(backfill_stage.snapshot_to_intervals) == 2 assert backfill_stage.deployability_index == DeployabilityIndex.all_deployable() - expected_backfill_interval = [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))] + expected_backfill_interval = [ + (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) + ] for intervals in backfill_stage.snapshot_to_intervals.values(): assert intervals == expected_backfill_interval @@ -777,14 +785,19 @@ def _get_snapshots(snapshot_ids: t.Iterable[SnapshotIdLike]): # note: we only clear the intervals from state for "a" in dev, we leave prod alone assert restatement_stage.snapshot_intervals_to_clear assert len(restatement_stage.snapshot_intervals_to_clear) == 1 - snapshot_name, clear_requests = list(restatement_stage.snapshot_intervals_to_clear.items())[0] + snapshot_name, clear_requests = list( + restatement_stage.snapshot_intervals_to_clear.items() + )[0] assert snapshot_name == '"a"' assert len(clear_requests) == 1 clear_request = clear_requests[0] assert isinstance(clear_request, SnapshotIntervalClearRequest) assert clear_request.snapshot_id == snapshot_a_dev.snapshot_id assert clear_request.snapshot == snapshot_a_dev.id_and_version - assert clear_request.interval == (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) + assert clear_request.interval == ( + to_timestamp("2023-01-01"), + to_timestamp("2023-01-02"), + ) # Verify EnvironmentRecordUpdateStage assert isinstance(stages[3], EnvironmentRecordUpdateStage) @@ -929,9 +942,13 @@ def test_build_plan_stages_restatement_dev_does_not_clear_intervals( assert isinstance(backfill_stage, BackfillStage) assert len(backfill_stage.snapshot_to_intervals) == 1 assert backfill_stage.deployability_index == DeployabilityIndex.all_deployable() - backfill_snapshot, backfill_intervals = list(backfill_stage.snapshot_to_intervals.items())[0] + backfill_snapshot, backfill_intervals = list( + backfill_stage.snapshot_to_intervals.items() + )[0] assert backfill_snapshot.snapshot_id == snapshot_a_dev.snapshot_id - assert backfill_intervals == [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))] + assert backfill_intervals == [ + (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) + ] # Verify EnvironmentRecordUpdateStage assert isinstance(stages[2], EnvironmentRecordUpdateStage) @@ -958,7 +975,9 @@ def test_build_plan_stages_forward_only( nodes={'"a"': new_snapshot_a.model}, ) new_snapshot_b.previous_versions = snapshot_b.all_versions - new_snapshot_b.categorize_as(SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only=True) + new_snapshot_b.categorize_as( + SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only=True + ) state_reader = mocker.Mock(spec=StateReader) state_reader.get_snapshots.return_value = {} @@ -1097,7 +1116,9 @@ def test_build_plan_stages_forward_only_dev( nodes={'"a"': new_snapshot_a.model}, ) new_snapshot_b.previous_versions = snapshot_b.all_versions - new_snapshot_b.categorize_as(SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only=True) + new_snapshot_b.categorize_as( + SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only=True + ) state_reader = mocker.Mock(spec=StateReader) state_reader.get_snapshots.return_value = {} @@ -1217,8 +1238,13 @@ def test_build_plan_stages_audit_only( new_snapshot_b.categorize_as(SnapshotChangeCategory.METADATA) new_snapshot_b.add_interval("2023-01-01", "2023-01-02") - def _get_snapshots(snapshot_ids: t.List[SnapshotId]) -> t.Dict[SnapshotId, Snapshot]: - if snapshot_a.snapshot_id in snapshot_ids and snapshot_b.snapshot_id in snapshot_ids: + def _get_snapshots( + snapshot_ids: t.List[SnapshotId], + ) -> t.Dict[SnapshotId, Snapshot]: + if ( + snapshot_a.snapshot_id in snapshot_ids + and snapshot_b.snapshot_id in snapshot_ids + ): return { snapshot_a.snapshot_id: snapshot_a, snapshot_b.snapshot_id: snapshot_b, @@ -1351,7 +1377,9 @@ def test_build_plan_stages_forward_only_ensure_finalized_snapshots( nodes={'"a"': new_snapshot_a.model}, ) new_snapshot_b.previous_versions = snapshot_b.all_versions - new_snapshot_b.categorize_as(SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only=True) + new_snapshot_b.categorize_as( + SnapshotChangeCategory.INDIRECT_NON_BREAKING, forward_only=True + ) state_reader = mocker.Mock(spec=StateReader) state_reader.get_snapshots.return_value = {} @@ -1725,7 +1753,10 @@ def test_build_plan_stages_virtual_environment_mode_filtering( end_at="2023-01-02", plan_id="test_plan", previous_plan_id=None, - promoted_snapshot_ids=[snapshot_full.snapshot_id, snapshot_dev_only.snapshot_id], + promoted_snapshot_ids=[ + snapshot_full.snapshot_id, + snapshot_dev_only.snapshot_id, + ], ) plan_dev = EvaluatablePlan( @@ -1745,7 +1776,10 @@ def test_build_plan_stages_virtual_environment_mode_filtering( end_bounded=False, ensure_finalized_snapshots=False, ignore_cron=False, - directly_modified_snapshots=[snapshot_full.snapshot_id, snapshot_dev_only.snapshot_id], + directly_modified_snapshots=[ + snapshot_full.snapshot_id, + snapshot_dev_only.snapshot_id, + ], indirectly_modified_snapshots={}, metadata_updated_snapshots=[], removed_snapshots=[], @@ -1779,7 +1813,10 @@ def test_build_plan_stages_virtual_environment_mode_filtering( end_at="2023-01-02", plan_id="test_plan", previous_plan_id=None, - promoted_snapshot_ids=[snapshot_full.snapshot_id, snapshot_dev_only.snapshot_id], + promoted_snapshot_ids=[ + snapshot_full.snapshot_id, + snapshot_dev_only.snapshot_id, + ], ) plan_prod = EvaluatablePlan( @@ -1799,7 +1836,10 @@ def test_build_plan_stages_virtual_environment_mode_filtering( end_bounded=False, ensure_finalized_snapshots=False, ignore_cron=False, - directly_modified_snapshots=[snapshot_full.snapshot_id, snapshot_dev_only.snapshot_id], + directly_modified_snapshots=[ + snapshot_full.snapshot_id, + snapshot_dev_only.snapshot_id, + ], indirectly_modified_snapshots={}, metadata_updated_snapshots=[], removed_snapshots=[], @@ -1830,7 +1870,10 @@ def test_build_plan_stages_virtual_environment_mode_filtering( end_at="2023-01-02", plan_id="previous_plan", previous_plan_id=None, - promoted_snapshot_ids=[snapshot_full.snapshot_id, snapshot_dev_only.snapshot_id], + promoted_snapshot_ids=[ + snapshot_full.snapshot_id, + snapshot_dev_only.snapshot_id, + ], finalized_ts=to_timestamp("2023-01-02"), ) state_reader.get_environment.return_value = existing_environment @@ -1879,12 +1922,16 @@ def test_build_plan_stages_virtual_environment_mode_filtering( # Find VirtualLayerUpdateStage virtual_stage_prod_demote = next( - stage for stage in stages_prod_demote if isinstance(stage, VirtualLayerUpdateStage) + stage + for stage in stages_prod_demote + if isinstance(stage, VirtualLayerUpdateStage) ) # In production environment, only FULL mode snapshots should be demoted assert len(virtual_stage_prod_demote.promoted_snapshots) == 0 - assert {s.name for s in virtual_stage_prod_demote.demoted_snapshots} == {'"full_model"'} + assert {s.name for s in virtual_stage_prod_demote.demoted_snapshots} == { + '"full_model"' + } assert ( virtual_stage_prod_demote.demoted_environment_naming_info == existing_environment.naming_info @@ -1953,7 +2000,9 @@ def test_build_plan_stages_virtual_environment_mode_no_updates( stages = build_plan_stages(plan, state_reader, None) # No VirtualLayerUpdateStage should be created since all snapshots are filtered out - virtual_stages = [stage for stage in stages if isinstance(stage, VirtualLayerUpdateStage)] + virtual_stages = [ + stage for stage in stages if isinstance(stage, VirtualLayerUpdateStage) + ] assert len(virtual_stages) == 0 @@ -1967,8 +2016,12 @@ def test_adjust_intervals_new_forward_only_dev_intervals( kind=dict(name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, time_column="ds"), ) ) - forward_only_snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) - forward_only_snapshot.intervals = [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))] + forward_only_snapshot.categorize_as( + SnapshotChangeCategory.BREAKING, forward_only=True + ) + forward_only_snapshot.intervals = [ + (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) + ] forward_only_snapshot.dev_intervals = [] @@ -1989,7 +2042,9 @@ def test_adjust_intervals_new_forward_only_dev_intervals( plan = EvaluatablePlan( start="2023-01-01", end="2023-01-02", - new_snapshots=[forward_only_snapshot], # This snapshot should have dev_intervals set + new_snapshots=[ + forward_only_snapshot + ], # This snapshot should have dev_intervals set environment=environment, no_gaps=False, skip_backfill=False, @@ -2092,17 +2147,23 @@ def test_adjust_intervals_restatement_removal( state_reader.refresh_snapshot_intervals.assert_called_once() - restatement_stages = [stage for stage in stages if isinstance(stage, RestatementStage)] + restatement_stages = [ + stage for stage in stages if isinstance(stage, RestatementStage) + ] assert len(restatement_stages) == 1 backfill_stages = [stage for stage in stages if isinstance(stage, BackfillStage)] assert len(backfill_stages) == 1 - (snapshot, intervals) = next(iter(backfill_stages[0].snapshot_to_intervals.items())) - assert snapshot.intervals == [(to_timestamp("2023-01-02"), to_timestamp("2023-01-04"))] + snapshot, intervals = next(iter(backfill_stages[0].snapshot_to_intervals.items())) + assert snapshot.intervals == [ + (to_timestamp("2023-01-02"), to_timestamp("2023-01-04")) + ] assert intervals == [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))] -def test_adjust_intervals_should_force_rebuild(make_snapshot, mocker: MockerFixture) -> None: +def test_adjust_intervals_should_force_rebuild( + make_snapshot, mocker: MockerFixture +) -> None: old_snapshot = make_snapshot( SqlModel( name="test_model", @@ -2126,7 +2187,12 @@ def test_adjust_intervals_should_force_rebuild(make_snapshot, mocker: MockerFixt state_reader = mocker.Mock(spec=StateReader) state_reader.refresh_snapshot_intervals = mocker.Mock() - state_reader.get_snapshots.side_effect = [{}, {old_snapshot.snapshot_id: old_snapshot}, {}, {}] + state_reader.get_snapshots.side_effect = [ + {}, + {old_snapshot.snapshot_id: old_snapshot}, + {}, + {}, + ] existing_environment = Environment( name="prod", @@ -2185,6 +2251,6 @@ def test_adjust_intervals_should_force_rebuild(make_snapshot, mocker: MockerFixt assert not new_snapshot.intervals backfill_stages = [stage for stage in stages if isinstance(stage, BackfillStage)] assert len(backfill_stages) == 1 - (snapshot, intervals) = next(iter(backfill_stages[0].snapshot_to_intervals.items())) + snapshot, intervals = next(iter(backfill_stages[0].snapshot_to_intervals.items())) assert not snapshot.intervals assert intervals == [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))] diff --git a/tests/core/test_reference.py b/tests/core/test_reference.py index fe08f3f71e..3688e3c567 100644 --- a/tests/core/test_reference.py +++ b/tests/core/test_reference.py @@ -26,7 +26,9 @@ def make(name, refs=None): def test_graph(make_model): graph = ReferenceGraph( [ - make_model("model_a", [("a", True), ("b", False), ("c", False), ("d", False)]), + make_model( + "model_a", [("a", True), ("b", False), ("c", False), ("d", False)] + ), make_model("model_b", [("a", True)]), make_model("model_c", [("d", True), ("e", True)]), make_model("model_d", [("e", True)]), @@ -38,7 +40,11 @@ def find_path(a, b): return [(r.model_name, r.name) for r in graph.find_path(a, b)] assert find_path("model_a", "model_b") == [("model_a", "a"), ("model_b", "a")] - assert find_path("model_a", "model_d") == [("model_a", "d"), ("model_c", "e"), ("model_d", "e")] + assert find_path("model_a", "model_d") == [ + ("model_a", "d"), + ("model_c", "e"), + ("model_d", "e"), + ] with pytest.raises(SQLMeshError): assert find_path("model_a", "model_e") diff --git a/tests/core/test_rule.py b/tests/core/test_rule.py index 785988932d..f8ee1ce56d 100644 --- a/tests/core/test_rule.py +++ b/tests/core/test_rule.py @@ -5,8 +5,9 @@ from unittest.mock import MagicMock import pytest -from sqlmesh.core.model import Model + from sqlmesh.core.linter.rule import Rule, RuleViolation +from sqlmesh.core.model import Model class TestRule(Rule): diff --git a/tests/core/test_scheduler.py b/tests/core/test_scheduler.py index cd32d2451d..9ddd7b6131 100644 --- a/tests/core/test_scheduler.py +++ b/tests/core/test_scheduler.py @@ -2,7 +2,7 @@ import pytest from pytest_mock.plugin import MockerFixture -from sqlglot import parse_one, parse +from sqlglot import parse, parse_one from sqlglot.helper import first from sqlmesh.core.context import Context, ExecutionContext @@ -10,31 +10,19 @@ from sqlmesh.core.macros import RuntimeStage from sqlmesh.core.model import load_sql_based_model from sqlmesh.core.model.definition import AuditResult, SqlModel -from sqlmesh.core.model.kind import ( - IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind, - TimeColumn, - SCDType2ByColumnKind, -) +from sqlmesh.core.model.kind import (IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, + SCDType2ByColumnKind, TimeColumn) from sqlmesh.core.node import IntervalUnit -from sqlmesh.core.scheduler import ( - Scheduler, - interval_diff, - compute_interval_params, - SnapshotToIntervals, - EvaluateNode, - SchedulingUnit, - DummyNode, -) +from sqlmesh.core.scheduler import (DummyNode, EvaluateNode, Scheduler, + SchedulingUnit, SnapshotToIntervals, + compute_interval_params, interval_diff) from sqlmesh.core.signal import signal -from sqlmesh.core.snapshot import ( - Snapshot, - SnapshotEvaluator, - SnapshotChangeCategory, - DeployabilityIndex, - snapshots_to_dag, -) -from sqlmesh.utils.date import to_datetime, to_timestamp, DatetimeRanges, TimeLike +from sqlmesh.core.snapshot import (DeployabilityIndex, Snapshot, + SnapshotChangeCategory, SnapshotEvaluator, + snapshots_to_dag) +from sqlmesh.utils.date import (DatetimeRanges, TimeLike, to_datetime, + to_timestamp) from sqlmesh.utils.errors import CircuitBreakerError, NodeAuditsErrors @@ -50,18 +38,24 @@ def orders(sushi_context_fixed_date: Context) -> Snapshot: @pytest.fixture def waiter_names(sushi_context_fixed_date: Context) -> Snapshot: - return sushi_context_fixed_date.get_snapshot("sushi.waiter_names", raise_if_missing=True) + return sushi_context_fixed_date.get_snapshot( + "sushi.waiter_names", raise_if_missing=True + ) @pytest.mark.slow -def test_interval_params(scheduler: Scheduler, sushi_context_fixed_date: Context, orders: Snapshot): +def test_interval_params( + scheduler: Scheduler, sushi_context_fixed_date: Context, orders: Snapshot +): waiter_revenue = sushi_context_fixed_date.get_snapshot( "sushi.waiter_revenue_by_day", raise_if_missing=True ) start_ds = "2022-01-01" end_ds = "2022-02-05" - assert compute_interval_params([orders, waiter_revenue], start=start_ds, end=end_ds) == { + assert compute_interval_params( + [orders, waiter_revenue], start=start_ds, end=end_ds + ) == { orders: [ (to_timestamp(start_ds), to_timestamp("2022-02-06")), ], @@ -74,14 +68,18 @@ def test_interval_params(scheduler: Scheduler, sushi_context_fixed_date: Context @pytest.fixture def get_batched_missing_intervals( mocker: MockerFixture, -) -> t.Callable[[Scheduler, TimeLike, TimeLike, t.Optional[TimeLike]], SnapshotToIntervals]: +) -> t.Callable[ + [Scheduler, TimeLike, TimeLike, t.Optional[TimeLike]], SnapshotToIntervals +]: def _get_batched_missing_intervals( scheduler: Scheduler, start: TimeLike, end: TimeLike, execution_time: t.Optional[TimeLike] = None, ) -> SnapshotToIntervals: - merged_intervals = scheduler.merged_missing_intervals(start, end, execution_time) + merged_intervals = scheduler.merged_missing_intervals( + start, end, execution_time + ) return scheduler.batch_intervals(merged_intervals, mocker.Mock(), mocker.Mock()) return _get_batched_missing_intervals @@ -102,7 +100,9 @@ def test_interval_params_nonconsecutive(scheduler: Scheduler, orders: Snapshot): @pytest.mark.slow -def test_interval_params_missing(scheduler: Scheduler, sushi_context_fixed_date: Context): +def test_interval_params_missing( + scheduler: Scheduler, sushi_context_fixed_date: Context +): waiters = sushi_context_fixed_date.get_snapshot( "sushi.waiter_as_customer_by_day", raise_if_missing=True ) @@ -119,7 +119,9 @@ def test_interval_params_missing(scheduler: Scheduler, sushi_context_fixed_date: @pytest.mark.slow def test_run(sushi_context_fixed_date: Context, scheduler: Scheduler): adapter = sushi_context_fixed_date.engine_adapter - snapshot = sushi_context_fixed_date.get_snapshot("sushi.items", raise_if_missing=True) + snapshot = sushi_context_fixed_date.get_snapshot( + "sushi.items", raise_if_missing=True + ) scheduler.run( EnvironmentNamingInfo(), "2022-01-01", @@ -127,11 +129,9 @@ def test_run(sushi_context_fixed_date: Context, scheduler: Scheduler): "2022-01-30", ) - assert adapter.fetchone( - f""" + assert adapter.fetchone(f""" SELECT id, name, price FROM sqlmesh__sushi.sushi__items__{snapshot.version} ORDER BY event_date LIMIT 1 - """ - ) == (0, "Hotate", 5.99) + """) == (0, "Hotate", 5.99) def test_incremental_by_unique_key_kind_dag( @@ -153,7 +153,9 @@ def test_incremental_by_unique_key_kind_dag( query=parse_one("SELECT id FROM VALUES (1), (2) AS t(id)"), ), ) - snapshot_evaluator = SnapshotEvaluator(adapters=mocker.MagicMock(), ddl_concurrent_tasks=1) + snapshot_evaluator = SnapshotEvaluator( + adapters=mocker.MagicMock(), ddl_concurrent_tasks=1 + ) mock_state_sync = mocker.MagicMock() scheduler = Scheduler( snapshots=[unique_by_key_snapshot], @@ -184,7 +186,9 @@ def test_incremental_time_self_reference_dag( incremental_self_snapshot: Snapshot = make_snapshot( SqlModel( name="name", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds"), batch_size=1), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds"), batch_size=1 + ), owner="owner", dialect="", cron="@daily", @@ -195,7 +199,9 @@ def test_incremental_time_self_reference_dag( incremental_self_snapshot.add_interval("2023-01-02", "2023-01-02") incremental_self_snapshot.add_interval("2023-01-05", "2023-01-05") - snapshot_evaluator = SnapshotEvaluator(adapters=mocker.MagicMock(), ddl_concurrent_tasks=1) + snapshot_evaluator = SnapshotEvaluator( + adapters=mocker.MagicMock(), ddl_concurrent_tasks=1 + ) scheduler = Scheduler( snapshots=[incremental_self_snapshot], snapshot_evaluator=snapshot_evaluator, @@ -295,7 +301,10 @@ def test_incremental_time_self_reference_dag( ): { EvaluateNode( snapshot_name='"test_model"', - interval=(to_timestamp("2023-01-01"), to_timestamp("2023-01-03")), + interval=( + to_timestamp("2023-01-01"), + to_timestamp("2023-01-03"), + ), batch_index=0, ), }, @@ -327,7 +336,10 @@ def test_incremental_time_self_reference_dag( ): { EvaluateNode( snapshot_name='"test_model"', - interval=(to_timestamp("2023-01-01"), to_timestamp("2023-01-02")), + interval=( + to_timestamp("2023-01-01"), + to_timestamp("2023-01-02"), + ), batch_index=0, ), }, @@ -338,7 +350,10 @@ def test_incremental_time_self_reference_dag( ): { EvaluateNode( snapshot_name='"test_model"', - interval=(to_timestamp("2023-01-02"), to_timestamp("2023-01-03")), + interval=( + to_timestamp("2023-01-02"), + to_timestamp("2023-01-03"), + ), batch_index=1, ), }, @@ -349,7 +364,10 @@ def test_incremental_time_self_reference_dag( ): { EvaluateNode( snapshot_name='"test_model"', - interval=(to_timestamp("2023-01-03"), to_timestamp("2023-01-04")), + interval=( + to_timestamp("2023-01-03"), + to_timestamp("2023-01-04"), + ), batch_index=2, ), }, @@ -429,7 +447,9 @@ def test_incremental_batch_concurrency( SqlModel( name="test_model", kind=IncrementalByTimeRangeKind( - time_column="ds", batch_size=batch_size, batch_concurrency=batch_concurrency + time_column="ds", + batch_size=batch_size, + batch_concurrency=batch_concurrency, ), cron="@daily", start=start, @@ -437,7 +457,9 @@ def test_incremental_batch_concurrency( ), ) - snapshot_evaluator = SnapshotEvaluator(adapters=mocker.MagicMock(), ddl_concurrent_tasks=1) + snapshot_evaluator = SnapshotEvaluator( + adapters=mocker.MagicMock(), ddl_concurrent_tasks=1 + ) mock_state_sync = mocker.MagicMock() scheduler = Scheduler( snapshots=[snapshot], @@ -478,7 +500,9 @@ def test_intervals_with_end_date_on_model( ) ) - snapshot_evaluator = SnapshotEvaluator(adapters=mocker.MagicMock(), ddl_concurrent_tasks=1) + snapshot_evaluator = SnapshotEvaluator( + adapters=mocker.MagicMock(), ddl_concurrent_tasks=1 + ) scheduler = Scheduler( snapshots=[snapshot], snapshot_evaluator=snapshot_evaluator, @@ -489,9 +513,9 @@ def test_intervals_with_end_date_on_model( # generate for 1 year to show that the returned batches should only cover # the range defined on the model itself - batches = get_batched_missing_intervals(scheduler, start="2023-01-01", end="2024-01-01")[ - snapshot - ] + batches = get_batched_missing_intervals( + scheduler, start="2023-01-01", end="2024-01-01" + )[snapshot] assert len(batches) == 31 # days in Jan 2023 assert batches[0] == (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) @@ -499,18 +523,18 @@ def test_intervals_with_end_date_on_model( # generate for less than 1 month to ensure that the scheduler end date # takes precedence over the model end date - batches = get_batched_missing_intervals(scheduler, start="2023-01-01", end="2023-01-10")[ - snapshot - ] + batches = get_batched_missing_intervals( + scheduler, start="2023-01-01", end="2023-01-10" + )[snapshot] assert len(batches) == 10 assert batches[0] == (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")) assert batches[-1] == (to_timestamp("2023-01-10"), to_timestamp("2023-01-11")) # generate for the last day of range - batches = get_batched_missing_intervals(scheduler, start="2023-01-31", end="2023-01-31")[ - snapshot - ] + batches = get_batched_missing_intervals( + scheduler, start="2023-01-31", end="2023-01-31" + )[snapshot] assert len(batches) == 1 assert batches[0] == (to_timestamp("2023-01-31"), to_timestamp("2023-02-01")) @@ -523,8 +547,7 @@ def test_intervals_with_end_date_on_model( def test_external_model_audit(mocker, make_snapshot): model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind EXTERNAL, @@ -533,8 +556,7 @@ def test_external_model_audit(mocker, make_snapshot): ); SELECT 1; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -565,15 +587,20 @@ def test_audit_failure_notifications( scheduler: Scheduler, waiter_names: Snapshot, mocker: MockerFixture ): evaluator_evaluate_mock = mocker.Mock() - mocker.patch("sqlmesh.core.scheduler.SnapshotEvaluator.evaluate", evaluator_evaluate_mock) + mocker.patch( + "sqlmesh.core.scheduler.SnapshotEvaluator.evaluate", evaluator_evaluate_mock + ) evaluator_audit_mock = mocker.Mock() mocker.patch("sqlmesh.core.scheduler.SnapshotEvaluator.audit", evaluator_audit_mock) notify_user_mock = mocker.Mock() mocker.patch( - "sqlmesh.core.notification_target.NotificationTargetManager.notify_user", notify_user_mock + "sqlmesh.core.notification_target.NotificationTargetManager.notify_user", + notify_user_mock, ) notify_mock = mocker.Mock() - mocker.patch("sqlmesh.core.notification_target.NotificationTargetManager.notify", notify_mock) + mocker.patch( + "sqlmesh.core.notification_target.NotificationTargetManager.notify", notify_mock + ) audit = first(waiter_names.model.audit_definitions.values()) query = waiter_names.model.render_query() @@ -674,11 +701,16 @@ def test_interval_diff(): ) == [(3, 4)] assert interval_diff([(1, 2), (2, 3)], [(1, 2)], uninterrupted=True) == [] - assert interval_diff([(1, 2), (2, 3)], [(3, 4)], uninterrupted=True) == [(1, 2), (2, 3)] + assert interval_diff([(1, 2), (2, 3)], [(3, 4)], uninterrupted=True) == [ + (1, 2), + (2, 3), + ] assert interval_diff([(1, 2), (2, 3)], [(2, 3)], uninterrupted=True) == [(1, 2)] -def test_signal_intervals(mocker: MockerFixture, make_snapshot, get_batched_missing_intervals): +def test_signal_intervals( + mocker: MockerFixture, make_snapshot, get_batched_missing_intervals +): @signal() def signal_a(batch: DatetimeRanges, context: ExecutionContext): if not hasattr(context, "engine_adapter"): @@ -693,8 +725,7 @@ def signal_b(batch: DatetimeRanges): a = make_snapshot( load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name a, kind FULL, @@ -703,16 +734,14 @@ def signal_b(batch: DatetimeRanges): ); SELECT 1 x; - """ - ), + """), # type: ignore signal_definitions=signals, ), ) b = make_snapshot( load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name b, kind FULL, @@ -722,8 +751,7 @@ def signal_b(batch: DatetimeRanges): ); SELECT 2 x; - """ - ), + """), # type: ignore signal_definitions=signals, ), nodes={a.name: a.model}, @@ -731,8 +759,7 @@ def signal_b(batch: DatetimeRanges): c = make_snapshot( load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name c, kind FULL, @@ -740,16 +767,14 @@ def signal_b(batch: DatetimeRanges): ); SELECT * FROM a UNION SELECT * FROM b - """ - ), + """), # type: ignore signal_definitions=signals, ), nodes={a.name: a.model, b.name: b.model}, ) d = make_snapshot( load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name d, kind FULL, @@ -757,14 +782,15 @@ def signal_b(batch: DatetimeRanges): ); SELECT * FROM c UNION SELECT * FROM d - """ - ), + """), # type: ignore signal_definitions=signals, ), nodes={a.name: a.model, b.name: b.model, c.name: c.model}, ) - snapshot_evaluator = SnapshotEvaluator(adapters=mocker.MagicMock(), ddl_concurrent_tasks=1) + snapshot_evaluator = SnapshotEvaluator( + adapters=mocker.MagicMock(), ddl_concurrent_tasks=1 + ) scheduler = Scheduler( snapshots=[a, b, c, d], snapshot_evaluator=snapshot_evaluator, @@ -795,8 +821,7 @@ def signal_base(batch: DatetimeRanges): snapshot_a = make_snapshot( load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name a, kind INCREMENTAL_BY_TIME_RANGE( @@ -807,16 +832,14 @@ def signal_base(batch: DatetimeRanges): signals SIGNAL_BASE(), ); SELECT @start_date AS dt; - """ - ), + """), # type: ignore signal_definitions=signals, ), ) snapshot_b = make_snapshot( load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name b, kind INCREMENTAL_BY_TIME_RANGE( @@ -826,16 +849,14 @@ def signal_base(batch: DatetimeRanges): start '2023-01-01' ); SELECT @start_date AS dt; - """ - ), + """), # type: ignore signal_definitions=signals, ) ) snapshot_c = make_snapshot( load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name c, kind INCREMENTAL_BY_TIME_RANGE( @@ -845,14 +866,15 @@ def signal_base(batch: DatetimeRanges): start '2023-01-01', ); SELECT * FROM a UNION SELECT * FROM b - """ - ), + """), # type: ignore signal_definitions=signals, ), nodes={snapshot_a.name: snapshot_a.model, snapshot_b.name: snapshot_b.model}, ) - snapshot_evaluator = SnapshotEvaluator(adapters=mocker.MagicMock(), ddl_concurrent_tasks=1) + snapshot_evaluator = SnapshotEvaluator( + adapters=mocker.MagicMock(), ddl_concurrent_tasks=1 + ) scheduler = Scheduler( snapshots=[snapshot_c, snapshot_b, snapshot_a], # reverse order snapshot_evaluator=snapshot_evaluator, @@ -920,7 +942,9 @@ def test_scd_type_2_batch_size( snapshot = make_snapshot(model) # Setup scheduler - snapshot_evaluator = SnapshotEvaluator(adapters=mocker.MagicMock(), ddl_concurrent_tasks=1) + snapshot_evaluator = SnapshotEvaluator( + adapters=mocker.MagicMock(), ddl_concurrent_tasks=1 + ) scheduler = Scheduler( snapshots=[snapshot], snapshot_evaluator=snapshot_evaluator, @@ -936,7 +960,9 @@ def test_scd_type_2_batch_size( assert batches == expected_batches -def test_before_all_environment_statements_called_first(mocker: MockerFixture, make_snapshot): +def test_before_all_environment_statements_called_first( + mocker: MockerFixture, make_snapshot +): model = SqlModel( name="test.model_items", query=parse_one("SELECT id, ds FROM raw.items"), @@ -956,7 +982,9 @@ def record_get_environment_statements(*args, **kwargs): call_order.append("get_environment_statements") return mock_state_sync.get_environment_statements.return_value - mock_state_sync.get_environment_statements.side_effect = record_get_environment_statements + mock_state_sync.get_environment_statements.side_effect = ( + record_get_environment_statements + ) mock_snapshot_evaluator = mocker.MagicMock() mock_adapter = mocker.MagicMock() @@ -966,7 +994,9 @@ def record_get_snapshots_to_create(*args, **kwargs): call_order.append("get_snapshots_to_create") return [] - mock_snapshot_evaluator.get_snapshots_to_create.side_effect = record_get_snapshots_to_create + mock_snapshot_evaluator.get_snapshots_to_create.side_effect = ( + record_get_snapshots_to_create + ) mock_execute_env_statements = mocker.patch( "sqlmesh.core.scheduler.execute_environment_statements" @@ -1048,7 +1078,9 @@ def test_dag_transitive_deps(mocker: MockerFixture, make_snapshot): snapshot_c: [(to_timestamp("2023-01-01"), to_timestamp("2023-01-02"))], } - deployability_index = DeployabilityIndex.create([snapshot_a, snapshot_b, snapshot_c]) + deployability_index = DeployabilityIndex.create( + [snapshot_a, snapshot_b, snapshot_c] + ) full_dag = snapshots_to_dag([snapshot_a, snapshot_b, snapshot_c]) @@ -1080,16 +1112,24 @@ def test_dag_multiple_chain_transitive_deps(mocker: MockerFixture, make_snapshot # Select A and E only snapshots = {} for name in ["a", "b", "c", "d", "e"]: - snapshots[name] = make_snapshot(SqlModel(name=name, query=parse_one("SELECT 1 as id"))) + snapshots[name] = make_snapshot( + SqlModel(name=name, query=parse_one("SELECT 1 as id")) + ) snapshots[name].categorize_as(SnapshotChangeCategory.BREAKING) # Set up dependencies - snapshots["b"] = snapshots["b"].model_copy(update={"parents": (snapshots["a"].snapshot_id,)}) - snapshots["c"] = snapshots["c"].model_copy(update={"parents": (snapshots["a"].snapshot_id,)}) + snapshots["b"] = snapshots["b"].model_copy( + update={"parents": (snapshots["a"].snapshot_id,)} + ) + snapshots["c"] = snapshots["c"].model_copy( + update={"parents": (snapshots["a"].snapshot_id,)} + ) snapshots["d"] = snapshots["d"].model_copy( update={"parents": (snapshots["b"].snapshot_id, snapshots["c"].snapshot_id)} ) - snapshots["e"] = snapshots["e"].model_copy(update={"parents": (snapshots["d"].snapshot_id,)}) + snapshots["e"] = snapshots["e"].model_copy( + update={"parents": (snapshots["d"].snapshot_id,)} + ) scheduler = Scheduler( snapshots=list(snapshots.values()), @@ -1128,7 +1168,9 @@ def test_dag_multiple_chain_transitive_deps(mocker: MockerFixture, make_snapshot } -def test_dag_upstream_dependency_caching_with_complex_diamond(mocker: MockerFixture, make_snapshot): +def test_dag_upstream_dependency_caching_with_complex_diamond( + mocker: MockerFixture, make_snapshot +): r""" Test that the upstream dependency caching correctly handles a complex diamond dependency graph. @@ -1148,19 +1190,29 @@ def test_dag_upstream_dependency_caching_with_complex_diamond(mocker: MockerFixt snapshots = {} for name in ["a", "b", "c", "d", "e", "f", "g", "h"]: - snapshots[name] = make_snapshot(SqlModel(name=name, query=parse_one("SELECT 1 as id"))) + snapshots[name] = make_snapshot( + SqlModel(name=name, query=parse_one("SELECT 1 as id")) + ) snapshots[name].categorize_as(SnapshotChangeCategory.BREAKING) # A is the root - snapshots["b"] = snapshots["b"].model_copy(update={"parents": (snapshots["a"].snapshot_id,)}) - snapshots["c"] = snapshots["c"].model_copy(update={"parents": (snapshots["a"].snapshot_id,)}) + snapshots["b"] = snapshots["b"].model_copy( + update={"parents": (snapshots["a"].snapshot_id,)} + ) + snapshots["c"] = snapshots["c"].model_copy( + update={"parents": (snapshots["a"].snapshot_id,)} + ) # Middle layer: D, E, F depend on B and/or C - snapshots["d"] = snapshots["d"].model_copy(update={"parents": (snapshots["b"].snapshot_id,)}) + snapshots["d"] = snapshots["d"].model_copy( + update={"parents": (snapshots["b"].snapshot_id,)} + ) snapshots["e"] = snapshots["e"].model_copy( update={"parents": (snapshots["b"].snapshot_id, snapshots["c"].snapshot_id)} ) - snapshots["f"] = snapshots["f"].model_copy(update={"parents": (snapshots["c"].snapshot_id,)}) + snapshots["f"] = snapshots["f"].model_copy( + update={"parents": (snapshots["c"].snapshot_id,)} + ) # Bottom layer: G and H depend on D/E and E/F respectively snapshots["g"] = snapshots["g"].model_copy( diff --git a/tests/core/test_schema_diff.py b/tests/core/test_schema_diff.py index 52bd6bb606..9c4b7137ae 100644 --- a/tests/core/test_schema_diff.py +++ b/tests/core/test_schema_diff.py @@ -4,17 +4,13 @@ from sqlglot import exp from sqlmesh.core.engine_adapter import create_engine_adapter -from sqlmesh.core.schema_diff import ( - SchemaDiffer, - TableAlterColumn, - TableAlterColumnPosition, - TableAlterOperation, - get_schema_differ, - TableAlterAddColumnOperation, - TableAlterDropColumnOperation, - TableAlterChangeColumnTypeOperation, - NestedSupport, -) +from sqlmesh.core.schema_diff import (NestedSupport, SchemaDiffer, + TableAlterAddColumnOperation, + TableAlterChangeColumnTypeOperation, + TableAlterColumn, + TableAlterColumnPosition, + TableAlterDropColumnOperation, + TableAlterOperation, get_schema_differ) from sqlmesh.utils.errors import SQLMeshError @@ -113,7 +109,12 @@ def test_schema_diff_calculate_type_transitions(): # # Add Tests # ########### # No diff - ("STRUCT", "STRUCT", [], {}), + ( + "STRUCT", + "STRUCT", + [], + {}, + ), # Add root level column at the end ( "STRUCT", @@ -244,7 +245,9 @@ def test_schema_diff_calculate_type_transitions(): TableAlterDropColumnOperation( target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("id")], - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ) ], @@ -272,7 +275,9 @@ def test_schema_diff_calculate_type_transitions(): TableAlterDropColumnOperation( target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("age")], - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ) ], @@ -302,7 +307,9 @@ def test_schema_diff_calculate_type_transitions(): TableAlterDropColumnOperation( target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("age")], - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ), ], @@ -336,11 +343,26 @@ def test_schema_diff_calculate_type_transitions(): # Move Tests ############# # Move root level column at the start - ("STRUCT", "STRUCT", [], {}), + ( + "STRUCT", + "STRUCT", + [], + {}, + ), # Move root level column in the middle - ("STRUCT", "STRUCT", [], {}), + ( + "STRUCT", + "STRUCT", + [], + {}, + ), # Move root level column at the end - ("STRUCT", "STRUCT", [], {}), + ( + "STRUCT", + "STRUCT", + [], + {}, + ), # ################### # # Type Change Tests # ################### @@ -428,7 +450,9 @@ def test_schema_diff_calculate_type_transitions(): position=TableAlterColumnPosition.first(), ), ], - dict(support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # Add a column to the end of a struct ( @@ -449,7 +473,9 @@ def test_schema_diff_calculate_type_transitions(): position=TableAlterColumnPosition.last(after="col_c"), ), ], - dict(support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # Add a column to the middle of a struct ( @@ -470,7 +496,9 @@ def test_schema_diff_calculate_type_transitions(): array_element_selector="", ), ], - dict(support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # Add two columns at the start of a struct ( @@ -504,7 +532,9 @@ def test_schema_diff_calculate_type_transitions(): array_element_selector="", ), ], - dict(support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # Add columns in different levels of nesting of structs ( @@ -535,7 +565,9 @@ def test_schema_diff_calculate_type_transitions(): array_element_selector="", ), ], - dict(support_positional_add=False, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=False, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # Remove a column from the start of a struct ( @@ -778,7 +810,9 @@ def test_schema_diff_calculate_type_transitions(): expected_table_struct=exp.DataType.build( "STRUCT>" ), - column_type=exp.DataType.build("STRUCT"), + column_type=exp.DataType.build( + "STRUCT" + ), array_element_selector="", is_part_of_destructive_change=True, ), @@ -873,7 +907,9 @@ def test_schema_diff_calculate_type_transitions(): array_element_selector="", ), ], - dict(support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # Remove column from array of structs ( @@ -951,7 +987,9 @@ def test_schema_diff_calculate_type_transitions(): array_element_selector="", ), ], - dict(support_positional_add=False, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=False, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # Add an array of primitives ( @@ -971,7 +1009,9 @@ def test_schema_diff_calculate_type_transitions(): array_element_selector="", ), ], - dict(support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP), + dict( + support_positional_add=True, nested_support=NestedSupport.ALL_BUT_DROP + ), ), # untyped array to support Snowflake ( @@ -995,7 +1035,9 @@ def test_schema_diff_calculate_type_transitions(): target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("ids")], column_type=exp.DataType.build("ARRAY"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", is_part_of_destructive_change=True, ), @@ -1145,7 +1187,9 @@ def test_schema_diff_calculate_type_transitions(): target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("address")], column_type=exp.DataType.build("VARCHAR"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), position=TableAlterColumnPosition.last("id"), array_element_selector="", is_part_of_destructive_change=True, @@ -1192,7 +1236,9 @@ def test_schema_diff_calculate_type_transitions(): column_parts=[TableAlterColumn.primitive("address")], column_type=exp.DataType.build("VARCHAR(2)"), current_type=exp.DataType.build("VARCHAR"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ) ], @@ -1217,7 +1263,9 @@ def test_schema_diff_calculate_type_transitions(): target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("address")], column_type=exp.DataType.build("VARCHAR"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), position=TableAlterColumnPosition.last("id"), array_element_selector="", is_part_of_destructive_change=True, @@ -1314,7 +1362,9 @@ def test_schema_diff_calculate_type_transitions(): column_parts=[TableAlterColumn.primitive("address")], column_type=exp.DataType.build("VARCHAR"), current_type=exp.DataType.build("VARCHAR(120)"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ) ], @@ -1368,7 +1418,9 @@ def test_schema_diff_calculate_type_transitions(): column_parts=[TableAlterColumn.primitive("address")], column_type=exp.DataType.build("TEXT"), current_type=exp.DataType.build("VARCHAR(120)"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ) ], @@ -1535,7 +1587,9 @@ def test_schema_diff_calculate_type_transitions(): column_parts=[TableAlterColumn.primitive("age")], column_type=exp.DataType.build("STRING"), current_type=exp.DataType.build("INT"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", is_part_of_destructive_change=True, ), @@ -1587,7 +1641,9 @@ def test_schema_diff_calculate_duckdb(duck_conn): }, ) - alter_expressions = engine_adapter.get_alter_operations("apply_to_table", "schema_from_table") + alter_expressions = engine_adapter.get_alter_operations( + "apply_to_table", "schema_from_table" + ) engine_adapter.alter_table(alter_expressions) assert engine_adapter.columns("apply_to_table") == { "id": exp.DataType.build("int"), @@ -1605,7 +1661,9 @@ def test_schema_diff_alter_op_column(): TableAlterColumn.primitive("col_a"), ], column_type=exp.DataType.build("INT"), - expected_table_struct=exp.DataType.build("STRUCT>>"), + expected_table_struct=exp.DataType.build( + "STRUCT>>" + ), position=TableAlterColumnPosition.last("id"), array_element_selector="", ) @@ -1870,14 +1928,20 @@ def test_ignore_destructive_compare_columns(): assert len(alter_expressions_ignore_destructive) == 2 # Only ADD + ALTER # Verify the operations are correct - operations_sql = [expr.expression.sql() for expr in alter_expressions_ignore_destructive] + operations_sql = [ + expr.expression.sql() for expr in alter_expressions_ignore_destructive + ] add_column_found = any("ADD COLUMN new_col DOUBLE" in op for op in operations_sql) - alter_column_found = any("ALTER COLUMN id SET DATA TYPE" in op for op in operations_sql) + alter_column_found = any( + "ALTER COLUMN id SET DATA TYPE" in op for op in operations_sql + ) drop_column_found = any("DROP COLUMN to_drop" in op for op in operations_sql) assert add_column_found, f"ADD COLUMN not found in: {operations_sql}" assert alter_column_found, f"ALTER COLUMN not found in: {operations_sql}" - assert not drop_column_found, f"DROP COLUMN should not be present in: {operations_sql}" + assert ( + not drop_column_found + ), f"DROP COLUMN should not be present in: {operations_sql}" def test_ignore_destructive_nested_struct_without_support(): @@ -1912,7 +1976,14 @@ def test_ignore_destructive_nested_struct_without_support(): def test_get_schema_differ(): # Test that known dialects return SchemaDiffer instances - for dialect in ["bigquery", "snowflake", "postgres", "databricks", "spark", "duckdb"]: + for dialect in [ + "bigquery", + "snowflake", + "postgres", + "databricks", + "spark", + "duckdb", + ]: schema_differ = get_schema_differ(dialect) assert isinstance(schema_differ, SchemaDiffer) @@ -2045,7 +2116,9 @@ def test_ignore_destructive_edge_cases(): target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("name")], column_type=exp.DataType.build("STRING"), - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ), TableAlterAddColumnOperation( @@ -2078,7 +2151,9 @@ def test_ignore_destructive_edge_cases(): TableAlterDropColumnOperation( target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("age")], - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ), ], @@ -2087,7 +2162,9 @@ def test_ignore_destructive_edge_cases(): TableAlterDropColumnOperation( target_table=exp.to_table("apply_to_table"), column_parts=[TableAlterColumn.primitive("age")], - expected_table_struct=exp.DataType.build("STRUCT"), + expected_table_struct=exp.DataType.build( + "STRUCT" + ), array_element_selector="", ), ], @@ -2243,7 +2320,9 @@ def test_ignore_additive_edge_cases(): # Test when all operations are additive - should result in empty list current_struct = "STRUCT" - new_struct = "STRUCT" # Add all columns + new_struct = ( + "STRUCT" # Add all columns + ) operations_ignore_additive = schema_differ._from_structs( exp.DataType.build(current_struct), @@ -2302,7 +2381,9 @@ def test_ignore_both_destructive_and_additive(): ) current_struct = "STRUCT" - new_struct = "STRUCT" # DROP name, ADD address, ALTER id + new_struct = ( + "STRUCT" # DROP name, ADD address, ALTER id + ) operations_ignore_both = schema_differ._from_structs( exp.DataType.build(current_struct), @@ -2322,7 +2403,9 @@ def test_ignore_additive_array_operations(): ) current_struct = "STRUCT>>" - new_struct = "STRUCT>>" + new_struct = ( + "STRUCT>>" + ) # With additive operations allowed - should add to array struct operations_with_additive = schema_differ._from_structs( diff --git a/tests/core/test_schema_loader.py b/tests/core/test_schema_loader.py index bb87f35b3d..69923d5709 100644 --- a/tests/core/test_schema_loader.py +++ b/tests/core/test_schema_loader.py @@ -1,9 +1,9 @@ -import pytest import typing as t from pathlib import Path from unittest.mock import patch import pandas as pd # noqa: TID253 +import pytest from pytest_mock.plugin import MockerFixture from sqlglot import exp, parse_one @@ -11,7 +11,8 @@ from sqlmesh.core.config import Config, DuckDBConnectionConfig, GatewayConfig from sqlmesh.core.context import Context from sqlmesh.core.dialect import parse -from sqlmesh.core.model import SqlModel, create_external_model, load_sql_based_model +from sqlmesh.core.model import (SqlModel, create_external_model, + load_sql_based_model) from sqlmesh.core.model.definition import ExternalModel from sqlmesh.core.schema_loader import create_external_models_file from sqlmesh.core.snapshot import SnapshotChangeCategory @@ -162,8 +163,14 @@ def _create_model(gateway: str): contents = yaml.load(tmp_path / c.EXTERNAL_MODELS_YAML) assert len(contents) == 2 - assert len([c for c in contents if c["name"] == '"memory"."landing"."dev_source"']) == 1 - assert len([c for c in contents if c["name"] == '"memory"."landing"."prod_source"']) == 1 + assert ( + len([c for c in contents if c["name"] == '"memory"."landing"."dev_source"']) + == 1 + ) + assert ( + len([c for c in contents if c["name"] == '"memory"."landing"."prod_source"']) + == 1 + ) def test_gateway_specific_external_models_mixed_with_others(tmp_path: Path): @@ -194,7 +201,9 @@ def _init_db(ctx: Context): """, ) - ctx = Context(paths=[tmp_path], config=config) # note: No explicitly defined gateway + ctx = Context( + paths=[tmp_path], config=config + ) # note: No explicitly defined gateway assert ctx.gateway is None assert ctx.selected_gateway == "dev" @@ -235,7 +244,9 @@ def _init_db(ctx: Context): # check that this doesnt present a problem on load prod_ctx.load() - external_models = [m for _, m in prod_ctx.models.items() if type(m) == ExternalModel] + external_models = [ + m for _, m in prod_ctx.models.items() if type(m) == ExternalModel + ] assert len(external_models) == 1 assert external_models[0].name == '"memory"."landing"."source_table"' assert external_models[0].gateway == "prod" @@ -305,7 +316,9 @@ def _load_external_models(): assert len(_load_external_models()) == 1 -def test_create_external_models_skips_models_from_external_models_directory(tmp_path: Path): +def test_create_external_models_skips_models_from_external_models_directory( + tmp_path: Path, +): config = Config(gateways={"": GatewayConfig(connection=DuckDBConnectionConfig())}) model_dir = tmp_path / c.MODELS @@ -347,7 +360,9 @@ def test_create_external_models_skips_models_from_external_models_directory(tmp_ assert yaml.load(tmp_path / c.EXTERNAL_MODELS_YAML) == [] ctx.load() - external_models = [model for model in ctx.models.values() if isinstance(model, ExternalModel)] + external_models = [ + model for model in ctx.models.values() if isinstance(model, ExternalModel) + ] assert len(external_models) == 1 assert external_models[0].fqn == '"memory"."landing"."source_table"' @@ -363,7 +378,9 @@ def test_no_internal_model_conversion(tmp_path: Path, mocker: MockerFixture): state_reader_mock.nodes_exist.return_value = {'"model_b"'} model_a = SqlModel(name="a", query=parse_one("select * FROM model_b, tbl_c")) - model_b = SqlModel(name="b", query=parse_one("select * FROM `tbl-d`", read="bigquery")) + model_b = SqlModel( + name="b", query=parse_one("select * FROM `tbl-d`", read="bigquery") + ) filename = tmp_path / c.EXTERNAL_MODELS_YAML create_external_models_file( @@ -408,7 +425,9 @@ def test_missing_table(tmp_path: Path): schema = yaml.load(filename) assert len(schema) == 0 - with pytest.raises(SQLMeshError, match=r"""Unable to get schema for '"tbl_source"'.*"""): + with pytest.raises( + SQLMeshError, match=r"""Unable to get schema for '"tbl_source"'.*""" + ): create_external_models_file( filename, {"a": model}, # type: ignore diff --git a/tests/core/test_selector_dbt.py b/tests/core/test_selector_dbt.py index 112c5740ac..eedf862f45 100644 --- a/tests/core/test_selector_dbt.py +++ b/tests/core/test_selector_dbt.py @@ -1,17 +1,18 @@ import typing as t + import pytest from pytest_mock import MockerFixture from sqlglot import exp -from sqlmesh.core.model.kind import SeedKind, ExternalKind, FullKind -from sqlmesh.core.model.seed import Seed -from sqlmesh.core.model.definition import SqlModel, SeedModel, ExternalModel + +import sqlmesh.core.dialect as d from sqlmesh.core.audit.definition import StandaloneAudit +from sqlmesh.core.model.definition import ExternalModel, SeedModel, SqlModel +from sqlmesh.core.model.kind import ExternalKind, FullKind, SeedKind +from sqlmesh.core.model.seed import Seed +from sqlmesh.core.selector import DbtSelector, ResourceType, parse from sqlmesh.core.snapshot.definition import Node -from sqlmesh.core.selector import DbtSelector -from sqlmesh.core.selector import parse, ResourceType -from sqlmesh.utils.errors import SQLMeshError -import sqlmesh.core.dialect as d from sqlmesh.utils import UniqueKeyDict +from sqlmesh.utils.errors import SQLMeshError def test_parse_resource_type(): @@ -37,17 +38,23 @@ def test_expand_model_selections_resource_type( query=d.parse_one("SELECT 'normal_model' AS what"), ), '"test"."seed_model"': SeedModel( - name="test.seed_model", kind=SeedKind(path="/tmp/foo"), seed=Seed(content="id,name") + name="test.seed_model", + kind=SeedKind(path="/tmp/foo"), + seed=Seed(content="id,name"), ), '"test"."standalone_audit"': StandaloneAudit( - name="test.standalone_audit", query=d.parse_one("SELECT 'standalone_audit' AS what") + name="test.standalone_audit", + query=d.parse_one("SELECT 'standalone_audit' AS what"), ), '"external"."model"': ExternalModel(name="external.model", kind=ExternalKind()), } selector = DbtSelector(state_reader=mocker.Mock(), models=UniqueKeyDict("models")) - assert selector.expand_model_selections([f"resource_type:{resource_type}"], models) == expected + assert ( + selector.expand_model_selections([f"resource_type:{resource_type}"], models) + == expected + ) def test_unsupported_resource_type(mocker: MockerFixture): diff --git a/tests/core/test_selector_native.py b/tests/core/test_selector_native.py index e8b6f8a7ad..8259acd302 100644 --- a/tests/core/test_selector_native.py +++ b/tests/core/test_selector_native.py @@ -1,12 +1,12 @@ from __future__ import annotations +import subprocess import typing as t from pathlib import Path from unittest.mock import call import pytest from pytest_mock.plugin import MockerFixture -import subprocess from sqlmesh.core import dialect as d from sqlmesh.core.audit import StandaloneAudit @@ -27,7 +27,9 @@ "test_catalog", ], ) -def test_select_models(mocker: MockerFixture, make_snapshot, default_catalog: t.Optional[str]): +def test_select_models( + mocker: MockerFixture, make_snapshot, default_catalog: t.Optional[str] +): added_model = SqlModel( name="db.added_model", query=d.parse_one("SELECT 1 AS a"), @@ -52,7 +54,8 @@ def test_select_models(mocker: MockerFixture, make_snapshot, default_catalog: t. default_catalog=default_catalog, ) standalone_audit = StandaloneAudit( - name="test_audit", query=d.parse_one("SELECT * FROM added_model WHERE a IS NULL") + name="test_audit", + query=d.parse_one("SELECT * FROM added_model WHERE a IS NULL"), ) modified_model_v1_snapshot = make_snapshot(modified_model_v1) @@ -69,7 +72,11 @@ def test_select_models(mocker: MockerFixture, make_snapshot, default_catalog: t. name=env_name, snapshots=[ s.table_info - for s in (modified_model_v1_snapshot, removed_model_snapshot, standalone_audit_snapshot) + for s in ( + modified_model_v1_snapshot, + removed_model_snapshot, + standalone_audit_snapshot, + ) ], start_at="2023-01-01", end_at="2023-02-01", @@ -90,7 +97,9 @@ def test_select_models(mocker: MockerFixture, make_snapshot, default_catalog: t. local_models[modified_model_v2.fqn] = modified_model_v2.copy( update={"mapping_schema": added_model_schema} ) - selector = NativeSelector(state_reader_mock, local_models, default_catalog=default_catalog) + selector = NativeSelector( + state_reader_mock, local_models, default_catalog=default_catalog + ) _assert_models_equal( selector.select_models(["db.added_model"], env_name), @@ -103,7 +112,9 @@ def test_select_models(mocker: MockerFixture, make_snapshot, default_catalog: t. }, ) _assert_models_equal( - selector.select_models(["db.modified_model"], "missing_env", fallback_env_name=env_name), + selector.select_models( + ["db.modified_model"], "missing_env", fallback_env_name=env_name + ), { modified_model_v2.fqn: modified_model_v2, removed_model.fqn: removed_model, @@ -117,7 +128,9 @@ def test_select_models(mocker: MockerFixture, make_snapshot, default_catalog: t. ) _assert_models_equal( selector.select_models( - ["db.added_model", "db.modified_model"], "missing_env", fallback_env_name=env_name + ["db.added_model", "db.modified_model"], + "missing_env", + fallback_env_name=env_name, ), { added_model.fqn: added_model, @@ -158,7 +171,9 @@ def test_select_models(mocker: MockerFixture, make_snapshot, default_catalog: t. local_models, ) _assert_models_equal( - selector.select_models(["tag:tag1", "tag:tag2"], "missing_env", fallback_env_name=env_name), + selector.select_models( + ["tag:tag1", "tag:tag2"], "missing_env", fallback_env_name=env_name + ), { added_model.fqn: added_model, modified_model_v2.fqn: modified_model_v2.copy( @@ -203,7 +218,8 @@ def test_select_models_expired_environment(mocker: MockerFixture, make_snapshot) query=d.parse_one("SELECT a FROM db.added_model"), ) standalone_audit = StandaloneAudit( - name="test_audit", query=d.parse_one("SELECT * FROM added_model WHERE a IS NULL") + name="test_audit", + query=d.parse_one("SELECT * FROM added_model WHERE a IS NULL"), ) modified_model_v1_snapshot = make_snapshot(modified_model_v1) @@ -224,7 +240,10 @@ def test_select_models_expired_environment(mocker: MockerFixture, make_snapshot) env_name = "test_env" dev_env = Environment( name=env_name, - snapshots=[modified_model_v1_snapshot.table_info, removed_model_snapshot.table_info], + snapshots=[ + modified_model_v1_snapshot.table_info, + removed_model_snapshot.table_info, + ], start_at="2023-01-01", end_at="2023-02-01", plan_id="test_plan_id", @@ -248,7 +267,9 @@ def test_select_models_expired_environment(mocker: MockerFixture, make_snapshot) selector = NativeSelector(state_reader_mock, local_models) _assert_models_equal( - selector.select_models(["*.modified_model"], env_name, fallback_env_name="prod"), + selector.select_models( + ["*.modified_model"], env_name, fallback_env_name="prod" + ), { removed_model.fqn: removed_model, modified_model_v2.fqn: modified_model_v2, @@ -257,7 +278,9 @@ def test_select_models_expired_environment(mocker: MockerFixture, make_snapshot) dev_env.expiration_ts = now_timestamp() - 1 _assert_models_equal( - selector.select_models(["*.modified_model"], env_name, fallback_env_name="prod"), + selector.select_models( + ["*.modified_model"], env_name, fallback_env_name="prod" + ), { modified_model_v2.fqn: modified_model_v2, }, @@ -304,7 +327,9 @@ def test_select_change_schema(mocker: MockerFixture, make_snapshot): } ) local_models[local_parent.fqn] = local_parent - local_child = child.copy(update={"mapping_schema": {'"db"': {'"parent"': {"b": "INT"}}}}) + local_child = child.copy( + update={"mapping_schema": {'"db"': {'"parent"': {"b": "INT"}}}} + ) local_models[local_child.fqn] = local_child selector = NativeSelector(state_reader_mock, local_models) @@ -365,25 +390,41 @@ def test_select_models_missing_env(mocker: MockerFixture, make_snapshot): [ # Direct matching only ( - [("model1", "tag1", None), ("model2", "tag2", None), ("model3", "tag3", None)], + [ + ("model1", "tag1", None), + ("model2", "tag2", None), + ("model3", "tag3", None), + ], ["tag:tag1", "tag:tag3"], {'"model1"', '"model3"'}, ), # Wildcard works ( - [("model1", "tag1", None), ("model2", "tag2", None), ("model3", "tag3", None)], + [ + ("model1", "tag1", None), + ("model2", "tag2", None), + ("model3", "tag3", None), + ], ["tag:tag*"], {'"model1"', '"model2"', '"model3"'}, ), # Downstream models are included ( - [("model1", "tag1", None), ("model2", "tag2", {"model1"}), ("model3", "tag3", None)], + [ + ("model1", "tag1", None), + ("model2", "tag2", {"model1"}), + ("model3", "tag3", None), + ], ["tag:tag1+"], {'"model1"', '"model2"'}, ), # Upstream models are included ( - [("model1", "tag1", None), ("model2", "tag2", None), ("model3", "tag3", {"model2"})], + [ + ("model1", "tag1", None), + ("model2", "tag2", None), + ("model3", "tag3", {"model2"}), + ], ["+tag:tag3"], {'"model2"', '"model3"'}, ), @@ -507,17 +548,29 @@ def test_select_models_missing_env(mocker: MockerFixture, make_snapshot): ), # negation ( - [("model1", "tag1", None), ("model2", "tag2", None), ("model3", "tag3", None)], + [ + ("model1", "tag1", None), + ("model2", "tag2", None), + ("model3", "tag3", None), + ], ["^tag:tag1"], {'"model2"', '"model3"'}, ), ( - [("model1", "tag1", None), ("model2", "tag2", None), ("model3", "tag3", None)], + [ + ("model1", "tag1", None), + ("model2", "tag2", None), + ("model3", "tag3", None), + ], ["^model1"], {'"model2"', '"model3"'}, ), ( - [("model1", "tag1", None), ("model2", "tag2", None), ("model3", "tag3", None)], + [ + ("model1", "tag1", None), + ("model2", "tag2", None), + ("model3", "tag3", None), + ], ["model* & ^(tag:tag1 | tag:tag2)"], {'"model3"'}, ), @@ -549,7 +602,11 @@ def test_select_models_missing_env(mocker: MockerFixture, make_snapshot): {'"model2"', '"model3"'}, ), ( - [("model2", "tag1", None), ("model2_1", "tag2", None), ("model2_2", "tag3", None)], + [ + ("model2", "tag1", None), + ("model2_1", "tag2", None), + ("model2_2", "tag3", None), + ], ["*2_*"], {'"model2_1"', '"model2_2"'}, ), @@ -561,7 +618,10 @@ def test_expand_model_selections( models: UniqueKeyDict[str, Model] = UniqueKeyDict("models") for model_name, tag, depends_on in model_defs: model = SqlModel( - name=model_name, query=d.parse_one("SELECT 1 AS a"), depends_on=depends_on, tags=[tag] + name=model_name, + query=d.parse_one("SELECT 1 AS a"), + depends_on=depends_on, + tags=[tag], ) models[model.fqn] = model @@ -589,7 +649,10 @@ def test_model_selection_normalized(mocker: MockerFixture, make_snapshot): (["git:main & +*model_c"], {'"test_model_c"'}), (["git:main+"], {'"test_model_a"', '"test_model_c"', '"test_model_d"'}), (["+git:main"], {'"test_model_a"', '"test_model_c"', '"test_model_b"'}), - (["+git:main+"], {'"test_model_a"', '"test_model_c"', '"test_model_b"', '"test_model_d"'}), + ( + ["+git:main+"], + {'"test_model_a"', '"test_model_c"', '"test_model_b"', '"test_model_d"'}, + ), ], ) def test_expand_git_selection( @@ -624,14 +687,19 @@ def test_expand_git_selection( git_client_mock = mocker.Mock() git_client_mock.list_untracked_files.return_value = [] git_client_mock.list_uncommitted_changed_files.return_value = [] - git_client_mock.list_committed_changed_files.return_value = [model_a._path, model_c._path] + git_client_mock.list_committed_changed_files.return_value = [ + model_a._path, + model_c._path, + ] selector = NativeSelector(mocker.Mock(), models) selector._git_client = git_client_mock assert selector.expand_model_selections(expressions) == expected_fqns - git_client_mock.list_committed_changed_files.assert_called_once_with(target_branch="main") + git_client_mock.list_committed_changed_files.assert_called_once_with( + target_branch="main" + ) git_client_mock.list_uncommitted_changed_files.assert_called_once() git_client_mock.list_untracked_files.assert_called_once() @@ -639,7 +707,9 @@ def test_expand_git_selection( def test_expand_git_selection_integration(tmp_path: Path, mocker: MockerFixture): repo_path = tmp_path / "test_repo" repo_path.mkdir() - subprocess.run(["git", "init", "-b", "main"], cwd=repo_path, check=True, capture_output=True) + subprocess.run( + ["git", "init", "-b", "main"], cwd=repo_path, check=True, capture_output=True + ) models: UniqueKeyDict[str, Model] = UniqueKeyDict("models") model_a_path = repo_path / "model_a.sql" @@ -682,7 +752,9 @@ def test_expand_git_selection_integration(tmp_path: Path, mocker: MockerFixture) assert selector.expand_model_selections([f"git:main"]) == {'"test_model_a"'} # stage model A, should still select it - subprocess.run(["git", "add", "model_a.sql"], cwd=repo_path, check=True, capture_output=True) + subprocess.run( + ["git", "add", "model_a.sql"], cwd=repo_path, check=True, capture_output=True + ) assert selector.expand_model_selections([f"git:main"]) == {'"test_model_a"'} # now add unstaged change to B and both should be selected @@ -746,7 +818,9 @@ def test_select_models_with_external_parent(mocker: MockerFixture): local_models: UniqueKeyDict[str, Model] = UniqueKeyDict("models") local_models[added_model.fqn] = added_model - selector = NativeSelector(state_reader_mock, local_models, default_catalog=default_catalog) + selector = NativeSelector( + state_reader_mock, local_models, default_catalog=default_catalog + ) expanded_selections = selector.expand_model_selections(["+*added_model*"]) assert expanded_selections == {added_model.fqn} @@ -844,7 +918,9 @@ def test_select_models_returns_selected_fqns(mocker: MockerFixture, make_snapsho assert local_model.fqn in selected_fqns # Mixed selection (active + deleted): both appear in selected_fqns. - _, selected_fqns = selector.select_models(["db.deleted_model", "db.local_model"], env_name) + _, selected_fqns = selector.select_models( + ["db.deleted_model", "db.local_model"], env_name + ) assert selected_fqns == {deleted_model.fqn, local_model.fqn} # Wildcard should match both local and env models. diff --git a/tests/core/test_snapshot.py b/tests/core/test_snapshot.py index 64bee7f472..919f77a733 100644 --- a/tests/core/test_snapshot.py +++ b/tests/core/test_snapshot.py @@ -1,5 +1,5 @@ -import pickle import json +import pickle import typing as t from copy import deepcopy from datetime import datetime, timedelta @@ -13,65 +13,45 @@ from sqlmesh.core import constants as c from sqlmesh.core.audit import StandaloneAudit -from sqlmesh.core.config import ( - AutoCategorizationMode, - CategorizerConfig, - EnvironmentSuffixTarget, -) +from sqlmesh.core.config import (AutoCategorizationMode, CategorizerConfig, + EnvironmentSuffixTarget) +from sqlmesh.core.config.common import VirtualEnvironmentMode +from sqlmesh.core.console import get_console from sqlmesh.core.context import Context from sqlmesh.core.dialect import parse, parse_one from sqlmesh.core.environment import EnvironmentNamingInfo from sqlmesh.core.macros import SQL -from sqlmesh.core.model import ( - FullKind, - IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind, - IncrementalUnmanagedKind, - Model, - Seed, - SeedKind, - SeedModel, - SqlModel, - create_seed_model, - load_sql_based_model, - CustomKind, -) -from sqlmesh.core.model.kind import TimeColumn, ModelKindName +from sqlmesh.core.model import (CustomKind, FullKind, + IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, + IncrementalUnmanagedKind, Model, Seed, + SeedKind, SeedModel, SqlModel, + create_seed_model, load_sql_based_model) +from sqlmesh.core.model.kind import ModelKindName, TimeColumn from sqlmesh.core.node import IntervalUnit from sqlmesh.core.signal import signal -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - QualifiedViewName, - Snapshot, - SnapshotId, - SnapshotIdAndVersion, - SnapshotChangeCategory, - SnapshotFingerprint, - SnapshotIntervals, - SnapshotTableInfo, - earliest_start_date, - fingerprint_from_node, - has_paused_forward_only, - missing_intervals, -) +from sqlmesh.core.snapshot import (DeployabilityIndex, QualifiedViewName, + Snapshot, SnapshotChangeCategory, + SnapshotFingerprint, SnapshotId, + SnapshotIdAndVersion, SnapshotIntervals, + SnapshotTableInfo, earliest_start_date, + fingerprint_from_node, + has_paused_forward_only, missing_intervals) from sqlmesh.core.snapshot.cache import SnapshotCache from sqlmesh.core.snapshot.categorizer import categorize_change -from sqlmesh.core.snapshot.definition import ( - apply_auto_restatements, - display_name, - get_next_model_interval_start, - check_ready_intervals, - _contiguous_intervals, - table_name, - TableNamingConvention, -) -from sqlmesh.core.config.common import VirtualEnvironmentMode +from sqlmesh.core.snapshot.definition import (TableNamingConvention, + _contiguous_intervals, + apply_auto_restatements, + check_ready_intervals, + display_name, + get_next_model_interval_start, + table_name) from sqlmesh.utils import AttributeDict -from sqlmesh.utils.date import DatetimeRanges, to_date, to_datetime, to_timestamp -from sqlmesh.utils.errors import SQLMeshError, SignalEvalError -from sqlmesh.utils.jinja import JinjaMacroRegistry, MacroInfo +from sqlmesh.utils.date import (DatetimeRanges, to_date, to_datetime, + to_timestamp) +from sqlmesh.utils.errors import SignalEvalError, SQLMeshError from sqlmesh.utils.hashing import md5 -from sqlmesh.core.console import get_console +from sqlmesh.utils.jinja import JinjaMacroRegistry, MacroInfo @pytest.fixture @@ -88,7 +68,11 @@ def parent_model(): def model(): return SqlModel( name="name", - kind=dict(time_column="ds", batch_size=30, name=ModelKindName.INCREMENTAL_BY_TIME_RANGE), + kind=dict( + time_column="ds", + batch_size=30, + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + ), owner="owner", dialect="spark", cron="1 0 * * *", @@ -171,7 +155,9 @@ def test_json(snapshot: Snapshot): "grants_target_layer": "virtual", }, "name": '"name"', - "parents": [{"name": '"parent"."tbl"', "identifier": snapshot.parents[0].identifier}], + "parents": [ + {"name": '"parent"."tbl"', "identifier": snapshot.parents[0].identifier} + ], "previous_versions": [], "table_naming_convention": "schema_and_table", "updated_ts": 1663891973000, @@ -187,7 +173,11 @@ def test_json_with_grants(make_snapshot: t.Callable): model = SqlModel( name="name", - kind=dict(time_column="ds", batch_size=30, name=ModelKindName.INCREMENTAL_BY_TIME_RANGE), + kind=dict( + time_column="ds", + batch_size=30, + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + ), owner="owner", dialect="spark", cron="1 0 * * *", @@ -208,14 +198,19 @@ def test_json_with_grants(make_snapshot: t.Callable): reparsed_snapshot = Snapshot.model_validate_json(json_str) assert isinstance(reparsed_snapshot.node, SqlModel) - assert reparsed_snapshot.node.grants == {"SELECT": ["role1", "role2"], "INSERT": ["role3"]} + assert reparsed_snapshot.node.grants == { + "SELECT": ["role1", "role2"], + "INSERT": ["role3"], + } assert reparsed_snapshot.node.grants_target_layer == GrantsTargetLayer.VIRTUAL def test_json_custom_materialization(make_snapshot: t.Callable): model = SqlModel( name="name", - kind=dict(name=ModelKindName.CUSTOM, materialization="non_existent_should_still_work"), + kind=dict( + name=ModelKindName.CUSTOM, materialization="non_existent_should_still_work" + ), owner="owner", dialect="spark", cron="1 0 * * *", @@ -247,10 +242,14 @@ def test_add_interval(snapshot: Snapshot, make_snapshot): snapshot.add_interval("2020-01-02", "2020-01-01") snapshot.add_interval("2020-01-01", "2020-01-01") - assert snapshot.intervals == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] + assert snapshot.intervals == [ + (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) + ] snapshot.add_interval("2020-01-02", "2020-01-02") - assert snapshot.intervals == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-03"))] + assert snapshot.intervals == [ + (to_timestamp("2020-01-01"), to_timestamp("2020-01-03")) + ] snapshot.add_interval("2020-01-04", "2020-01-05") assert snapshot.intervals == [ @@ -306,22 +305,34 @@ def test_add_interval_dev(snapshot: Snapshot, make_snapshot): snapshot.forward_only = True snapshot.add_interval("2020-01-01", "2020-01-01") - assert snapshot.intervals == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] + assert snapshot.intervals == [ + (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) + ] snapshot.add_interval("2020-01-02", "2020-01-02", is_dev=True) - assert snapshot.intervals == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] - assert snapshot.dev_intervals == [(to_timestamp("2020-01-02"), to_timestamp("2020-01-03"))] + assert snapshot.intervals == [ + (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) + ] + assert snapshot.dev_intervals == [ + (to_timestamp("2020-01-02"), to_timestamp("2020-01-03")) + ] new_snapshot = make_snapshot(snapshot.model) new_snapshot.merge_intervals(snapshot) - assert new_snapshot.intervals == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] + assert new_snapshot.intervals == [ + (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) + ] assert new_snapshot.dev_intervals == [] new_snapshot = make_snapshot(snapshot.model) new_snapshot.dev_version_ = snapshot.dev_version new_snapshot.merge_intervals(snapshot) - assert new_snapshot.intervals == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] - assert new_snapshot.dev_intervals == [(to_timestamp("2020-01-02"), to_timestamp("2020-01-03"))] + assert new_snapshot.intervals == [ + (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) + ] + assert new_snapshot.dev_intervals == [ + (to_timestamp("2020-01-02"), to_timestamp("2020-01-03")) + ] def test_add_interval_partial(snapshot: Snapshot, make_snapshot): @@ -362,7 +373,10 @@ def test_get_next_model_interval_start(make_snapshot): daily_snapshot = make_snapshot( SqlModel( - name="early", kind=FullKind(), query=parse_one("SELECT 1, ds FROM name"), cron="@daily" + name="early", + kind=FullKind(), + query=parse_one("SELECT 1, ds FROM name"), + cron="@daily", ) ) @@ -412,7 +426,9 @@ def test_missing_intervals(snapshot: Snapshot): (to_timestamp("2020-01-06"), to_timestamp("2020-01-07")), (to_timestamp("2020-01-07"), to_timestamp("2020-01-08")), ] - assert snapshot.missing_intervals("2020-01-03 00:00:01", "2020-01-05 00:00:02") == [] + assert ( + snapshot.missing_intervals("2020-01-03 00:00:01", "2020-01-05 00:00:02") == [] + ) assert snapshot.missing_intervals("2020-01-03 00:00:01", "2020-01-07 00:00:02") == [ (to_timestamp("2020-01-06"), to_timestamp("2020-01-07")), (to_timestamp("2020-01-07"), to_timestamp("2020-01-08")), @@ -437,17 +453,25 @@ def test_missing_intervals_partial(make_snapshot): (to_timestamp(start), end_ts), ] assert snapshot.missing_intervals(start, end_ts, execution_time=end_ts) == [] - assert snapshot.missing_intervals(start, end_ts, execution_time=end_ts, ignore_cron=True) == [ - (to_timestamp(start), end_ts) - ] + assert snapshot.missing_intervals( + start, end_ts, execution_time=end_ts, ignore_cron=True + ) == [(to_timestamp(start), end_ts)] assert snapshot.missing_intervals(start, end_ts, execution_time="2023-01-02") == [ (to_timestamp(start), end_ts) ] assert snapshot.missing_intervals(start, start) == [ (to_timestamp(start), to_timestamp("2023-01-02")), ] - assert snapshot.missing_intervals(start, start, execution_time=start, ignore_cron=True) == [] - assert snapshot.missing_intervals(start, start, execution_time=end_ts, end_bounded=True) == [] + assert ( + snapshot.missing_intervals(start, start, execution_time=start, ignore_cron=True) + == [] + ) + assert ( + snapshot.missing_intervals( + start, start, execution_time=end_ts, end_bounded=True + ) + == [] + ) assert snapshot.missing_intervals(start, to_timestamp("2023-01-02 12:00:00")) == [ (to_timestamp(start), to_timestamp("2023-01-02")), @@ -459,7 +483,9 @@ def test_missing_intervals_end_bounded_with_lookback(make_snapshot): snapshot = make_snapshot( SqlModel( name="test_model", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds"), lookback=1), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds"), lookback=1 + ), owner="owner", cron="@daily", query=parse_one("SELECT 1, ds FROM name"), @@ -509,7 +535,11 @@ def test_missing_intervals_end_bounded_with_ignore_cron(make_snapshot): == [] ) assert snapshot.missing_intervals( - start, to_datetime(end), execution_time=execution_ts, ignore_cron=True, end_bounded=True + start, + to_datetime(end), + execution_time=execution_ts, + ignore_cron=True, + end_bounded=True, ) == [ (to_timestamp("2023-01-02"), to_timestamp(end)), ] @@ -519,7 +549,9 @@ def test_missing_intervals_past_end_date_with_lookback(make_snapshot): snapshot: Snapshot = make_snapshot( # type: ignore SqlModel( name="test_model", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds"), lookback=2), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds"), lookback=2 + ), owner="owner", cron="@daily", query=parse_one("SELECT 1, ds FROM name"), @@ -538,7 +570,9 @@ def test_missing_intervals_past_end_date_with_lookback(make_snapshot): ) # baseline - all intervals missing - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [ + assert snapshot.missing_intervals( + start_time, end_time, execution_time=end_time + ) == [ (to_timestamp("2023-01-01"), to_timestamp("2023-01-02")), (to_timestamp("2023-01-02"), to_timestamp("2023-01-03")), (to_timestamp("2023-01-03"), to_timestamp("2023-01-04")), @@ -551,13 +585,19 @@ def test_missing_intervals_past_end_date_with_lookback(make_snapshot): # even though lookback=2, because every interval has been filled, # there should be no missing intervals - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + assert ( + snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + ) # however, when running for a new interval, this triggers lookback # in this case, we remove the most recent interval (the one for 2023-01-05) to simulate it being new # since lookback=2 days, this triggers missing intervals for 2023-01-03, 2023-01-04, 2023-01-05 - snapshot.remove_interval(interval=(to_timestamp("2023-01-05"), to_timestamp("2023-01-06"))) - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [ + snapshot.remove_interval( + interval=(to_timestamp("2023-01-05"), to_timestamp("2023-01-06")) + ) + assert snapshot.missing_intervals( + start_time, end_time, execution_time=end_time + ) == [ (to_timestamp("2023-01-03"), to_timestamp("2023-01-04")), (to_timestamp("2023-01-04"), to_timestamp("2023-01-05")), (to_timestamp("2023-01-05"), to_timestamp("2023-01-06")), @@ -565,13 +605,17 @@ def test_missing_intervals_past_end_date_with_lookback(make_snapshot): # put the interval we just removed back to make the model fully backfilled again snapshot.add_interval(to_timestamp("2023-01-05"), to_timestamp("2023-01-06")) - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + assert ( + snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + ) # running on the end date + 1 day (2023-01-07) # 2023-01-06 "would" run and since lookback=2 this pulls in 2023-01-04 and 2023-01-05 as well # however, only 2023-01-04 and 2023-01-05 are within the model end date end_time = to_timestamp("2023-01-07") - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [ + assert snapshot.missing_intervals( + start_time, end_time, execution_time=end_time + ) == [ (to_timestamp("2023-01-04"), to_timestamp("2023-01-05")), (to_timestamp("2023-01-05"), to_timestamp("2023-01-06")), ] @@ -580,24 +624,29 @@ def test_missing_intervals_past_end_date_with_lookback(make_snapshot): # 2023-01-07 "would" run and since lookback=2 this pulls in 2023-01-06 and 2023-01-05 as well # however, only 2023-01-05 is within the model end date end_time = to_timestamp("2023-01-08") - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [ - (to_timestamp("2023-01-05"), to_timestamp("2023-01-06")) - ] + assert snapshot.missing_intervals( + start_time, end_time, execution_time=end_time + ) == [(to_timestamp("2023-01-05"), to_timestamp("2023-01-06"))] # running on the end date + 3 days (2023-01-09) # no missing intervals because subtracting 2 days for lookback exceeds the models end date end_time = to_timestamp("2023-01-09") - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + assert ( + snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + ) # running way in the future, no missing intervals because subtracting 2 days for lookback still exceeds the models end date end_time = to_timestamp("2024-01-01") - assert snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + assert ( + snapshot.missing_intervals(start_time, end_time, execution_time=end_time) == [] + ) -def test_missing_intervals_start_override_per_model(make_snapshot: t.Callable[..., Snapshot]): +def test_missing_intervals_start_override_per_model( + make_snapshot: t.Callable[..., Snapshot], +): snapshot = make_snapshot( - load_sql_based_model( - parse(""" + load_sql_based_model(parse(""" MODEL ( name a, kind FULL, @@ -605,14 +654,15 @@ def test_missing_intervals_start_override_per_model(make_snapshot: t.Callable[.. cron '@daily' ); SELECT 1; - """) - ), + """)), version="a", ) # base case - no override assert list( - missing_intervals(execution_time="2023-02-08 00:05:07", snapshots=[snapshot]).values() + missing_intervals( + execution_time="2023-02-08 00:05:07", snapshots=[snapshot] + ).values() )[0] == [ (to_timestamp("2023-02-01"), to_timestamp("2023-02-02")), (to_timestamp("2023-02-02"), to_timestamp("2023-02-03")), @@ -629,7 +679,9 @@ def test_missing_intervals_start_override_per_model(make_snapshot: t.Callable[.. start="1 day ago", execution_time="2023-02-08 00:05:07", snapshots=[snapshot], - start_override_per_model={snapshot.name: to_datetime("2023-02-05 00:00:00")}, + start_override_per_model={ + snapshot.name: to_datetime("2023-02-05 00:00:00") + }, ).values() )[0] == [ (to_timestamp("2023-02-05"), to_timestamp("2023-02-06")), @@ -642,7 +694,9 @@ def test_incremental_time_self_reference(make_snapshot): snapshot = make_snapshot( SqlModel( name="name", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds"), batch_size=1), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds"), batch_size=1 + ), owner="owner", dialect="", cron="@daily", @@ -654,7 +708,9 @@ def test_incremental_time_self_reference(make_snapshot): assert snapshot.missing_intervals(to_date("1 week ago"), to_date("1 day ago")) == [] # Remove should take away not only 3 days ago but also everything after since this model # depends on past - interval = snapshot.get_removal_interval(to_date("3 days ago"), to_date("3 days ago")) + interval = snapshot.get_removal_interval( + to_date("3 days ago"), to_date("3 days ago") + ) snapshot.remove_interval(interval) assert snapshot.missing_intervals(to_date("1 week ago"), to_date("1 day ago")) == [ (to_timestamp(to_date("3 days ago")), to_timestamp(to_date("2 days ago"))), @@ -700,17 +756,27 @@ def test_lookback(make_snapshot): ] snapshot.add_interval("2023-01-28", "2023-01-29") - assert snapshot.missing_intervals("2023-01-27", "2023-01-27", "2023-01-30 05:00:00") == [ + assert snapshot.missing_intervals( + "2023-01-27", "2023-01-27", "2023-01-30 05:00:00" + ) == [ (to_timestamp("2023-01-27"), to_timestamp("2023-01-28")), ] - assert snapshot.missing_intervals("2023-01-28", "2023-01-29", "2023-01-31 05:00:00") == [] - assert snapshot.missing_intervals("2023-01-28", "2023-01-30", "2023-01-31 05:00:00") == [ + assert ( + snapshot.missing_intervals("2023-01-28", "2023-01-29", "2023-01-31 05:00:00") + == [] + ) + assert snapshot.missing_intervals( + "2023-01-28", "2023-01-30", "2023-01-31 05:00:00" + ) == [ (to_timestamp("2023-01-28"), to_timestamp("2023-01-29")), (to_timestamp("2023-01-29"), to_timestamp("2023-01-30")), (to_timestamp("2023-01-30"), to_timestamp("2023-01-31")), ] - assert snapshot.missing_intervals("2023-01-28", "2023-01-30", "2023-01-31 04:00:00") == [] + assert ( + snapshot.missing_intervals("2023-01-28", "2023-01-30", "2023-01-31 04:00:00") + == [] + ) def test_lookback_custom_materialization(make_snapshot): @@ -719,8 +785,7 @@ def test_lookback_custom_materialization(make_snapshot): class MyTestStrategy(CustomMaterialization): pass - expressions = parse( - """ + expressions = parse(""" MODEL ( name name, kind CUSTOM ( @@ -732,8 +797,7 @@ class MyTestStrategy(CustomMaterialization): ); SELECT ds FROM parent.tbl - """ - ) + """) snapshot = make_snapshot(load_sql_based_model(expressions)) @@ -781,10 +845,13 @@ def test_missing_interval_smaller_than_interval_unit(make_snapshot): ) ) - assert snapshot_daily.missing_intervals("2020-01-01 00:00:05", "2020-01-01 23:59:59") == [] - assert snapshot_daily.missing_intervals("2020-01-01 00:00:00", "2020-01-02 00:00:00") == [ - (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) - ] + assert ( + snapshot_daily.missing_intervals("2020-01-01 00:00:05", "2020-01-01 23:59:59") + == [] + ) + assert snapshot_daily.missing_intervals( + "2020-01-01 00:00:00", "2020-01-02 00:00:00" + ) == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] snapshot_hourly = make_snapshot( SqlModel( @@ -798,10 +865,13 @@ def test_missing_interval_smaller_than_interval_unit(make_snapshot): ) ) - assert snapshot_hourly.missing_intervals("2020-01-01 00:00:00", "2020-01-01 00:59:00") == [] - assert snapshot_hourly.missing_intervals("2020-01-01 00:00:00", "2020-01-01 01:00:00") == [ - (to_timestamp("2020-01-01"), to_timestamp("2020-01-01 01:00:00")) - ] + assert ( + snapshot_hourly.missing_intervals("2020-01-01 00:00:00", "2020-01-01 00:59:00") + == [] + ) + assert snapshot_hourly.missing_intervals( + "2020-01-01 00:00:00", "2020-01-01 01:00:00" + ) == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-01 01:00:00"))] snapshot_end_categorical = make_snapshot( SqlModel( @@ -815,9 +885,9 @@ def test_missing_interval_smaller_than_interval_unit(make_snapshot): ) ) - assert snapshot_end_categorical.missing_intervals("2020-01-01 00:00:00", "2020-01-01") == [ - (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) - ] + assert snapshot_end_categorical.missing_intervals( + "2020-01-01 00:00:00", "2020-01-01" + ) == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] snapshot_partial = make_snapshot( SqlModel( @@ -832,12 +902,12 @@ def test_missing_interval_smaller_than_interval_unit(make_snapshot): ) ) - assert snapshot_partial.missing_intervals("2020-01-01 00:00:05", "2020-01-01 23:59:59") == [ - (to_timestamp("2020-01-01"), to_timestamp("2020-01-01 23:59:59")) - ] - assert snapshot_partial.missing_intervals("2020-01-01 00:00:00", "2020-01-02 00:00:00") == [ - (to_timestamp("2020-01-01"), to_timestamp("2020-01-02")) - ] + assert snapshot_partial.missing_intervals( + "2020-01-01 00:00:05", "2020-01-01 23:59:59" + ) == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-01 23:59:59"))] + assert snapshot_partial.missing_intervals( + "2020-01-01 00:00:00", "2020-01-02 00:00:00" + ) == [(to_timestamp("2020-01-01"), to_timestamp("2020-01-02"))] def test_remove_intervals(snapshot): @@ -848,7 +918,9 @@ def test_remove_intervals(snapshot): snapshot.add_interval("2020-01-01", "2020-01-01") snapshot.add_interval("2020-01-03", "2020-01-03") snapshot.remove_interval(snapshot.get_removal_interval("2020-01-01", "2020-01-01")) - assert snapshot.intervals == [(to_timestamp("2020-01-03"), to_timestamp("2020-01-04"))] + assert snapshot.intervals == [ + (to_timestamp("2020-01-03"), to_timestamp("2020-01-04")) + ] snapshot.remove_interval(snapshot.get_removal_interval("2020-01-01", "2020-01-05")) assert snapshot.intervals == [] @@ -989,7 +1061,9 @@ def test_fingerprint(model: Model, parent_model: Model): ) assert fingerprint == original_fingerprint - with_parent_fingerprint = fingerprint_from_node(model, nodes={'"parent"."tbl"': parent_model}) + with_parent_fingerprint = fingerprint_from_node( + model, nodes={'"parent"."tbl"': parent_model} + ) assert with_parent_fingerprint != fingerprint assert int(with_parent_fingerprint.parent_data_hash) > 0 assert int(with_parent_fingerprint.parent_metadata_hash) > 0 @@ -998,7 +1072,9 @@ def test_fingerprint(model: Model, parent_model: Model): fingerprint_from_node( model, nodes={ - '"parent"."tbl"': SqlModel(**{**model.dict(), "query": parse_one("select 2, ds")}) + '"parent"."tbl"': SqlModel( + **{**model.dict(), "query": parse_one("select 2, ds")} + ) }, ) != with_parent_fingerprint @@ -1009,7 +1085,9 @@ def test_fingerprint(model: Model, parent_model: Model): assert new_fingerprint.data_hash != fingerprint.data_hash assert new_fingerprint.metadata_hash != fingerprint.metadata_hash - model = SqlModel(**{**model.dict(), "query": parse_one("select 1, ds -- annotation")}) + model = SqlModel( + **{**model.dict(), "query": parse_one("select 1, ds -- annotation")} + ) fingerprint = fingerprint_from_node(model, nodes={}) assert new_fingerprint != fingerprint assert new_fingerprint.data_hash != fingerprint.data_hash @@ -1024,7 +1102,9 @@ def test_fingerprint(model: Model, parent_model: Model): assert new_fingerprint.metadata_hash != fingerprint.metadata_hash assert fingerprint.metadata_hash != original_fingerprint.metadata_hash - model = SqlModel(**{**original_model.dict(), "post_statements": [parse_one("DROP TABLE test")]}) + model = SqlModel( + **{**original_model.dict(), "post_statements": [parse_one("DROP TABLE test")]} + ) fingerprint = fingerprint_from_node(model, nodes={}) assert new_fingerprint != fingerprint assert new_fingerprint.data_hash != fingerprint.data_hash @@ -1033,23 +1113,23 @@ def test_fingerprint(model: Model, parent_model: Model): def test_fingerprint_seed_model(): - expressions = parse( - """ + expressions = parse(""" MODEL ( name db.seed, kind SEED ( path '../seeds/waiter_names.csv' ) ); - """ - ) + """) expected_fingerprint = SnapshotFingerprint( data_hash="2112858704", metadata_hash="2674364560", ) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) actual_fingerprint = fingerprint_from_node(model, nodes={}) assert actual_fingerprint == expected_fingerprint @@ -1067,7 +1147,9 @@ def test_fingerprint_seed_model(): ) updated_actual_fingerprint = fingerprint_from_node(updated_model, nodes={}) assert updated_actual_fingerprint.data_hash != expected_fingerprint.data_hash - assert updated_actual_fingerprint.metadata_hash == expected_fingerprint.metadata_hash + assert ( + updated_actual_fingerprint.metadata_hash == expected_fingerprint.metadata_hash + ) def test_fingerprint_jinja_macros(model: Model): @@ -1077,7 +1159,8 @@ def test_fingerprint_jinja_macros(model: Model): "jinja_macros": JinjaMacroRegistry( root_macros={ "test_macro": MacroInfo( - definition="{% macro test_macro() %}a{% endmacro %}", depends_on=[] + definition="{% macro test_macro() %}a{% endmacro %}", + depends_on=[], ) } ), @@ -1122,7 +1205,10 @@ def test_fingerprint_builtin_audits(model: Model, parent_model: Model): fingerprint = fingerprint_from_node(model, nodes={}) model = SqlModel.parse_obj( - {**model.dict(), "audits": [("unique_values", {"columns": exp.convert([to_column("a")])})]} + { + **model.dict(), + "audits": [("unique_values", {"columns": exp.convert([to_column("a")])})], + } ) new_fingerprint = fingerprint_from_node(model, nodes={}) assert new_fingerprint != fingerprint @@ -1139,7 +1225,9 @@ def test_fingerprint_standalone_audits(parent_model: Model): name="test_standalone_audit", query=parse_one(f"SELECT colb FROM {parent_model.name} WHERE colb IS NULL"), ) - new_fingerprint = fingerprint_from_node(new_audit, nodes={parent_model.name: parent_model}) + new_fingerprint = fingerprint_from_node( + new_audit, nodes={parent_model.name: parent_model} + ) assert new_fingerprint != fingerprint assert new_fingerprint.data_hash == fingerprint.data_hash @@ -1152,7 +1240,9 @@ def test_fingerprint_virtual_properties(model: Model, parent_model: Model): updated_model = SqlModel( **original_model.dict(), - virtual_properties=parse_one("(labels = [('test-virtual-label', 'label-virtual-value')],)"), + virtual_properties=parse_one( + "(labels = [('test-virtual-label', 'label-virtual-value')],)" + ), ) assert "labels" in updated_model.virtual_properties updated_fingerprint = fingerprint_from_node(updated_model, nodes={}) @@ -1182,9 +1272,13 @@ def test_fingerprint_grants(model: Model, parent_model: Model): **original_model.dict(), grants={"SELECT": ["role3"], "INSERT": ["role4"]}, ) - different_grants_fingerprint = fingerprint_from_node(different_grants_model, nodes={}) + different_grants_fingerprint = fingerprint_from_node( + different_grants_model, nodes={} + ) - assert different_grants_fingerprint.metadata_hash != updated_fingerprint.metadata_hash + assert ( + different_grants_fingerprint.metadata_hash != updated_fingerprint.metadata_hash + ) assert different_grants_fingerprint.metadata_hash != fingerprint.metadata_hash target_layer_model = SqlModel( @@ -1254,11 +1348,19 @@ def test_snapshot_table_name(snapshot: Snapshot, make_snapshot: t.Callable): ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) assert snapshot.table_naming_convention == TableNamingConvention.SCHEMA_AND_TABLE - assert snapshot.data_version.table_naming_convention == TableNamingConvention.SCHEMA_AND_TABLE + assert ( + snapshot.data_version.table_naming_convention + == TableNamingConvention.SCHEMA_AND_TABLE + ) snapshot.previous_versions = () - assert snapshot.table_name(is_deployable=True) == "sqlmesh__default.name__3078928823" - assert snapshot.table_name(is_deployable=False) == "sqlmesh__default.name__3078928823__dev" + assert ( + snapshot.table_name(is_deployable=True) == "sqlmesh__default.name__3078928823" + ) + assert ( + snapshot.table_name(is_deployable=False) + == "sqlmesh__default.name__3078928823__dev" + ) assert snapshot.dev_version == snapshot.fingerprint.to_version() @@ -1270,9 +1372,14 @@ def test_snapshot_table_name(snapshot: Snapshot, make_snapshot: t.Callable): ) snapshot.previous_versions = (previous_data_version,) snapshot.categorize_as(SnapshotChangeCategory.INDIRECT_NON_BREAKING) - assert snapshot.table_name(is_deployable=True) == "sqlmesh__default.name__3078928823" + assert ( + snapshot.table_name(is_deployable=True) == "sqlmesh__default.name__3078928823" + ) # Indirect non-breaking snapshots reuse the dev table as well. - assert snapshot.table_name(is_deployable=False) == "sqlmesh__default.name__3078928823__dev" + assert ( + snapshot.table_name(is_deployable=False) + == "sqlmesh__default.name__3078928823__dev" + ) assert snapshot.dev_version != snapshot.fingerprint.to_version() assert snapshot.dev_version == previous_data_version.dev_version @@ -1282,8 +1389,13 @@ def test_snapshot_table_name(snapshot: Snapshot, make_snapshot: t.Callable): ) snapshot.previous_versions = (previous_data_version,) snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) - assert snapshot.table_name(is_deployable=True) == "sqlmesh__default.name__3078928823" - assert snapshot.table_name(is_deployable=False) == "sqlmesh__default.name__3049392110__dev" + assert ( + snapshot.table_name(is_deployable=True) == "sqlmesh__default.name__3078928823" + ) + assert ( + snapshot.table_name(is_deployable=False) + == "sqlmesh__default.name__3049392110__dev" + ) fully_qualified_snapshot = make_snapshot( SqlModel(name='"my-catalog".db.table', query=parse_one("select 1, ds")) @@ -1295,7 +1407,9 @@ def test_snapshot_table_name(snapshot: Snapshot, make_snapshot: t.Callable): ) non_fully_qualified_snapshot = make_snapshot( SqlModel( - name="db.table", query=parse_one("select 1, ds"), default_catalog='"other-catalog"' + name="db.table", + query=parse_one("select 1, ds"), + default_catalog='"other-catalog"', ) ) non_fully_qualified_snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -1305,7 +1419,9 @@ def test_snapshot_table_name(snapshot: Snapshot, make_snapshot: t.Callable): ) -def test_table_name_naming_convention_table_only(make_snapshot: t.Callable[..., Snapshot]): +def test_table_name_naming_convention_table_only( + make_snapshot: t.Callable[..., Snapshot], +): # 3-part naming snapshot = make_snapshot( SqlModel(name='"foo"."bar"."baz"', query=parse_one("select 1")), @@ -1313,11 +1429,18 @@ def test_table_name_naming_convention_table_only(make_snapshot: t.Callable[..., ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) assert snapshot.table_naming_convention == TableNamingConvention.TABLE_ONLY - assert snapshot.data_version.table_naming_convention == TableNamingConvention.TABLE_ONLY + assert ( + snapshot.data_version.table_naming_convention + == TableNamingConvention.TABLE_ONLY + ) - assert snapshot.table_name(is_deployable=True) == f"foo.sqlmesh__bar.baz__{snapshot.version}" assert ( - snapshot.table_name(is_deployable=False) == f"foo.sqlmesh__bar.baz__{snapshot.version}__dev" + snapshot.table_name(is_deployable=True) + == f"foo.sqlmesh__bar.baz__{snapshot.version}" + ) + assert ( + snapshot.table_name(is_deployable=False) + == f"foo.sqlmesh__bar.baz__{snapshot.version}__dev" ) # 2-part naming @@ -1327,11 +1450,19 @@ def test_table_name_naming_convention_table_only(make_snapshot: t.Callable[..., ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - assert snapshot.table_name(is_deployable=True) == f"sqlmesh__foo.bar__{snapshot.version}" - assert snapshot.table_name(is_deployable=False) == f"sqlmesh__foo.bar__{snapshot.version}__dev" + assert ( + snapshot.table_name(is_deployable=True) + == f"sqlmesh__foo.bar__{snapshot.version}" + ) + assert ( + snapshot.table_name(is_deployable=False) + == f"sqlmesh__foo.bar__{snapshot.version}__dev" + ) -def test_table_name_naming_convention_hash_md5(make_snapshot: t.Callable[..., Snapshot]): +def test_table_name_naming_convention_hash_md5( + make_snapshot: t.Callable[..., Snapshot], +): # 3-part naming snapshot = make_snapshot( SqlModel(name='"foo"."bar"."baz"', query=parse_one("select 1")), @@ -1339,13 +1470,19 @@ def test_table_name_naming_convention_hash_md5(make_snapshot: t.Callable[..., Sn ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) assert snapshot.table_naming_convention == TableNamingConvention.HASH_MD5 - assert snapshot.data_version.table_naming_convention == TableNamingConvention.HASH_MD5 + assert ( + snapshot.data_version.table_naming_convention == TableNamingConvention.HASH_MD5 + ) hash = md5(f"foo.sqlmesh__bar.bar__baz__{snapshot.version}") - assert snapshot.table_name(is_deployable=True) == f"foo.sqlmesh__bar.sqlmesh_md5__{hash}" + assert ( + snapshot.table_name(is_deployable=True) + == f"foo.sqlmesh__bar.sqlmesh_md5__{hash}" + ) hash_dev = md5(f"foo.sqlmesh__bar.bar__baz__{snapshot.version}__dev") assert ( - snapshot.table_name(is_deployable=False) == f"foo.sqlmesh__bar.sqlmesh_md5__{hash_dev}__dev" + snapshot.table_name(is_deployable=False) + == f"foo.sqlmesh__bar.sqlmesh_md5__{hash_dev}__dev" ) # 2-part naming @@ -1356,13 +1493,20 @@ def test_table_name_naming_convention_hash_md5(make_snapshot: t.Callable[..., Sn snapshot.categorize_as(SnapshotChangeCategory.BREAKING) hash = md5(f"sqlmesh__foo.foo__bar__{snapshot.version}") - assert snapshot.table_name(is_deployable=True) == f"sqlmesh__foo.sqlmesh_md5__{hash}" + assert ( + snapshot.table_name(is_deployable=True) == f"sqlmesh__foo.sqlmesh_md5__{hash}" + ) hash_dev = md5(f"sqlmesh__foo.foo__bar__{snapshot.version}__dev") - assert snapshot.table_name(is_deployable=False) == f"sqlmesh__foo.sqlmesh_md5__{hash_dev}__dev" + assert ( + snapshot.table_name(is_deployable=False) + == f"sqlmesh__foo.sqlmesh_md5__{hash_dev}__dev" + ) -def test_table_naming_convention_passed_around_correctly(make_snapshot: t.Callable[..., Snapshot]): +def test_table_naming_convention_passed_around_correctly( + make_snapshot: t.Callable[..., Snapshot], +): snapshot = make_snapshot( SqlModel(name='"foo"."bar"."baz"', query=parse_one("select 1")), table_naming_convention=TableNamingConvention.HASH_MD5, @@ -1370,12 +1514,18 @@ def test_table_naming_convention_passed_around_correctly(make_snapshot: t.Callab snapshot.categorize_as(SnapshotChangeCategory.BREAKING) assert snapshot.table_naming_convention == TableNamingConvention.HASH_MD5 - assert snapshot.data_version.table_naming_convention == TableNamingConvention.HASH_MD5 + assert ( + snapshot.data_version.table_naming_convention == TableNamingConvention.HASH_MD5 + ) assert snapshot.table_info.table_naming_convention == TableNamingConvention.HASH_MD5 assert ( - snapshot.table_info.data_version.table_naming_convention == TableNamingConvention.HASH_MD5 + snapshot.table_info.data_version.table_naming_convention + == TableNamingConvention.HASH_MD5 + ) + assert ( + snapshot.table_info.table_info.table_naming_convention + == TableNamingConvention.HASH_MD5 ) - assert snapshot.table_info.table_info.table_naming_convention == TableNamingConvention.HASH_MD5 assert ( snapshot.table_info.table_info.data_version.table_naming_convention == TableNamingConvention.HASH_MD5 @@ -1384,10 +1534,15 @@ def test_table_naming_convention_passed_around_correctly(make_snapshot: t.Callab def test_table_name_view(make_snapshot: t.Callable): # Mimic a direct breaking change. - snapshot = make_snapshot(SqlModel(name="name", query=parse_one("select 1"), kind="VIEW")) + snapshot = make_snapshot( + SqlModel(name="name", query=parse_one("select 1"), kind="VIEW") + ) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) snapshot.previous_versions = () - assert snapshot.table_name(is_deployable=True) == f"sqlmesh__default.name__{snapshot.version}" + assert ( + snapshot.table_name(is_deployable=True) + == f"sqlmesh__default.name__{snapshot.version}" + ) assert ( snapshot.table_name(is_deployable=False) == f"sqlmesh__default.name__{snapshot.dev_version}__dev" @@ -1396,12 +1551,15 @@ def test_table_name_view(make_snapshot: t.Callable): assert snapshot.dev_version == snapshot.fingerprint.to_version() # Mimic an indirect non-breaking change. - new_snapshot = make_snapshot(SqlModel(name="name", query=parse_one("select 2"), kind="VIEW")) + new_snapshot = make_snapshot( + SqlModel(name="name", query=parse_one("select 2"), kind="VIEW") + ) previous_data_version = snapshot.data_version new_snapshot.previous_versions = (previous_data_version,) new_snapshot.categorize_as(SnapshotChangeCategory.INDIRECT_NON_BREAKING) assert ( - new_snapshot.table_name(is_deployable=True) == f"sqlmesh__default.name__{snapshot.version}" + new_snapshot.table_name(is_deployable=True) + == f"sqlmesh__default.name__{snapshot.version}" ) # Indirect non-breaking view snapshots should not reuse the dev table. assert ( @@ -1421,8 +1579,14 @@ def test_table_naming_convention_change_reuse_previous_version(make_snapshot): ) original_snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - assert original_snapshot.table_naming_convention == TableNamingConvention.SCHEMA_AND_TABLE - assert original_snapshot.table_name() == f"sqlmesh__default.a__{original_snapshot.version}" + assert ( + original_snapshot.table_naming_convention + == TableNamingConvention.SCHEMA_AND_TABLE + ) + assert ( + original_snapshot.table_name() + == f"sqlmesh__default.a__{original_snapshot.version}" + ) changed_snapshot: Snapshot = make_snapshot( SqlModel(name="a", query=parse_one("select 1, 'forward_only' as a, ds")), @@ -1435,12 +1599,18 @@ def test_table_naming_convention_change_reuse_previous_version(make_snapshot): changed_snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) # inherited from previous version even though changed_snapshot was created with TableNamingConvention.HASH_MD5 - assert changed_snapshot.table_naming_convention == TableNamingConvention.SCHEMA_AND_TABLE + assert ( + changed_snapshot.table_naming_convention + == TableNamingConvention.SCHEMA_AND_TABLE + ) assert ( changed_snapshot.previous_version.table_naming_convention == TableNamingConvention.SCHEMA_AND_TABLE ) - assert changed_snapshot.table_name() == f"sqlmesh__default.a__{changed_snapshot.version}" + assert ( + changed_snapshot.table_name() + == f"sqlmesh__default.a__{changed_snapshot.version}" + ) def test_categorize_change_sql(make_snapshot): @@ -1463,7 +1633,8 @@ def test_categorize_change_sql(make_snapshot): categorize_change( new=make_snapshot( SqlModel( - name="a", query=parse_one("select 1, fun(another_fun(a + 1) * 2)::INT, ds") + name="a", + query=parse_one("select 1, fun(another_fun(a + 1) * 2)::INT, ds"), ) ), old=old_snapshot, @@ -1475,7 +1646,9 @@ def test_categorize_change_sql(make_snapshot): # Multiple projections have been added. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select 1, 2, a, b, ds"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select 1, 2, a, b, ds")) + ), old=old_snapshot, config=config, ) @@ -1519,7 +1692,9 @@ def test_categorize_change_sql(make_snapshot): # A WHERE clause has been added. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds WHERE a = 2"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select 1, ds WHERE a = 2")) + ), old=old_snapshot, config=config, ) @@ -1529,7 +1704,9 @@ def test_categorize_change_sql(make_snapshot): # A FROM clause has been added. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds FROM test_table"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select 1, ds FROM test_table")) + ), old=old_snapshot, config=config, ) @@ -1539,7 +1716,9 @@ def test_categorize_change_sql(make_snapshot): # DISTINCT has been added. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select DISTINCT 1, ds"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select DISTINCT 1, ds")) + ), old=old_snapshot, config=config, ) @@ -1549,7 +1728,9 @@ def test_categorize_change_sql(make_snapshot): # An EXPLODE projection has been added. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds, explode(a)"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select 1, ds, explode(a)")) + ), old=old_snapshot, config=config, ) @@ -1571,7 +1752,9 @@ def test_categorize_change_sql(make_snapshot): # A POSEXPLODE projection has been added. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds, posexplode(a)"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select 1, ds, posexplode(a)")) + ), old=old_snapshot, config=config, ) @@ -1593,7 +1776,9 @@ def test_categorize_change_sql(make_snapshot): # An UNNEST projection has been added. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds, unnest(a)"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select 1, ds, unnest(a)")) + ), old=old_snapshot, config=config, ) @@ -1603,7 +1788,9 @@ def test_categorize_change_sql(make_snapshot): # A metadata change occurred assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("select 1, ds"), owner="foo")), + new=make_snapshot( + SqlModel(name="a", query=parse_one("select 1, ds"), owner="foo") + ), old=old_snapshot, config=config, ) @@ -1614,7 +1801,10 @@ def test_categorize_change_sql(make_snapshot): assert ( categorize_change( new=make_snapshot( - SqlModel(name="a", query=parse_one("select 1, ds, (select x from unnest(a) x)")) + SqlModel( + name="a", + query=parse_one("select 1, ds, (select x from unnest(a) x)"), + ) ), old=old_snapshot, config=config, @@ -1624,7 +1814,10 @@ def test_categorize_change_sql(make_snapshot): assert ( categorize_change( new=make_snapshot( - SqlModel(name="a", query=parse_one("select 1, ds, (select x from posexplode(a) x)")) + SqlModel( + name="a", + query=parse_one("select 1, ds, (select x from posexplode(a) x)"), + ) ), old=old_snapshot, config=config, @@ -1644,7 +1837,9 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): assert ( categorize_change( new=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::DATE, x::TEXT, s::TEXT FROM t")) + SqlModel( + name="a", query=parse_one("SELECT a::DATE, x::TEXT, s::TEXT FROM t") + ) ), old=old_snapshot, config=config, @@ -1656,7 +1851,9 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): assert ( categorize_change( new=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT s::TEXT, a::DATE, s::TEXT FROM t")) + SqlModel( + name="a", query=parse_one("SELECT s::TEXT, a::DATE, s::TEXT FROM t") + ) ), old=old_snapshot, config=config, @@ -1669,7 +1866,8 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): categorize_change( new=make_snapshot( SqlModel( - name="a", query=parse_one("SELECT a::DATE, x::TEXT, s::TEXT, y::INT FROM t") + name="a", + query=parse_one("SELECT a::DATE, x::TEXT, s::TEXT, y::INT FROM t"), ) ), old=old_snapshot, @@ -1682,7 +1880,9 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): assert ( categorize_change( new=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::INT, x::TEXT, s::TEXT FROM t")) + SqlModel( + name="a", query=parse_one("SELECT a::INT, x::TEXT, s::TEXT FROM t") + ) ), old=old_snapshot, config=config, @@ -1705,7 +1905,9 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): # An existing projection has been removed: undetermined. assert ( categorize_change( - new=make_snapshot(SqlModel(name="a", query=parse_one("SELECT s::TEXT FROM t"))), + new=make_snapshot( + SqlModel(name="a", query=parse_one("SELECT s::TEXT FROM t")) + ), old=old_snapshot, config=config, ) @@ -1718,11 +1920,16 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): new=make_snapshot( SqlModel( name="a", - query=parse_one("SELECT a::DATE, x::TEXT, s::TEXT FROM t WHERE a = 2"), + query=parse_one( + "SELECT a::DATE, x::TEXT, s::TEXT FROM t WHERE a = 2" + ), ) ), old=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::DATE, s::TEXT FROM t WHERE a = 1")) + SqlModel( + name="a", + query=parse_one("SELECT a::DATE, s::TEXT FROM t WHERE a = 1"), + ) ), config=config, ) @@ -1735,11 +1942,16 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): new=make_snapshot( SqlModel( name="a", - query=parse_one("SELECT a::DATE, x::TEXT, s::TEXT FROM t ORDER BY 2"), + query=parse_one( + "SELECT a::DATE, x::TEXT, s::TEXT FROM t ORDER BY 2" + ), ) ), old=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::DATE, s::TEXT FROM t ORDER BY 2")) + SqlModel( + name="a", + query=parse_one("SELECT a::DATE, s::TEXT FROM t ORDER BY 2"), + ) ), config=config, ) @@ -1752,11 +1964,16 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): new=make_snapshot( SqlModel( name="a", - query=parse_one("SELECT a::DATE, s::TEXT, x::TEXT FROM t ORDER BY 2"), + query=parse_one( + "SELECT a::DATE, s::TEXT, x::TEXT FROM t ORDER BY 2" + ), ) ), old=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::DATE, s::TEXT FROM t ORDER BY 2")) + SqlModel( + name="a", + query=parse_one("SELECT a::DATE, s::TEXT FROM t ORDER BY 2"), + ) ), config=config, ) @@ -1769,11 +1986,16 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): new=make_snapshot( SqlModel( name="a", - query=parse_one("SELECT a::DATE, x::TEXT, s::TEXT FROM t GROUP BY 2"), + query=parse_one( + "SELECT a::DATE, x::TEXT, s::TEXT FROM t GROUP BY 2" + ), ) ), old=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::DATE, s::TEXT FROM t GROUP BY 2")) + SqlModel( + name="a", + query=parse_one("SELECT a::DATE, s::TEXT FROM t GROUP BY 2"), + ) ), config=config, ) @@ -1792,7 +2014,10 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): ) ), old=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::DATE AS a, s::TEXT AS s FROM t")) + SqlModel( + name="a", + query=parse_one("SELECT a::DATE AS a, s::TEXT AS s FROM t"), + ) ), config=config, ) @@ -1811,7 +2036,10 @@ def test_categorize_change_sql_additive_projection_edge_cases(make_snapshot): ) ), old=make_snapshot( - SqlModel(name="a", query=parse_one("SELECT a::DATE AS a, s::TEXT AS s FROM t")) + SqlModel( + name="a", + query=parse_one("SELECT a::DATE AS a, s::TEXT AS s FROM t"), + ) ), config=config, ) @@ -1910,23 +2138,19 @@ def test_categorize_change_seed(make_snapshot, tmp_path): model_kind = SeedKind(path=str(seed_path.absolute())) with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_c 1,text_a,1.0 -2,text_b,2.0""" - ) +2,text_b,2.0""") original_snapshot = make_snapshot(create_seed_model(model_name, model_kind)) # New column. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_d,col_b,col_c 1,,text_a,1.0 -2,test,text_b,2.0""" - ) +2,test,text_b,2.0""") assert ( categorize_change( @@ -1939,12 +2163,10 @@ def test_categorize_change_seed(make_snapshot, tmp_path): # Column removed. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_c 1,1.0 -2,2.0""" - ) +2,2.0""") assert ( categorize_change( @@ -1957,12 +2179,10 @@ def test_categorize_change_seed(make_snapshot, tmp_path): # Column renamed. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_d 1,text_a,1.0 -2,text_b,2.0""" - ) +2,text_b,2.0""") assert ( categorize_change( @@ -1975,13 +2195,11 @@ def test_categorize_change_seed(make_snapshot, tmp_path): # New row. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_c 1,text_a,1.0 2,text_b,2.0 -3,text_c,3.0""" - ) +3,text_c,3.0""") assert ( categorize_change( @@ -1994,11 +2212,9 @@ def test_categorize_change_seed(make_snapshot, tmp_path): # Deleted row. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_c -1,text_a,1.0""" - ) +1,text_a,1.0""") assert ( categorize_change( @@ -2011,12 +2227,10 @@ def test_categorize_change_seed(make_snapshot, tmp_path): # Numeric column changed. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_c 1,text_a,1.0 -2,text_b,3.0""" - ) +2,text_b,3.0""") assert ( categorize_change( @@ -2029,12 +2243,10 @@ def test_categorize_change_seed(make_snapshot, tmp_path): # Text column changed. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_c 1,text_a,1.0 -2,text_c,2.0""" - ) +2,text_c,2.0""") assert ( categorize_change( @@ -2047,12 +2259,10 @@ def test_categorize_change_seed(make_snapshot, tmp_path): # Column type changed. with open(seed_path, "w", encoding="utf-8") as fd: - fd.write( - """ + fd.write(""" col_a,col_b,col_c 1,text_a,1.0 -2.0,text_b,2.0""" - ) +2.0,text_b,2.0""") assert ( categorize_change( @@ -2128,7 +2338,9 @@ def test_inclusive_exclusive_monthly(make_snapshot): snapshot = make_snapshot( SqlModel( name="name", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds"), batch_size=1), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds"), batch_size=1 + ), owner="owner", dialect="", cron="@monthly", @@ -2157,7 +2369,9 @@ def test_inclusive_exclusive_hourly(make_snapshot): snapshot = make_snapshot( SqlModel( name="name", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds"), batch_size=1), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds"), batch_size=1 + ), owner="owner", dialect="", cron="@hourly", @@ -2196,7 +2410,9 @@ def test_model_custom_cron(make_snapshot): (to_timestamp("2023-01-29"), to_timestamp("2023-01-30")), ] assert snapshot.missing_intervals( - to_timestamp("2023-01-29"), to_timestamp("2023-01-30"), execution_time="2023-01-30 05:00:00" + to_timestamp("2023-01-29"), + to_timestamp("2023-01-30"), + execution_time="2023-01-30 05:00:00", ) == [ (to_timestamp("2023-01-29"), to_timestamp("2023-01-30")), ] @@ -2208,14 +2424,18 @@ def test_model_custom_cron(make_snapshot): (to_timestamp("2023-01-29"), to_timestamp("2023-01-30")), ] assert snapshot.missing_intervals( - to_timestamp("2023-01-29"), to_timestamp("2023-01-30"), execution_time="2023-01-30 05:01:00" + to_timestamp("2023-01-29"), + to_timestamp("2023-01-30"), + execution_time="2023-01-30 05:01:00", ) == [ (to_timestamp("2023-01-29"), to_timestamp("2023-01-30")), ] # Run at 4:59AM assert ( - snapshot.missing_intervals("2023-01-29", "2023-01-29", execution_time="2023-01-30 04:59:00") + snapshot.missing_intervals( + "2023-01-29", "2023-01-29", execution_time="2023-01-30 04:59:00" + ) == [] ) assert ( @@ -2229,7 +2449,10 @@ def test_model_custom_cron(make_snapshot): # Run at 4:59AM and ignore cron assert snapshot.missing_intervals( - "2023-01-29", "2023-01-29", execution_time="2023-01-30 04:59:00", ignore_cron=True + "2023-01-29", + "2023-01-29", + execution_time="2023-01-30 04:59:00", + ignore_cron=True, ) == [ (to_timestamp("2023-01-29"), to_timestamp("2023-01-30")), ] @@ -2297,12 +2520,16 @@ def test_is_valid_start(make_snapshot): ), ( QualifiedViewName(catalog="a-b", schema_name="c-d", table="e-f"), - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.TABLE), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.TABLE + ), '"a-b"."c-d"."e-f__dev"', ), ( QualifiedViewName(catalog="a-b", schema_name="c-d", table="e-f"), - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.SCHEMA), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.SCHEMA + ), '"a-b"."c-d__dev"."e-f"', ), ( @@ -2337,7 +2564,9 @@ def test_is_valid_start(make_snapshot): '"default-foo__dev".sqlmesh_example.full_model', ), ( - QualifiedViewName(catalog="default", schema_name="sqlmesh_example", table="full_model"), + QualifiedViewName( + catalog="default", schema_name="sqlmesh_example", table="full_model" + ), EnvironmentNamingInfo( name=c.PROD, catalog_name_override=None, @@ -2352,10 +2581,16 @@ def test_qualified_view_name(qualified_view_name, environment_naming_info, expec def test_qualified_view_name_with_dialect(): - qualified_view_name = QualifiedViewName(catalog="catalog", schema_name="db", table="table") - environment_naming_info = EnvironmentNamingInfo(name="dev", catalog_name_override="override") + qualified_view_name = QualifiedViewName( + catalog="catalog", schema_name="db", table="table" + ) + environment_naming_info = EnvironmentNamingInfo( + name="dev", catalog_name_override="override" + ) assert ( - qualified_view_name.for_environment(environment_naming_info, dialect="snowflake") + qualified_view_name.for_environment( + environment_naming_info, dialect="snowflake" + ) == "OVERRIDE.db__DEV.table" ) @@ -2399,7 +2634,9 @@ def test_multi_interval_merge(make_snapshot): def test_earliest_start_date(sushi_context: Context): model_name = "sushi.waiter_names" fqn_name = '"memory"."sushi"."waiter_names"' - assert sushi_context.get_snapshot(model_name, raise_if_missing=True).node.start is None + assert ( + sushi_context.get_snapshot(model_name, raise_if_missing=True).node.start is None + ) cache: t.Dict[str, datetime] = {} earliest_start_date(sushi_context.snapshots.values(), cache) @@ -2471,7 +2708,9 @@ def test_deployability_index(make_snapshot): none_deployable_index = deployability_index.none_deployable() assert all(not none_deployable_index.is_deployable(s) for s in snapshots.values()) - assert all(not none_deployable_index.is_representative(s) for s in snapshots.values()) + assert all( + not none_deployable_index.is_representative(s) for s in snapshots.values() + ) def test_deployability_index_unpaused_forward_only(make_snapshot): @@ -2644,15 +2883,30 @@ def test_deployability_index_missing_parent(make_snapshot): "sqlmesh__foo.foo__bar__1234", ), ( - dict(physical_schema="sqlmesh__foo", name="bar", version="1234", catalog="foo"), + dict( + physical_schema="sqlmesh__foo", + name="bar", + version="1234", + catalog="foo", + ), "foo.sqlmesh__foo.bar__1234", ), ( - dict(physical_schema="sqlmesh__foo", name="bar.baz", version="1234", catalog="foo"), + dict( + physical_schema="sqlmesh__foo", + name="bar.baz", + version="1234", + catalog="foo", + ), "foo.sqlmesh__foo.bar__baz__1234", ), ( - dict(physical_schema="sqlmesh__foo", name="bar.baz", version="1234", suffix="dev"), + dict( + physical_schema="sqlmesh__foo", + name="bar.baz", + version="1234", + suffix="dev", + ), "sqlmesh__foo.bar__baz__1234__dev", ), ( @@ -2820,14 +3074,18 @@ def test_table_name(call_kwargs: t.Dict[str, t.Any], expected: str): ), ( "test_db.test_model", - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.SCHEMA), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.SCHEMA + ), None, "duckdb", "test_db__dev.test_model", ), ( "test_db.test_model", - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.TABLE), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.TABLE + ), None, "duckdb", "test_db.test_model__dev", @@ -2845,7 +3103,9 @@ def test_table_name(call_kwargs: t.Dict[str, t.Any], expected: str): ), ( "original_catalog.test_db.test_model", - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.TABLE), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.TABLE + ), "default_catalog", "duckdb", "original_catalog.test_db.test_model__dev", @@ -2874,7 +3134,9 @@ def test_table_name(call_kwargs: t.Dict[str, t.Any], expected: str): ), ( "test_db.test_model", - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.TABLE), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.TABLE + ), "default_catalog", "duckdb", "test_db.test_model__dev", @@ -2893,14 +3155,18 @@ def test_table_name(call_kwargs: t.Dict[str, t.Any], expected: str): # EnvironmentSuffixTarget.CATALOG ( "test_db.test_model", - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.CATALOG), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.CATALOG + ), "default_catalog", "duckdb", "default_catalog__dev.test_db.test_model", ), ( "test_db.test_model", - EnvironmentNamingInfo(name="dev", suffix_target=EnvironmentSuffixTarget.CATALOG), + EnvironmentNamingInfo( + name="dev", suffix_target=EnvironmentSuffixTarget.CATALOG + ), "default_catalog", "snowflake", "DEFAULT_CATALOG__DEV.test_db.test_model", @@ -2908,7 +3174,12 @@ def test_table_name(call_kwargs: t.Dict[str, t.Any], expected: str): ), ) def test_display_name( - make_snapshot, model_name, environment_naming_info, default_catalog, dialect, expected + make_snapshot, + model_name, + environment_naming_info, + default_catalog, + dialect, + expected, ): input_model = SqlModel( name=model_name, @@ -2918,7 +3189,9 @@ def test_display_name( ) input_snapshot = make_snapshot(input_model) assert ( - display_name(input_snapshot, environment_naming_info, default_catalog, dialect=dialect) + display_name( + input_snapshot, environment_naming_info, default_catalog, dialect=dialect + ) == expected ) @@ -2938,19 +3211,31 @@ def test_missing_intervals_node_start_end(make_snapshot): assert missing_intervals([snapshot], start="2024-03-12")[snapshot] == [ (to_timestamp("2024-03-12"), to_timestamp("2024-03-13")) ] - assert missing_intervals([snapshot], start="2024-03-12", end="2024-03-12")[snapshot] == [ - (to_timestamp("2024-03-12"), to_timestamp("2024-03-13")) - ] - assert missing_intervals([snapshot], start="2024-03-12", end="2024-03-14")[snapshot] == [ - (to_timestamp("2024-03-12"), to_timestamp("2024-03-13")) - ] - assert missing_intervals([snapshot], start="2024-03-01", end="2024-03-30")[snapshot] == [ - (to_timestamp("2024-03-12"), to_timestamp("2024-03-13")) - ] - assert missing_intervals([snapshot], start="2024-03-11", end=to_datetime("2024-03-12")) == {} - assert missing_intervals([snapshot], start="2024-03-12", end=to_datetime("2024-03-12")) == {} - assert missing_intervals([snapshot], start="2024-03-01", end=to_datetime("2024-03-12")) == {} - assert missing_intervals([snapshot], start="2024-03-01", end=to_datetime("2024-03-10")) == {} + assert missing_intervals([snapshot], start="2024-03-12", end="2024-03-12")[ + snapshot + ] == [(to_timestamp("2024-03-12"), to_timestamp("2024-03-13"))] + assert missing_intervals([snapshot], start="2024-03-12", end="2024-03-14")[ + snapshot + ] == [(to_timestamp("2024-03-12"), to_timestamp("2024-03-13"))] + assert missing_intervals([snapshot], start="2024-03-01", end="2024-03-30")[ + snapshot + ] == [(to_timestamp("2024-03-12"), to_timestamp("2024-03-13"))] + assert ( + missing_intervals([snapshot], start="2024-03-11", end=to_datetime("2024-03-12")) + == {} + ) + assert ( + missing_intervals([snapshot], start="2024-03-12", end=to_datetime("2024-03-12")) + == {} + ) + assert ( + missing_intervals([snapshot], start="2024-03-01", end=to_datetime("2024-03-12")) + == {} + ) + assert ( + missing_intervals([snapshot], start="2024-03-01", end=to_datetime("2024-03-10")) + == {} + ) assert missing_intervals([snapshot], start="2024-03-13", end="2024-03-14") == {} assert missing_intervals([snapshot], start="2024-03-14", end="2024-03-30") == {} @@ -3040,7 +3325,9 @@ def _loader(snapshot_ids: t.Set[SnapshotId]) -> t.Collection[Snapshot]: ) assert loader_called_times == 1 - cached_snapshot = cache.get_or_load([snapshot.snapshot_id], _loader)[0][snapshot.snapshot_id] + cached_snapshot = cache.get_or_load([snapshot.snapshot_id], _loader)[0][ + snapshot.snapshot_id + ] assert cached_snapshot.model._query_renderer._optimized_cache is not None assert cached_snapshot.model._data_hash is not None assert cached_snapshot.model._metadata_hash is not None @@ -3136,7 +3423,9 @@ def test_physical_version_pin(make_snapshot): SqlModel( name="a", kind=dict( - time_column="ds", name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True + time_column="ds", + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, ), query=parse_one("SELECT 1, ds"), physical_version="1234", @@ -3152,7 +3441,9 @@ def test_physical_version_pin_for_new_forward_only_models(make_snapshot): SqlModel( name="a", kind=dict( - time_column="ds", name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True + time_column="ds", + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, ), query=parse_one("SELECT 1, ds"), ), @@ -3164,7 +3455,9 @@ def test_physical_version_pin_for_new_forward_only_models(make_snapshot): SqlModel( name="a", kind=dict( - time_column="ds", name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True + time_column="ds", + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, ), query=parse_one("SELECT 2, ds"), ), @@ -3179,7 +3472,9 @@ def test_physical_version_pin_for_new_forward_only_models(make_snapshot): SqlModel( name="a", kind=dict( - time_column="ds", name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True + time_column="ds", + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, ), query=parse_one("SELECT 3, ds"), ), @@ -3209,7 +3504,9 @@ def test_physical_version_pin_for_new_forward_only_models(make_snapshot): SqlModel( name="a", kind=dict( - time_column="ds", name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True + time_column="ds", + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, ), query=parse_one("SELECT 5, ds"), ), @@ -3225,7 +3522,9 @@ def test_physical_version_pin_for_new_forward_only_models(make_snapshot): SqlModel( name="a", kind=dict( - time_column="ds", name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, forward_only=True + time_column="ds", + name=ModelKindName.INCREMENTAL_BY_TIME_RANGE, + forward_only=True, ), query=parse_one("SELECT 5, ds"), physical_version="1234", @@ -3252,7 +3551,9 @@ def test_contiguous_intervals(): def test_check_ready_intervals(mocker: MockerFixture): def assert_always_signal(intervals): assert ( - check_ready_intervals(lambda _: True, intervals, mocker.Mock(), mocker.Mock()) + check_ready_intervals( + lambda _: True, intervals, mocker.Mock(), mocker.Mock() + ) == intervals ) @@ -3262,7 +3563,12 @@ def assert_always_signal(intervals): assert_always_signal([(0, 1), (2, 3)]) def assert_never_signal(intervals): - assert check_ready_intervals(lambda _: False, intervals, mocker.Mock(), mocker.Mock()) == [] + assert ( + check_ready_intervals( + lambda _: False, intervals, mocker.Mock(), mocker.Mock() + ) + == [] + ) assert_never_signal([]) assert_never_signal([(0, 1)]) @@ -3270,7 +3576,10 @@ def assert_never_signal(intervals): assert_never_signal([(0, 1), (2, 3)]) def assert_empty_signal(intervals): - assert check_ready_intervals(lambda _: [], intervals, mocker.Mock(), mocker.Mock()) == [] + assert ( + check_ready_intervals(lambda _: [], intervals, mocker.Mock(), mocker.Mock()) + == [] + ) assert_empty_signal([]) assert_empty_signal([(0, 1)]) @@ -3363,9 +3672,14 @@ def test_get_next_auto_restatement_interval( snapshot.add_interval("2020-01-01", "2020-01-05") snapshot.next_auto_restatement_ts = to_timestamp("2020-01-06 10:00:00") - assert snapshot.get_next_auto_restatement_interval(to_timestamp("2020-01-06 09:59:00")) is None + assert ( + snapshot.get_next_auto_restatement_interval(to_timestamp("2020-01-06 09:59:00")) + is None + ) - assert snapshot.get_next_auto_restatement_interval(to_timestamp("2020-01-06 10:01:00")) == ( + assert snapshot.get_next_auto_restatement_interval( + to_timestamp("2020-01-06 10:01:00") + ) == ( to_timestamp(expected_auto_restatement_start), to_timestamp("2020-01-06"), ) @@ -3500,7 +3814,10 @@ def test_apply_auto_restatements(make_snapshot): intervals=[], dev_intervals=[], pending_restatement_intervals=[ - (to_timestamp("2020-01-05 10:00:00"), to_timestamp("2020-01-06 10:00:00")) + ( + to_timestamp("2020-01-05 10:00:00"), + to_timestamp("2020-01-06 10:00:00"), + ) ], ), SnapshotIntervals( @@ -3533,7 +3850,10 @@ def test_apply_auto_restatements(make_snapshot): intervals=[], dev_intervals=[], pending_restatement_intervals=[ - (to_timestamp("2020-01-05 10:00:00"), to_timestamp("2020-01-06 10:00:00")) + ( + to_timestamp("2020-01-05 10:00:00"), + to_timestamp("2020-01-06 10:00:00"), + ) ], ), SnapshotIntervals( @@ -3544,7 +3864,10 @@ def test_apply_auto_restatements(make_snapshot): intervals=[], dev_intervals=[], pending_restatement_intervals=[ - (to_timestamp("2020-01-06 05:00:00"), to_timestamp("2020-01-06 10:00:00")) + ( + to_timestamp("2020-01-06 05:00:00"), + to_timestamp("2020-01-06 10:00:00"), + ) ], ), SnapshotIntervals( @@ -3633,7 +3956,10 @@ def test_apply_auto_restatements_disable_restatement_downstream(make_snapshot): intervals=[], dev_intervals=[], pending_restatement_intervals=[ - (to_timestamp("2020-01-05 10:00:00"), to_timestamp("2020-01-06 10:00:00")) + ( + to_timestamp("2020-01-05 10:00:00"), + to_timestamp("2020-01-06 10:00:00"), + ) ], ), ] @@ -3769,7 +4095,9 @@ def test_auto_restatement_triggers(make_snapshot): def test_render_signal(make_snapshot, mocker): @signal() - def check_types(batch, env: str, sql: list[SQL], table: exp.Table, default: int = 0): + def check_types( + batch, env: str, sql: list[SQL], table: exp.Table, default: int = 0 + ): if not ( env == "in_memory" and default == 0 @@ -3781,15 +4109,13 @@ def check_types(batch, env: str, sql: list[SQL], table: exp.Table, default: int return True sql_model = load_sql_based_model( - parse( - """ + parse(""" MODEL ( name test_schema.test_model, signals check_types(env := @gateway, sql := [a.b], table := b.c) ); SELECT a FROM tbl; - """ - ), + """), variables={ c.GATEWAY: "in_memory", }, @@ -3800,16 +4126,14 @@ def check_types(batch, env: str, sql: list[SQL], table: exp.Table, default: int def test_partitioned_by_roundtrip(make_snapshot: t.Callable): - sql_model = load_sql_based_model( - parse(""" + sql_model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind full, partitioned_by (a, bucket(4, b), truncate(3, c), month(d)) ); SELECT a, b, c, d FROM tbl; - """) - ) + """)) snapshot = make_snapshot(sql_model) assert isinstance(snapshot, Snapshot) assert isinstance(snapshot.node, SqlModel) @@ -3897,7 +4221,11 @@ def test_snapshot_id_and_version_fingerprint_lazy_init(): # can also be supplied as a SnapshotFingerprint to begin with instead of a str snapshot = SnapshotIdAndVersion( - name="a", identifier="1234", version="2345", dev_version=None, fingerprint=fingerprint + name="a", + identifier="1234", + version="2345", + dev_version=None, + fingerprint=fingerprint, ) assert isinstance(snapshot.fingerprint_, SnapshotFingerprint) diff --git a/tests/core/test_snapshot_evaluator.py b/tests/core/test_snapshot_evaluator.py index 27bcbe05ae..c793d0e457 100644 --- a/tests/core/test_snapshot_evaluator.py +++ b/tests/core/test_snapshot_evaluator.py @@ -1,88 +1,66 @@ from __future__ import annotations -import typing as t -from typing_extensions import Self -from unittest.mock import call, patch, Mock import contextlib -import re +import json import logging -import pytest +import re +import typing as t +from pathlib import Path +from unittest.mock import Mock, call, patch + import pandas as pd # noqa: TID253 -import json +import pytest from pydantic import model_validator -from pathlib import Path from pytest_mock.plugin import MockerFixture from sqlglot import expressions as exp from sqlglot import parse, parse_one, select +from typing_extensions import Self -from sqlmesh.core.audit import ModelAudit, StandaloneAudit from sqlmesh.core import dialect as d +from sqlmesh.core.audit import ModelAudit, StandaloneAudit from sqlmesh.core.dialect import schema_, to_schema -from sqlmesh.core.engine_adapter import EngineAdapter, create_engine_adapter, BigQueryEngineAdapter -from sqlmesh.core.engine_adapter.base import MERGE_SOURCE_ALIAS, MERGE_TARGET_ALIAS -from sqlmesh.core.engine_adapter.shared import ( - DataObject, - DataObjectType, - InsertOverwriteStrategy, -) +from sqlmesh.core.engine_adapter import (BigQueryEngineAdapter, EngineAdapter, + create_engine_adapter) +from sqlmesh.core.engine_adapter.base import (MERGE_SOURCE_ALIAS, + MERGE_TARGET_ALIAS) +from sqlmesh.core.engine_adapter.shared import (DataObject, DataObjectType, + InsertOverwriteStrategy) from sqlmesh.core.environment import EnvironmentNamingInfo -from sqlmesh.core.macros import RuntimeStage, macro, MacroEvaluator, MacroFunc -from sqlmesh.core.model import ( - Model, - FullKind, - IncrementalByTimeRangeKind, - IncrementalUnmanagedKind, - IncrementalByPartitionKind, - IncrementalByUniqueKeyKind, - PythonModel, - SqlModel, - TimeColumn, - ViewKind, - CustomKind, - load_sql_based_model, - ExternalModel, - model, - create_sql_model, -) -from sqlmesh.core.model.kind import OnDestructiveChange, ExternalKind, OnAdditiveChange +from sqlmesh.core.macros import MacroEvaluator, MacroFunc, RuntimeStage, macro +from sqlmesh.core.model import (CustomKind, ExternalModel, FullKind, + IncrementalByPartitionKind, + IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, + IncrementalUnmanagedKind, Model, PythonModel, + SqlModel, TimeColumn, ViewKind, + create_sql_model, load_sql_based_model, model) +from sqlmesh.core.model.kind import (ExternalKind, OnAdditiveChange, + OnDestructiveChange) from sqlmesh.core.model.meta import GrantsTargetLayer from sqlmesh.core.node import IntervalUnit -from sqlmesh.core.snapshot import ( - DeployabilityIndex, - Intervals, - Snapshot, - SnapshotDataVersion, - SnapshotFingerprint, - SnapshotChangeCategory, - SnapshotEvaluator, - SnapshotTableCleanupTask, -) +from sqlmesh.core.snapshot import (DeployabilityIndex, Intervals, Snapshot, + SnapshotChangeCategory, SnapshotDataVersion, + SnapshotEvaluator, SnapshotFingerprint, + SnapshotTableCleanupTask) from sqlmesh.core.snapshot.definition import to_view_mapping -from sqlmesh.core.snapshot.evaluator import ( - CustomMaterialization, - EngineManagedStrategy, - FullRefreshStrategy, - IncrementalByPartitionStrategy, - IncrementalByTimeRangeStrategy, - IncrementalByUniqueKeyStrategy, - IncrementalUnmanagedStrategy, - MaterializableStrategy, - SCDType2Strategy, - SnapshotCreationFailedError, - ViewStrategy, -) +from sqlmesh.core.snapshot.evaluator import (CustomMaterialization, + EngineManagedStrategy, + FullRefreshStrategy, + IncrementalByPartitionStrategy, + IncrementalByTimeRangeStrategy, + IncrementalByUniqueKeyStrategy, + IncrementalUnmanagedStrategy, + MaterializableStrategy, + SCDType2Strategy, + SnapshotCreationFailedError, + ViewStrategy) from sqlmesh.utils.concurrency import NodeExecutionFailedError from sqlmesh.utils.date import to_timestamp -from sqlmesh.utils.errors import ( - ConfigError, - SQLMeshError, - DestructiveChangeError, - AdditiveChangeError, -) +from sqlmesh.utils.errors import (AdditiveChangeError, ConfigError, + DestructiveChangeError, SQLMeshError) from sqlmesh.utils.metaprogramming import Executable from sqlmesh.utils.pydantic import list_of_fields_validator - if t.TYPE_CHECKING: from sqlmesh.core.engine_adapter._typing import QueryOrDF @@ -183,8 +161,7 @@ def x(evaluator, y=None) -> None: evaluator.locals["payload"]["y"] = y model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE (time_column a), @@ -204,8 +181,7 @@ def x(evaluator, y=None) -> None: @DEF(b, 2); @x(['a', 2, TRUE]); - """ - ), + """), # type: ignore macros=macro.get_registry(), ) @@ -256,8 +232,7 @@ def increment_stage_counter(evaluator) -> None: print(f"RuntimeStage value: {evaluator.locals['runtime_stage']}") model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind FULL, @@ -267,8 +242,7 @@ def increment_stage_counter(evaluator) -> None: @if(@runtime_stage = 'evaluating', ALTER TABLE test_schema.foo MODIFY COLUMN c SET MASKING POLICY p); SELECT 1 AS a, @runtime_stage AS b; - """ - ), + """), # type: ignore macros=macro.get_registry(), ) @@ -279,15 +253,26 @@ def increment_stage_counter(evaluator) -> None: snapshot.model.render_pre_statements() - assert f"RuntimeStage value: {RuntimeStage.LOADING.value}" in capsys.readouterr().out + assert ( + f"RuntimeStage value: {RuntimeStage.LOADING.value}" in capsys.readouterr().out + ) evaluator.create([snapshot], {}) - assert f"RuntimeStage value: {RuntimeStage.CREATING.value}" in capsys.readouterr().out + assert ( + f"RuntimeStage value: {RuntimeStage.CREATING.value}" in capsys.readouterr().out + ) evaluator.evaluate( - snapshot, start="2020-01-01", end="2020-01-02", execution_time="2020-01-02", snapshots={} + snapshot, + start="2020-01-01", + end="2020-01-02", + execution_time="2020-01-02", + snapshots={}, + ) + assert ( + f"RuntimeStage value: {RuntimeStage.EVALUATING.value}" + in capsys.readouterr().out ) - assert f"RuntimeStage value: {RuntimeStage.EVALUATING.value}" in capsys.readouterr().out empty_call = call([]) non_empty_calls = [c for c in adapter_mock.execute.mock_calls if c != empty_call] @@ -296,7 +281,9 @@ def increment_stage_counter(evaluator) -> None: [parse_one("ALTER TABLE test_schema.foo MODIFY COLUMN c SET MASKING POLICY p")] ) - assert snapshot.model.render_query().sql() == '''SELECT 1 AS "a", 'loading' AS "b"''' + assert ( + snapshot.model.render_query().sql() == '''SELECT 1 AS "a", 'loading' AS "b"''' + ) assert ( snapshot.model.render_query(runtime_stage=RuntimeStage.CREATING).sql() == '''SELECT 1 AS "a", 'creating' AS "b"''' @@ -324,7 +311,9 @@ def test_promote(mocker: MockerFixture, adapter_mock, make_snapshot): adapter_mock.transaction.assert_called() adapter_mock.session.assert_called() - adapter_mock.create_schema.assert_called_once_with(to_schema("test_schema__test_env")) + adapter_mock.create_schema.assert_called_once_with( + to_schema("test_schema__test_env") + ) adapter_mock.create_view.assert_called_once_with( "test_schema__test_env.test_model", parse_one( @@ -463,7 +452,11 @@ def create_and_cleanup(name: str, dev_table_only: bool): evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) evaluator.cleanup( - [SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=dev_table_only)], + [ + SnapshotTableCleanupTask( + snapshot=snapshot.table_info, dev_table_only=dev_table_only + ) + ], on_complete=on_cleanup_mock, ) assert on_cleanup_mock.call_count == 1 if dev_table_only else 2 @@ -485,7 +478,10 @@ def create_and_cleanup(name: str, dev_table_only: bool): f"sqlmesh__test_schema.test_schema__test_model__{snapshot.fingerprint.to_version()}__dev", cascade=True, ), - call(f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}", cascade=True), + call( + f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}", + cascade=True, + ), ] ) adapter_mock.reset_mock() @@ -516,7 +512,9 @@ def test_cleanup_view(adapter_mock, make_snapshot): snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) - evaluator.cleanup([SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True)]) + evaluator.cleanup( + [SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True)] + ) adapter_mock.get_data_object.assert_not_called() adapter_mock.drop_view.assert_called_once_with( @@ -541,7 +539,9 @@ def test_cleanup_materialized_view(adapter_mock, make_snapshot): adapter_mock.drop_view.side_effect = [RuntimeError("failed to drop view"), None] evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) - evaluator.cleanup([SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True)]) + evaluator.cleanup( + [SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True)] + ) adapter_mock.get_data_object.assert_not_called() adapter_mock.drop_view.assert_has_calls( @@ -580,7 +580,11 @@ def test_cleanup_fails(adapter_mock, make_snapshot): evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) with pytest.raises(SQLMeshError) as exc_info: evaluator.cleanup( - [SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True)] + [ + SnapshotTableCleanupTask( + snapshot=snapshot.table_info, dev_table_only=True + ) + ] ) assert "test_error" in str(exc_info.value) @@ -604,7 +608,9 @@ def test_cleanup_skip_missing_table(adapter_mock, make_snapshot): snapshot.version = "test_version" evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) - evaluator.cleanup([SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True)]) + evaluator.cleanup( + [SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True)] + ) adapter_mock.get_data_object.assert_called_once_with( f"catalog.sqlmesh__test_schema.test_schema__test_model__{snapshot.fingerprint.to_version()}__dev" @@ -630,7 +636,11 @@ def create_and_cleanup_external_model(name: str, dev_table_only: bool): evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) evaluator.cleanup( - [SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=dev_table_only)] + [ + SnapshotTableCleanupTask( + snapshot=snapshot.table_info, dev_table_only=dev_table_only + ) + ] ) return snapshot @@ -659,8 +669,12 @@ def test_cleanup_symbolic_and_audit_snapshots_no_callback( evaluator.cleanup( [ - SnapshotTableCleanupTask(snapshot=external_snapshot.table_info, dev_table_only=False), - SnapshotTableCleanupTask(snapshot=audit_snapshot.table_info, dev_table_only=False), + SnapshotTableCleanupTask( + snapshot=external_snapshot.table_info, dev_table_only=False + ), + SnapshotTableCleanupTask( + snapshot=audit_snapshot.table_info, dev_table_only=False + ), ], on_complete=on_complete_mock, ) @@ -679,8 +693,7 @@ def test_evaluate_materialized_view( evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ( @@ -689,8 +702,7 @@ def test_evaluate_materialized_view( ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -734,8 +746,7 @@ def test_evaluate_materialized_view_not_recreated_on_evaluation( evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ( @@ -744,8 +755,7 @@ def test_evaluate_materialized_view_not_recreated_on_evaluation( ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -811,8 +821,7 @@ def test_evaluate_materialized_view_with_execution_time_macro( evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ( @@ -821,8 +830,7 @@ def test_evaluate_materialized_view_with_execution_time_macro( ); SELECT a::int FROM tbl WHERE ds < @execution_ds; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -909,7 +917,10 @@ def test_evaluate_incremental_unmanaged_no_intervals( snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - table_columns = {"one": exp.DataType.build("int"), "ds": exp.DataType.build("timestamp")} + table_columns = { + "one": exp.DataType.build("int"), + "ds": exp.DataType.build("timestamp"), + } adapter_mock.columns.return_value = table_columns evaluator = SnapshotEvaluator(adapter_mock) @@ -940,16 +951,14 @@ def test_evaluate_incremental_unmanaged_no_intervals( def test_create_prod_table_exists(mocker: MockerFixture, adapter_mock, make_snapshot): model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -982,10 +991,11 @@ def test_pre_hook_forward_only_clone( """ Verifies that pre-statements are executed when creating a clone of a forward-only model. """ - pre_statement = """CREATE TEMPORARY FUNCTION "example_udf"("x" BIGINT) AS ("x" + 1)""" + pre_statement = ( + """CREATE TEMPORARY FUNCTION "example_udf"("x" BIGINT) AS ("x" + 1)""" + ) model = load_sql_based_model( - parse( # type: ignore - f""" + parse(f""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -996,8 +1006,7 @@ def test_pre_hook_forward_only_clone( {pre_statement}; SELECT a::int, ds::string FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1021,22 +1030,24 @@ def test_pre_hook_forward_only_clone( evaluator = SnapshotEvaluator(adapter) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) adapter.cursor.execute.assert_any_call(pre_statement) -def test_create_only_dev_table_exists(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_create_only_dev_table_exists( + mocker: MockerFixture, adapter_mock, make_snapshot +): model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1052,7 +1063,9 @@ def test_create_only_dev_table_exists(mocker: MockerFixture, adapter_mock, make_ adapter_mock.table_exists.return_value = True evaluator = SnapshotEvaluator(adapter_mock) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) adapter_mock.create_view.assert_not_called() adapter_mock.get_data_objects.assert_called_once_with( schema_("sqlmesh__test_schema"), @@ -1063,10 +1076,11 @@ def test_create_only_dev_table_exists(mocker: MockerFixture, adapter_mock, make_ ) -def test_create_new_forward_only_model(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_create_new_forward_only_model( + mocker: MockerFixture, adapter_mock, make_snapshot +): model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1076,8 +1090,7 @@ def test_create_new_forward_only_model(mocker: MockerFixture, adapter_mock, make ); SELECT a::int, '2024-01-01' as ds FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1087,7 +1100,9 @@ def test_create_new_forward_only_model(mocker: MockerFixture, adapter_mock, make adapter_mock.table_exists.return_value = False evaluator = SnapshotEvaluator(adapter_mock) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) # Only non-deployable table should be created adapter_mock.create_table.assert_called_once_with( f"sqlmesh__test_schema.test_schema__test_model__{snapshot.dev_version}__dev", @@ -1117,7 +1132,11 @@ def test_create_new_forward_only_model(mocker: MockerFixture, adapter_mock, make "deployability_index, snapshot_category, forward_only", [ (DeployabilityIndex.all_deployable(), SnapshotChangeCategory.BREAKING, False), - (DeployabilityIndex.all_deployable(), SnapshotChangeCategory.NON_BREAKING, False), + ( + DeployabilityIndex.all_deployable(), + SnapshotChangeCategory.NON_BREAKING, + False, + ), (DeployabilityIndex.all_deployable(), SnapshotChangeCategory.BREAKING, True), ( DeployabilityIndex.all_deployable(), @@ -1207,18 +1226,18 @@ def test_create_tables_exist( adapter_mock.create_table.assert_not_called() -def test_create_prod_table_exists_forward_only(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_create_prod_table_exists_forward_only( + mocker: MockerFixture, adapter_mock, make_snapshot +): model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind FULL ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1245,18 +1264,18 @@ def test_create_prod_table_exists_forward_only(mocker: MockerFixture, adapter_mo adapter_mock.create_table.assert_not_called() -def test_create_view_non_deployable_snapshot(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_create_view_non_deployable_snapshot( + mocker: MockerFixture, adapter_mock, make_snapshot +): model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1288,8 +1307,7 @@ def test_create_materialized_view(mocker: MockerFixture, adapter_mock, make_snap evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ( @@ -1298,8 +1316,7 @@ def test_create_materialized_view(mocker: MockerFixture, adapter_mock, make_snap ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1321,19 +1338,23 @@ def test_create_materialized_view(mocker: MockerFixture, adapter_mock, make_snap ) adapter_mock.create_view.assert_called_once_with( - snapshot.table_name(), model.render_query(), column_descriptions={}, **common_kwargs + snapshot.table_name(), + model.render_query(), + column_descriptions={}, + **common_kwargs, ) -def test_create_view_with_properties(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_create_view_with_properties( + mocker: MockerFixture, adapter_mock, make_snapshot +): adapter_mock.get_data_objects.return_value = [] adapter_mock.table_exists.return_value = False evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ( @@ -1345,8 +1366,7 @@ def test_create_view_with_properties(mocker: MockerFixture, adapter_mock, make_s ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1370,7 +1390,10 @@ def test_create_view_with_properties(mocker: MockerFixture, adapter_mock, make_s ) adapter_mock.create_view.assert_called_once_with( - snapshot.table_name(), model.render_query(), column_descriptions={}, **common_kwargs + snapshot.table_name(), + model.render_query(), + column_descriptions={}, + **common_kwargs, ) @@ -1387,8 +1410,7 @@ def test_create_materialized_view_with_audits_sets_has_audits( evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind VIEW ( @@ -1400,8 +1422,7 @@ def test_create_materialized_view_with_audits_sets_has_audits( ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1433,10 +1454,14 @@ def test_promote_model_info(mocker: MockerFixture, make_snapshot): evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) - adapter_mock.create_schema.assert_called_once_with(to_schema("test_schema__test_env")) + adapter_mock.create_schema.assert_called_once_with( + to_schema("test_schema__test_env") + ) adapter_mock.create_view.assert_called_once_with( "test_schema__test_env.test_model", - parse_one(f"SELECT * FROM physical_schema.test_schema__test_model__{snapshot.version}"), + parse_one( + f"SELECT * FROM physical_schema.test_schema__test_model__{snapshot.version}" + ), table_description=None, column_descriptions=None, view_properties={}, @@ -1481,7 +1506,9 @@ def test_promote_deployable(mocker: MockerFixture, make_snapshot): adapter_mock.get_data_objects.return_value = [] evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) - adapter_mock.create_schema.assert_called_once_with(to_schema("test_schema__test_env")) + adapter_mock.create_schema.assert_called_once_with( + to_schema("test_schema__test_env") + ) adapter_mock.create_view.assert_called_once_with( "test_schema__test_env.test_model", parse_one( @@ -1556,7 +1583,9 @@ def columns(table_name): session_spy.assert_called_once() -def test_migrate_missing_table(mocker: MockerFixture, make_snapshot, make_mocked_engine_adapter): +def test_migrate_missing_table( + mocker: MockerFixture, make_snapshot, make_mocked_engine_adapter +): adapter = make_mocked_engine_adapter(EngineAdapter) adapter.table_exists = lambda _: False # type: ignore adapter.with_settings = lambda **kwargs: adapter # type: ignore @@ -1579,7 +1608,9 @@ def test_migrate_missing_table(mocker: MockerFixture, make_snapshot, make_mocked snapshot.forward_only = True snapshot.previous_versions = snapshot.all_versions - evaluator.migrate([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.migrate( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) adapter.cursor.execute.assert_not_called() @@ -1702,7 +1733,9 @@ def assert_tables_exist() -> None: snapshots={}, ) assert_tables_exist() - assert duck_conn.execute(f"SELECT * FROM sqlmesh__db.db__model__{version}").fetchall() == [(1,)] + assert duck_conn.execute( + f"SELECT * FROM sqlmesh__db.db__model__{version}" + ).fetchall() == [(1,)] # test that existing tables work evaluator.evaluate( @@ -1713,7 +1746,9 @@ def assert_tables_exist() -> None: snapshots={}, ) assert_tables_exist() - assert duck_conn.execute(f"SELECT * FROM sqlmesh__db.db__model__{version}").fetchall() == [(1,)] + assert duck_conn.execute( + f"SELECT * FROM sqlmesh__db.db__model__{version}" + ).fetchall() == [(1,)] def test_migrate_duckdb(snapshot: Snapshot, duck_conn, make_snapshot): @@ -1812,7 +1847,9 @@ def test_audit_unversioned(mocker: MockerFixture, adapter_mock, make_snapshot): ), ], ) -def test_snapshot_evaluator_yield_pd(adapter_mock, make_snapshot, input_dfs, output_dict): +def test_snapshot_evaluator_yield_pd( + adapter_mock, make_snapshot, input_dfs, output_dict +): adapter_mock.is_pyspark_df.return_value = False adapter_mock.INSERT_OVERWRITE_STRATEGY = InsertOverwriteStrategy.INSERT_OVERWRITE adapter_mock.try_get_df = lambda x: x @@ -1822,7 +1859,9 @@ def test_snapshot_evaluator_yield_pd(adapter_mock, make_snapshot, input_dfs, out PythonModel( name="db.model", entrypoint="python_func", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds", format="%Y-%m-%d")), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds", format="%Y-%m-%d") + ), columns={ "a": "INT", "ds": "STRING", @@ -1854,7 +1893,10 @@ def python_func(**kwargs): snapshots={}, ) - assert adapter_mock.insert_overwrite_by_time_partition.call_args[0][1].to_dict() == output_dict + assert ( + adapter_mock.insert_overwrite_by_time_partition.call_args[0][1].to_dict() + == output_dict + ) def test_snapshot_evaluator_yield_empty_pd(adapter_mock, make_snapshot): @@ -1867,7 +1909,9 @@ def test_snapshot_evaluator_yield_empty_pd(adapter_mock, make_snapshot): PythonModel( name="db.model", entrypoint="python_func", - kind=IncrementalByTimeRangeKind(time_column=TimeColumn(column="ds", format="%Y-%m-%d")), + kind=IncrementalByTimeRangeKind( + time_column=TimeColumn(column="ds", format="%Y-%m-%d") + ), columns={ "a": "INT", "ds": "STRING", @@ -1906,8 +1950,7 @@ def test_create_clone_in_dev(mocker: MockerFixture, adapter_mock, make_snapshot) evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1916,19 +1959,23 @@ def test_create_clone_in_dev(mocker: MockerFixture, adapter_mock, make_snapshot) ); SELECT 1::INT as a, ds::DATE FROM a; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) snapshot.previous_versions = snapshot.all_versions - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) adapter_mock.create_table.assert_called_once_with( f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}__dev_schema_tmp", - target_columns_to_types={"a": exp.DataType.build("int"), "ds": exp.DataType.build("date")}, + target_columns_to_types={ + "a": exp.DataType.build("int"), + "ds": exp.DataType.build("date"), + }, table_format=None, storage_format=None, partitioned_by=[exp.to_column("ds", quoted=True)], @@ -1959,7 +2006,9 @@ def test_create_clone_in_dev(mocker: MockerFixture, adapter_mock, make_snapshot) ) -def test_drop_clone_in_dev_when_migration_fails(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_drop_clone_in_dev_when_migration_fails( + mocker: MockerFixture, adapter_mock, make_snapshot +): adapter_mock.SUPPORTS_CLONING = True adapter_mock.get_alter_operations.return_value = [] evaluator = SnapshotEvaluator(adapter_mock) @@ -1967,8 +2016,7 @@ def test_drop_clone_in_dev_when_migration_fails(mocker: MockerFixture, adapter_m adapter_mock.alter_table.side_effect = DestructiveChangeError("Migration failed") model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -1977,8 +2025,7 @@ def test_drop_clone_in_dev_when_migration_fails(mocker: MockerFixture, adapter_m ); SELECT 1::INT as a, ds::DATE FROM a; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -1986,7 +2033,9 @@ def test_drop_clone_in_dev_when_migration_fails(mocker: MockerFixture, adapter_m snapshot.previous_versions = snapshot.all_versions with pytest.raises(SnapshotCreationFailedError): - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) adapter_mock.clone_table.assert_called_once_with( f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}__dev", @@ -2008,7 +2057,9 @@ def test_drop_clone_in_dev_when_migration_fails(mocker: MockerFixture, adapter_m call( f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}__dev_schema_tmp" ), - call(f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}__dev"), + call( + f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}__dev" + ), ] ) @@ -2023,8 +2074,7 @@ def test_create_clone_in_dev_self_referencing( from_table = "test_schema.test_model" if not use_this_model else "@this_model" model = load_sql_based_model( - parse( # type: ignore - f""" + parse(f""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -2033,19 +2083,23 @@ def test_create_clone_in_dev_self_referencing( ); SELECT 1::INT as a, ds::DATE FROM {from_table}; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) snapshot.previous_versions = snapshot.all_versions - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) adapter_mock.create_table.assert_called_once_with( f"sqlmesh__test_schema.test_schema__test_model__{snapshot.version}__dev_schema_tmp", - target_columns_to_types={"a": exp.DataType.build("int"), "ds": exp.DataType.build("date")}, + target_columns_to_types={ + "a": exp.DataType.build("int"), + "ds": exp.DataType.build("date"), + }, table_format=None, storage_format=None, partitioned_by=[exp.to_column("ds", quoted=True)], @@ -2164,8 +2218,12 @@ def test_on_additive_change_runtime_check( # SQLMesh default: ERROR model = SqlModel( name="test_schema.test_model", - kind=IncrementalByTimeRangeKind(time_column="a", on_additive_change=OnAdditiveChange.ERROR), - query=parse_one("SELECT c, a, b FROM tbl WHERE ds BETWEEN @start_ds and @end_ds"), + kind=IncrementalByTimeRangeKind( + time_column="a", on_additive_change=OnAdditiveChange.ERROR + ), + query=parse_one( + "SELECT c, a, b FROM tbl WHERE ds BETWEEN @start_ds and @end_ds" + ), ) snapshot = make_snapshot(model, version="1") snapshot.change_category = SnapshotChangeCategory.BREAKING @@ -2297,13 +2355,14 @@ def mock_temp_table(query_or_df, name="diff", **kwargs): assert str(temp_table_name_captured.db) == "sqlmesh__test_schema" -def test_forward_only_snapshot_for_added_model(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_forward_only_snapshot_for_added_model( + mocker: MockerFixture, adapter_mock, make_snapshot +): adapter_mock.SUPPORTS_CLONING = False evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE ( @@ -2312,17 +2371,21 @@ def test_forward_only_snapshot_for_added_model(mocker: MockerFixture, adapter_mo ); SELECT 1::INT as a, ds::DATE FROM a; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) common_create_args = dict( - target_columns_to_types={"a": exp.DataType.build("int"), "ds": exp.DataType.build("date")}, + target_columns_to_types={ + "a": exp.DataType.build("int"), + "ds": exp.DataType.build("date"), + }, table_format=None, storage_format=None, partitioned_by=[exp.to_column("ds", quoted=True)], @@ -2344,9 +2407,7 @@ def test_forward_only_snapshot_for_added_model(mocker: MockerFixture, adapter_mo def test_create_scd_type_2_by_time(adapter_mock, make_snapshot): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind SCD_TYPE_2 ( @@ -2356,14 +2417,14 @@ def test_create_scd_type_2_by_time(adapter_mock, make_snapshot): ); SELECT id::int, name::string, updated_at::timestamp FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) common_kwargs = dict( target_columns_to_types={ @@ -2396,9 +2457,7 @@ def test_create_scd_type_2_by_time(adapter_mock, make_snapshot): def test_create_ctas_scd_type_2_by_time(adapter_mock, make_snapshot, mocker): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind SCD_TYPE_2 ( @@ -2410,14 +2469,14 @@ def test_create_ctas_scd_type_2_by_time(adapter_mock, make_snapshot, mocker): ); SELECT * FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) source_query = parse_one('SELECT * FROM "tbl" AS "tbl"') query = parse_one( @@ -2493,9 +2552,7 @@ def test_insert_into_scd_type_2_by_time( adapter_mock, make_snapshot, intervals: Intervals, truncate: bool ): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind SCD_TYPE_2 ( @@ -2504,9 +2561,7 @@ def test_insert_into_scd_type_2_by_time( ); SELECT id::int, name::string, updated_at::timestamp FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -2557,9 +2612,7 @@ def test_insert_into_scd_type_2_by_time( def test_create_scd_type_2_by_column(adapter_mock, make_snapshot): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind SCD_TYPE_2_BY_COLUMN ( @@ -2570,14 +2623,14 @@ def test_create_scd_type_2_by_column(adapter_mock, make_snapshot): ); SELECT id::int, name::string, FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) common_kwargs = dict( target_columns_to_types={ @@ -2608,9 +2661,7 @@ def test_create_scd_type_2_by_column(adapter_mock, make_snapshot): def test_create_ctas_scd_type_2_by_column(adapter_mock, make_snapshot): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind SCD_TYPE_2_BY_COLUMN ( @@ -2620,14 +2671,14 @@ def test_create_ctas_scd_type_2_by_column(adapter_mock, make_snapshot): ); SELECT * FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) query = parse_one( """SELECT *, CAST(NULL AS TIMESTAMP) AS valid_from, CAST(NULL AS TIMESTAMP) AS valid_to FROM "tbl" AS "tbl" WHERE FALSE LIMIT 0""" @@ -2667,9 +2718,7 @@ def test_insert_into_scd_type_2_by_column( adapter_mock, make_snapshot, intervals: Intervals, truncate: bool ): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind SCD_TYPE_2_BY_COLUMN ( @@ -2679,9 +2728,7 @@ def test_insert_into_scd_type_2_by_column( ); SELECT id::int, name::string FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -2731,9 +2778,7 @@ def test_insert_into_scd_type_2_by_column( def test_create_incremental_by_unique_key_updated_at_exp(adapter_mock, make_snapshot): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -2743,9 +2788,7 @@ def test_create_incremental_by_unique_key_updated_at_exp(adapter_mock, make_snap ); SELECT id::int, name::string, updated_at::timestamp FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -2779,11 +2822,19 @@ def test_create_incremental_by_unique_key_updated_at_exp(adapter_mock, make_snap exp.column("name", MERGE_TARGET_ALIAS, quoted=True).eq( exp.column("name", MERGE_SOURCE_ALIAS, quoted=True) ), - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True).eq( + exp.column( + "updated_at", MERGE_TARGET_ALIAS, quoted=True + ).eq( exp.Coalesce( - this=exp.column("updated_at", MERGE_SOURCE_ALIAS, quoted=True), + this=exp.column( + "updated_at", MERGE_SOURCE_ALIAS, quoted=True + ), expressions=[ - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True) + exp.column( + "updated_at", + MERGE_TARGET_ALIAS, + quoted=True, + ) ], ) ), @@ -2797,11 +2848,11 @@ def test_create_incremental_by_unique_key_updated_at_exp(adapter_mock, make_snap ) -def test_create_incremental_by_unique_key_multiple_updated_at_exp(adapter_mock, make_snapshot): +def test_create_incremental_by_unique_key_multiple_updated_at_exp( + adapter_mock, make_snapshot +): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -2812,9 +2863,7 @@ def test_create_incremental_by_unique_key_multiple_updated_at_exp(adapter_mock, ); SELECT id::int, name::string, updated_at::timestamp FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -2850,11 +2899,19 @@ def test_create_incremental_by_unique_key_multiple_updated_at_exp(adapter_mock, exp.column("name", MERGE_TARGET_ALIAS, quoted=True).eq( exp.column("name", MERGE_SOURCE_ALIAS, quoted=True) ), - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True).eq( + exp.column( + "updated_at", MERGE_TARGET_ALIAS, quoted=True + ).eq( exp.Coalesce( - this=exp.column("updated_at", MERGE_SOURCE_ALIAS, quoted=True), + this=exp.column( + "updated_at", MERGE_SOURCE_ALIAS, quoted=True + ), expressions=[ - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True) + exp.column( + "updated_at", + MERGE_TARGET_ALIAS, + quoted=True, + ) ], ) ), @@ -2869,11 +2926,19 @@ def test_create_incremental_by_unique_key_multiple_updated_at_exp(adapter_mock, exp.column("name", MERGE_TARGET_ALIAS, quoted=True).eq( exp.column("name", MERGE_SOURCE_ALIAS, quoted=True) ), - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True).eq( + exp.column( + "updated_at", MERGE_TARGET_ALIAS, quoted=True + ).eq( exp.Coalesce( - this=exp.column("updated_at", MERGE_SOURCE_ALIAS, quoted=True), + this=exp.column( + "updated_at", MERGE_SOURCE_ALIAS, quoted=True + ), expressions=[ - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True) + exp.column( + "updated_at", + MERGE_TARGET_ALIAS, + quoted=True, + ) ], ) ), @@ -2889,9 +2954,7 @@ def test_create_incremental_by_unique_key_multiple_updated_at_exp(adapter_mock, def test_create_incremental_by_unique_no_intervals(adapter_mock, make_snapshot): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -2900,9 +2963,7 @@ def test_create_incremental_by_unique_no_intervals(adapter_mock, make_snapshot): ); SELECT id::int, name::string, updated_at::timestamp FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -2941,9 +3002,7 @@ def test_create_incremental_by_unique_no_intervals(adapter_mock, make_snapshot): def test_create_incremental_by_unique_key_merge_filter(adapter_mock, make_snapshot): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - d.parse( - """ + model = load_sql_based_model(d.parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_UNIQUE_KEY ( @@ -2954,9 +3013,7 @@ def test_create_incremental_by_unique_key_merge_filter(adapter_mock, make_snapsh ); SELECT id::int, updated_at::timestamp FROM tbl; - """ - ) - ) + """)) # At load time macros should remain unresolved assert model.merge_filter == exp.And( @@ -3003,11 +3060,19 @@ def test_create_incremental_by_unique_key_merge_filter(adapter_mock, make_snapsh matched=True, then=exp.Update( expressions=[ - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True).eq( + exp.column( + "updated_at", MERGE_TARGET_ALIAS, quoted=True + ).eq( exp.Coalesce( - this=exp.column("updated_at", MERGE_SOURCE_ALIAS, quoted=True), + this=exp.column( + "updated_at", MERGE_SOURCE_ALIAS, quoted=True + ), expressions=[ - exp.column("updated_at", MERGE_TARGET_ALIAS, quoted=True) + exp.column( + "updated_at", + MERGE_TARGET_ALIAS, + quoted=True, + ) ], ) ), @@ -3038,8 +3103,7 @@ def test_create_incremental_by_unique_key_merge_filter(adapter_mock, make_snapsh def test_create_seed(mocker: MockerFixture, adapter_mock, make_snapshot): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -3047,10 +3111,11 @@ def test_create_seed(mocker: MockerFixture, adapter_mock, make_snapshot): batch_size 5, ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -3114,8 +3179,7 @@ def test_create_seed(mocker: MockerFixture, adapter_mock, make_snapshot): def test_create_seed_on_error(mocker: MockerFixture, adapter_mock, make_snapshot): adapter_mock.insert_append.side_effect = Exception("test error") - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( @@ -3123,10 +3187,11 @@ def test_create_seed_on_error(mocker: MockerFixture, adapter_mock, make_snapshot batch_size 5, ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -3154,22 +3219,24 @@ def test_create_seed_on_error(mocker: MockerFixture, adapter_mock, make_snapshot source_columns=["id", "name"], ) - adapter_mock.drop_table.assert_called_once_with(f"sqlmesh__db.db__seed__{snapshot.version}") + adapter_mock.drop_table.assert_called_once_with( + f"sqlmesh__db.db__seed__{snapshot.version}" + ) def test_create_seed_no_intervals(mocker: MockerFixture, adapter_mock, make_snapshot): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL ( name db.seed, kind SEED ( path '../seeds/waiter_names.csv', ) ); - """ - ) + """) - model = load_sql_based_model(expressions, path=Path("./examples/sushi/models/test_model.sql")) + model = load_sql_based_model( + expressions, path=Path("./examples/sushi/models/test_model.sql") + ) snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -3213,7 +3280,9 @@ def test_create_seed_no_intervals(mocker: MockerFixture, adapter_mock, make_snap def test_standalone_audit(mocker: MockerFixture, adapter_mock, make_snapshot): evaluator = SnapshotEvaluator(adapter_mock) - audit = StandaloneAudit(name="test_standalone_audit", query=parse_one("SELECT NULL LIMIT 0")) + audit = StandaloneAudit( + name="test_standalone_audit", query=parse_one("SELECT NULL LIMIT 0") + ) snapshot = make_snapshot(audit) snapshot.categorize_as(SnapshotChangeCategory.NON_BREAKING) @@ -3283,7 +3352,9 @@ def test_standalone_audit(mocker: MockerFixture, adapter_mock, make_snapshot): adapter_mock.session.assert_not_called() -def test_audit_wap(adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot]) -> None: +def test_audit_wap( + adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot] +) -> None: evaluator = SnapshotEvaluator(adapter_mock) custom_audit = ModelAudit( @@ -3376,8 +3447,7 @@ def blocking_value(evaluator): return False model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_table, kind FULL, @@ -3387,8 +3457,7 @@ def blocking_value(evaluator): ); SELECT a::int FROM tbl - """ - ), + """), # type: ignore audit_definitions={always_failing_audit.name: always_failing_audit}, ) snapshot = make_snapshot(model) @@ -3427,8 +3496,7 @@ def test_create_post_statements_use_non_deployable_table( evaluator = SnapshotEvaluator(adapter_mock) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind FULL, @@ -3440,8 +3508,7 @@ def test_create_post_statements_use_non_deployable_table( SELECT a::int FROM tbl; CREATE INDEX IF NOT EXISTS test_idx ON test_schema.test_model(a); - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -3523,7 +3590,9 @@ def model_with_statements(context, **kwargs): assert post_calls[0].sql(dialect="postgres") == expected_call -def test_on_virtual_update_statements(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_on_virtual_update_statements( + mocker: MockerFixture, adapter_mock, make_snapshot +): evaluator = SnapshotEvaluator(adapter_mock) @macro() @@ -3531,8 +3600,7 @@ def create_log_table(evaluator, view_name): return f"CREATE OR REPLACE TABLE log_table AS SELECT '{view_name}' as fqn_this_model, '{evaluator.this_model}' as eval_this_model" model = load_sql_based_model( - d.parse( - """ + d.parse(""" MODEL ( name test_schema.test_model, kind FULL, @@ -3551,8 +3619,7 @@ def create_log_table(evaluator, view_name): @create_log_table(@this_model); ON_VIRTUAL_UPDATE_END; - """ - ), + """), ) snapshot = make_snapshot(model) @@ -3600,7 +3667,9 @@ def create_log_table(evaluator, view_name): ) -def test_on_virtual_update_python_model_macro(mocker: MockerFixture, adapter_mock, make_snapshot): +def test_on_virtual_update_python_model_macro( + mocker: MockerFixture, adapter_mock, make_snapshot +): evaluator = SnapshotEvaluator(adapter_mock) @macro() @@ -3668,7 +3737,9 @@ def model_with_statements(context, **kwargs): ) -def test_evaluate_incremental_by_partition(mocker: MockerFixture, make_snapshot, adapter_mock): +def test_evaluate_incremental_by_partition( + mocker: MockerFixture, make_snapshot, adapter_mock +): model = SqlModel( name="test_schema.test_model", query=parse_one("SELECT 1, ds, b FROM tbl_a"), @@ -3755,9 +3826,7 @@ def insert( custom_insert_query_or_df = query_or_df evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind CUSTOM ( @@ -3769,9 +3838,7 @@ def insert( ); SELECT * FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -3792,7 +3859,9 @@ def insert( assert custom_insert_query_or_df.sql() == 'SELECT * FROM "tbl" AS "tbl"' -def test_custom_materialization_strategy_with_custom_properties(adapter_mock, make_snapshot): +def test_custom_materialization_strategy_with_custom_properties( + adapter_mock, make_snapshot +): custom_insert_kind = None class TestCustomKind(CustomKind): @@ -3829,9 +3898,7 @@ def insert( evaluator = SnapshotEvaluator(adapter_mock) with pytest.raises(ConfigError, match=r"primary_key must be specified"): - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind CUSTOM ( @@ -3840,13 +3907,9 @@ def insert( ); SELECT * FROM tbl; - """ - ) - ) + """)) # type: ignore - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind CUSTOM ( @@ -3858,9 +3921,7 @@ def insert( ); SELECT * FROM tbl; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -3892,9 +3953,7 @@ class TestCustomKind: class TestCustomMaterializationStrategy(CustomMaterialization[TestCustomKind]): # type: ignore NAME = "custom_materialization_test_2" - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind CUSTOM ( @@ -3903,9 +3962,7 @@ class TestCustomMaterializationStrategy(CustomMaterialization[TestCustomKind]): ); SELECT * FROM tbl; - """ - ) - ) + """)) # type: ignore with pytest.raises( SQLMeshError, match=r"kind 'TestCustomKind' must be a subclass of CustomKind" @@ -3916,9 +3973,7 @@ class TestCustomMaterializationStrategy(CustomMaterialization[TestCustomKind]): def test_create_managed(adapter_mock, make_snapshot, mocker: MockerFixture): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind MANAGED, @@ -3930,14 +3985,14 @@ def test_create_managed(adapter_mock, make_snapshot, mocker: MockerFixture): ); select a, b from foo; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.all_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.all_deployable() + ) adapter_mock.create_managed_table.assert_called_with( table_name=snapshot.table_name(), @@ -3955,9 +4010,7 @@ def test_create_managed(adapter_mock, make_snapshot, mocker: MockerFixture): def test_create_managed_dev(adapter_mock, make_snapshot, mocker: MockerFixture): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind MANAGED, @@ -3969,14 +4022,14 @@ def test_create_managed_dev(adapter_mock, make_snapshot, mocker: MockerFixture): ); select a, b from foo; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) - evaluator.create([snapshot], {}, deployability_index=DeployabilityIndex.none_deployable()) + evaluator.create( + [snapshot], {}, deployability_index=DeployabilityIndex.none_deployable() + ) adapter_mock.ctas.assert_called_once_with( f"{snapshot.table_name()}__dev", @@ -3996,9 +4049,7 @@ def test_create_managed_dev(adapter_mock, make_snapshot, mocker: MockerFixture): def test_evaluate_managed(adapter_mock, make_snapshot, mocker: MockerFixture): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind MANAGED, @@ -4010,9 +4061,7 @@ def test_evaluate_managed(adapter_mock, make_snapshot, mocker: MockerFixture): ); select a, b from foo; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -4066,15 +4115,15 @@ def test_evaluate_managed(adapter_mock, make_snapshot, mocker: MockerFixture): column_descriptions=model.column_descriptions, source_columns=None, ) - adapter_mock.columns.assert_called_once_with(snapshot.table_name(is_deployable=False)) + adapter_mock.columns.assert_called_once_with( + snapshot.table_name(is_deployable=False) + ) def test_cleanup_managed(adapter_mock, make_snapshot, mocker: MockerFixture): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind MANAGED, @@ -4086,16 +4135,16 @@ def test_cleanup_managed(adapter_mock, make_snapshot, mocker: MockerFixture): ); select a, b from foo; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) adapter_mock.assert_not_called() - cleanup_task = SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=False) + cleanup_task = SnapshotTableCleanupTask( + snapshot=snapshot.table_info, dev_table_only=False + ) evaluator.cleanup(target_snapshots=[cleanup_task]) @@ -4109,18 +4158,14 @@ def test_create_managed_forward_only_with_previous_version_doesnt_clone_for_dev_ ): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind MANAGED ); select a, b from foo; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) @@ -4149,10 +4194,14 @@ def test_create_managed_forward_only_with_previous_version_doesnt_clone_for_dev_ # The table gets created using ctas() because the model column types arent known adapter_mock.ctas.assert_called_once() - assert adapter_mock.ctas.call_args_list[0].args[0] == snapshot.table_name(is_deployable=False) + assert adapter_mock.ctas.call_args_list[0].args[0] == snapshot.table_name( + is_deployable=False + ) -def test_migrate_snapshot(snapshot: Snapshot, mocker: MockerFixture, adapter_mock, make_snapshot): +def test_migrate_snapshot( + snapshot: Snapshot, mocker: MockerFixture, adapter_mock, make_snapshot +): adapter_mock = mocker.patch("sqlmesh.core.engine_adapter.EngineAdapter") adapter_mock.dialect = "duckdb" adapter_mock.with_settings.return_value = adapter_mock @@ -4219,7 +4268,9 @@ def test_migrate_snapshot(snapshot: Snapshot, mocker: MockerFixture, adapter_moc adapter_mock.fetchall.assert_has_calls( [ call( - parse_one('SELECT CAST("a" AS INT) AS "a" FROM "tbl" AS "tbl" WHERE FALSE LIMIT 0') + parse_one( + 'SELECT CAST("a" AS INT) AS "a" FROM "tbl" AS "tbl" WHERE FALSE LIMIT 0' + ) ), call( parse_one( @@ -4297,18 +4348,14 @@ def apply_side_effect(snapshot_iterable, fn, *_args, **_kwargs): def test_migrate_managed(adapter_mock, make_snapshot, mocker: MockerFixture): evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind MANAGED ); select a, b from foo; - """ - ) - ) + """)) # type: ignore snapshot: Snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING, forward_only=True) snapshot.previous_versions = snapshot.all_versions @@ -4356,7 +4403,11 @@ def test_migrate_managed(adapter_mock, make_snapshot, mocker: MockerFixture): def test_multiple_engine_creation(snapshot: Snapshot, adapters, make_snapshot): - engine_adapters = {"default": adapters[0], "secondary": adapters[1], "third": adapters[2]} + engine_adapters = { + "default": adapters[0], + "secondary": adapters[1], + "third": adapters[2], + } evaluator = SnapshotEvaluator(engine_adapters) assert len(evaluator.adapters) == 3 @@ -4366,8 +4417,7 @@ def test_multiple_engine_creation(snapshot: Snapshot, adapters, make_snapshot): assert evaluator.get_adapter("secondary") == engine_adapters["secondary"] model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind FULL, @@ -4377,8 +4427,7 @@ def test_multiple_engine_creation(snapshot: Snapshot, adapters, make_snapshot): CREATE INDEX IF NOT EXISTS test_idx ON test_schema.test_model(a); SELECT a::int FROM tbl; CREATE INDEX IF NOT EXISTS test_idx ON test_schema.test_model(a); - """ - ), + """), # type: ignore ) snapshot_2 = make_snapshot(model) @@ -4559,7 +4608,9 @@ def columns(table_name): adapter_one.cursor.execute.assert_has_calls( [ - call('ALTER TABLE "sqlmesh__test_schema"."test_schema__test_model__1" DROP COLUMN "b"'), + call( + 'ALTER TABLE "sqlmesh__test_schema"."test_schema__test_model__1" DROP COLUMN "b"' + ), call( 'ALTER TABLE "sqlmesh__test_schema"."test_schema__test_model__1" ADD COLUMN "a" INT' ), @@ -4580,16 +4631,14 @@ def test_multiple_engine_cleanup(snapshot: Snapshot, adapters, make_snapshot): evaluator = SnapshotEvaluator(engine_adapters) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind FULL, gateway secondary, ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot_2 = make_snapshot(model) @@ -4607,7 +4656,9 @@ def test_multiple_engine_cleanup(snapshot: Snapshot, adapters, make_snapshot): evaluator.cleanup( [ SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True), - SnapshotTableCleanupTask(snapshot=snapshot_2.table_info, dev_table_only=True), + SnapshotTableCleanupTask( + snapshot=snapshot_2.table_info, dev_table_only=True + ), ], ) @@ -4616,7 +4667,8 @@ def test_multiple_engine_cleanup(snapshot: Snapshot, adapters, make_snapshot): f"sqlmesh__db.db__model__{snapshot.version}__dev", cascade=True ) engine_adapters["secondary"].drop_table.assert_called_once_with( - f"sqlmesh__test_schema.test_schema__test_model__{snapshot_2.version}__dev", cascade=True + f"sqlmesh__test_schema.test_schema__test_model__{snapshot_2.version}__dev", + cascade=True, ) @@ -4625,16 +4677,14 @@ def test_cleanup_skips_unavailable_gateway(snapshot: Snapshot, adapters, make_sn evaluator = SnapshotEvaluator(engine_adapters) model_with_missing_gw = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind FULL, gateway nonexistent_gateway, ); SELECT a::int FROM tbl; - """ - ), + """), # type: ignore ) snapshot_missing_gw = make_snapshot(model_with_missing_gw) @@ -4646,7 +4696,9 @@ def test_cleanup_skips_unavailable_gateway(snapshot: Snapshot, adapters, make_sn evaluator.cleanup( [ SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True), - SnapshotTableCleanupTask(snapshot=snapshot_missing_gw.table_info, dev_table_only=True), + SnapshotTableCleanupTask( + snapshot=snapshot_missing_gw.table_info, dev_table_only=True + ), ], ) @@ -4706,7 +4758,9 @@ def model_with_statements(context, **kwargs): # Validate model-specific gateway usage during table creation create_args = engine_adapters["secondary"].create_table.call_args_list assert len(create_args) == 1 - assert create_args[0][0] == (f"sqlmesh__db.db__multi_engine_test_model__{snapshot.version}",) + assert create_args[0][0] == ( + f"sqlmesh__db.db__multi_engine_test_model__{snapshot.version}", + ) environment_naming_info = EnvironmentNamingInfo(name="test_env") evaluator.promote([snapshot], environment_naming_info) @@ -4731,7 +4785,9 @@ def model_with_statements(context, **kwargs): cascade=False, ) - environment_naming_info_gw = EnvironmentNamingInfo(name="test_env", gateway_managed=True) + 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 @@ -4747,12 +4803,15 @@ def model_with_statements(context, **kwargs): def test_multiple_engine_virtual_layer(snapshot: Snapshot, adapters, make_snapshot): - engine_adapters = {"default": adapters[0], "secondary": adapters[1], "third": adapters[2]} + engine_adapters = { + "default": adapters[0], + "secondary": adapters[1], + "third": adapters[2], + } evaluator = SnapshotEvaluator(engine_adapters) model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.test_model, kind FULL, @@ -4761,8 +4820,7 @@ def test_multiple_engine_virtual_layer(snapshot: Snapshot, adapters, make_snapsh ); SELECT a::int FROM tbl; CREATE INDEX IF NOT EXISTS test_idx ON test_schema.test_model(a); - """ - ), + """), # type: ignore ) snapshot_2 = make_snapshot(model) @@ -4782,7 +4840,9 @@ 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=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) @@ -4903,7 +4963,9 @@ def test_wap_model_wap_supported( ) -def test_wap_no_wap_support(adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot]) -> None: +def test_wap_no_wap_support( + adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot] +) -> None: evaluator = SnapshotEvaluator(adapter_mock) model = SqlModel( @@ -4945,14 +5007,20 @@ def test_wap_non_materialized_snapshot( adapter_mock.wap_supported.return_value = True wap_id = evaluator.evaluate( - snapshot, start="2020-01-01", end="2020-01-01", execution_time="2020-01-01", snapshots={} + snapshot, + start="2020-01-01", + end="2020-01-01", + execution_time="2020-01-01", + snapshots={}, ) assert wap_id is None adapter_mock.wap_prepare.assert_not_called() -def test_wap_publish_snapshot(adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot]) -> None: +def test_wap_publish_snapshot( + adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot] +) -> None: evaluator = SnapshotEvaluator(adapter_mock) model = SqlModel( @@ -4972,7 +5040,9 @@ def test_wap_publish_snapshot(adapter_mock: Mock, make_snapshot: t.Callable[..., adapter_mock.wap_publish.assert_called_once_with(expected_table_name, wap_id) -def test_wap_during_audit(adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot]) -> None: +def test_wap_during_audit( + adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot] +) -> None: evaluator = SnapshotEvaluator(adapter_mock) custom_audit = ModelAudit( @@ -5006,7 +5076,9 @@ def test_wap_during_audit(adapter_mock: Mock, make_snapshot: t.Callable[..., Sna adapter_mock.wap_publish.assert_called_once_with(snapshot.table_name(), wap_id) -def test_wap_prepare_failure(adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot]) -> None: +def test_wap_prepare_failure( + adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot] +) -> None: evaluator = SnapshotEvaluator(adapter_mock) model = SqlModel( @@ -5032,7 +5104,9 @@ def test_wap_prepare_failure(adapter_mock: Mock, make_snapshot: t.Callable[..., ) -def test_wap_publish_failure(adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot]) -> None: +def test_wap_publish_failure( + adapter_mock: Mock, make_snapshot: t.Callable[..., Snapshot] +) -> None: """Test error handling when WAP publish fails.""" evaluator = SnapshotEvaluator(adapter_mock) @@ -5097,8 +5171,7 @@ def mutate_view_properties(*args, **kwargs): # create a view model with SECURITY INVOKER physical property # AND self referenctial to trigger two create statements model = load_sql_based_model( - parse( # type: ignore - """ + parse(""" MODEL ( name test_schema.security_view, kind VIEW, @@ -5108,8 +5181,7 @@ def mutate_view_properties(*args, **kwargs): ); SELECT 1 as col from test_schema.security_view; - """ - ), + """), # type: ignore ) snapshot = make_snapshot(model) @@ -5139,10 +5211,11 @@ def _create_grants_test_model( grants=None, kind="FULL", grants_target_layer=None, virtual_environment_mode=None ): if kind == "SEED": + import os + import tempfile + from sqlmesh.core.model.definition import create_seed_model from sqlmesh.core.model.kind import SeedKind - import tempfile - import os # Create a temporary CSV file for the test temp_csv = tempfile.NamedTemporaryFile(mode="w", suffix=".csv", delete=False) @@ -5204,7 +5277,9 @@ def _create_grants_test_model( return create_sql_model( "test_model", - parse_one("SELECT 1 as id, CURRENT_DATE as ds, CURRENT_TIMESTAMP as updated_at"), + parse_one( + "SELECT 1 as id, CURRENT_DATE as ds, CURRENT_TIMESTAMP as updated_at" + ), **kwargs, ) @@ -5362,7 +5437,10 @@ def test_grants_update( evaluator.create([new_snapshot], {}) sync_grants_mock.assert_called_once() - assert sync_grants_mock.call_args[0][1] == {"select": ["user2", "user3"], "insert": ["admin"]} + assert sync_grants_mock.call_args[0][1] == { + "select": ["user2", "user3"], + "insert": ["admin"], + } # Update model query AND remove grants updated_model_dict = model.dict() @@ -5388,9 +5466,7 @@ def test_grants_create_and_evaluate( evaluator = SnapshotEvaluator(adapter_mock) - model = load_sql_based_model( - parse( # type: ignore - """ + model = load_sql_based_model(parse(""" MODEL ( name test_schema.test_model, kind INCREMENTAL_BY_TIME_RANGE (time_column ds), @@ -5401,9 +5477,7 @@ def test_grants_create_and_evaluate( grants_target_layer 'all' ); SELECT ds::DATE, value::INT FROM source WHERE ds BETWEEN @start_ds AND @end_ds; - """ - ) - ) + """)) # type: ignore snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -5417,7 +5491,11 @@ def test_grants_create_and_evaluate( sync_grants_mock.reset_mock() evaluator.evaluate( - snapshot, start="2020-01-01", end="2020-01-02", execution_time="2020-01-02", snapshots={} + snapshot, + start="2020-01-01", + end="2020-01-02", + execution_time="2020-01-02", + snapshots={}, ) # Evaluate should not reapply grants sync_grants_mock.assert_not_called() @@ -5447,7 +5525,9 @@ def test_grants_materializable_strategy_migrate( sync_grants_mock = mocker.patch.object(adapter_mock, "sync_grants_config") strategy = strategy_class(adapter_mock) grants = {"select": ["user1"]} - model = _create_grants_test_model(grants=grants, grants_target_layer=GrantsTargetLayer.ALL) + model = _create_grants_test_model( + grants=grants, grants_target_layer=GrantsTargetLayer.ALL + ) snapshot = make_snapshot(model) strategy.migrate( @@ -5472,7 +5552,9 @@ def test_grants_clone_snapshot_in_dev( evaluator = SnapshotEvaluator(adapter_mock) grants = {"select": ["user1", "user2"]} - model = _create_grants_test_model(grants=grants, grants_target_layer=GrantsTargetLayer.ALL) + model = _create_grants_test_model( + grants=grants, grants_target_layer=GrantsTargetLayer.ALL + ) snapshot = make_snapshot(model) snapshot.categorize_as(SnapshotChangeCategory.BREAKING) @@ -5625,7 +5707,8 @@ def test_grants_in_production_with_dev_only_vde( adapter_mock.SUPPORTS_GRANTS = True sync_grants_mock = mocker.patch.object(adapter_mock, "sync_grants_config") - from sqlmesh.core.model.meta import VirtualEnvironmentMode, GrantsTargetLayer + from sqlmesh.core.model.meta import (GrantsTargetLayer, + VirtualEnvironmentMode) from sqlmesh.core.snapshot.definition import DeployabilityIndex model_virtual_grants = _create_grants_test_model( @@ -5642,7 +5725,10 @@ def test_grants_in_production_with_dev_only_vde( evaluator.create([snapshot], {}, deployability_index=deployability_index) sync_grants_mock.assert_called_once() - assert sync_grants_mock.call_args[0][1] == {"select": ["user1"], "insert": ["role1"]} + assert sync_grants_mock.call_args[0][1] == { + "select": ["user1"], + "insert": ["role1"], + } # Non-deployable (dev) env sync_grants_mock.reset_mock() @@ -5653,4 +5739,7 @@ def test_grants_in_production_with_dev_only_vde( else: # Should still apply grants to physical table when target layer is ALL or PHYSICAL sync_grants_mock.assert_called_once() - assert sync_grants_mock.call_args[0][1] == {"select": ["user1"], "insert": ["role1"]} + assert sync_grants_mock.call_args[0][1] == { + "select": ["user1"], + "insert": ["role1"], + } diff --git a/tests/core/test_table_diff.py b/tests/core/test_table_diff.py index c2e293e4c2..031931653a 100644 --- a/tests/core/test_table_diff.py +++ b/tests/core/test_table_diff.py @@ -1,18 +1,21 @@ +import typing as t +from io import StringIO + +import numpy as np # noqa: TID253 +import pandas as pd # noqa: TID253 import pytest from pytest_mock.plugin import MockerFixture -import pandas as pd # noqa: TID253 +from rich.console import Console from sqlglot import exp + from sqlmesh.core import dialect as d -import typing as t -from io import StringIO -from rich.console import Console +from sqlmesh.core.config import (AutoCategorizationMode, CategorizerConfig, + DuckDBConnectionConfig) from sqlmesh.core.console import TerminalConsole from sqlmesh.core.context import Context -from sqlmesh.core.config import AutoCategorizationMode, CategorizerConfig, DuckDBConnectionConfig from sqlmesh.core.model import SqlModel, load_sql_based_model from sqlmesh.core.model.common import ParsableSql -from sqlmesh.core.table_diff import TableDiff, SchemaDiff -import numpy as np # noqa: TID253 +from sqlmesh.core.table_diff import SchemaDiff, TableDiff from sqlmesh.utils.errors import SQLMeshError from sqlmesh.utils.rich import strip_ansi_codes @@ -47,14 +50,16 @@ def capture_console_output(method_name: str, **kwargs) -> str: def test_data_diff(sushi_context_fixed_date, capsys, caplog): - model = sushi_context_fixed_date.models['"memory"."sushi"."customer_revenue_by_day"'] + model = sushi_context_fixed_date.models[ + '"memory"."sushi"."customer_revenue_by_day"' + ] sushi_context_fixed_date.upsert_model( model, query_=ParsableSql( - sql=model.query.select(exp.cast("'1'", "VARCHAR").as_("modified_col"), "1 AS y").sql( - model.dialect - ) + sql=model.query.select( + exp.cast("'1'", "VARCHAR").as_("modified_col"), "1 AS y" + ).sql(model.dialect) ), ) @@ -66,7 +71,9 @@ def test_data_diff(sushi_context_fixed_date, capsys, caplog): start="2023-01-31", end="2023-01-31", ) - model = sushi_context_fixed_date.models['"memory"."sushi"."customer_revenue_by_day"'] + model = sushi_context_fixed_date.models[ + '"memory"."sushi"."customer_revenue_by_day"' + ] for column in model.query.find_all(exp.Column): if column.name == "total": @@ -112,13 +119,17 @@ def test_data_diff(sushi_context_fixed_date, capsys, caplog): diff = sushi_context_fixed_date.table_diff( source="source_dev", target="target_dev", - on=exp.condition("s.customer_id = t.customer_id AND s.event_date = t.event_date"), + on=exp.condition( + "s.customer_id = t.customer_id AND s.event_date = t.event_date" + ), select_models={"sushi.customer_revenue_by_day"}, )[0] # verify queries were actually logged to the log file, this helps immensely with debugging console_output = capsys.readouterr() - assert "__sqlmesh_join_key" not in console_output # they should not go to the console + assert ( + "__sqlmesh_join_key" not in console_output + ) # they should not go to the console assert "__sqlmesh_join_key" in caplog.text schema_diff = diff.schema_diff() @@ -265,7 +276,11 @@ def test_data_diff_decimals_on_numeric(): for decimals in range(5, 0, -1): table_diff = TableDiff( - adapter=engine_adapter, source="src", target="target", on=["key"], decimals=decimals + adapter=engine_adapter, + source="src", + target="target", + on=["key"], + decimals=decimals, ) diff = table_diff.row_diff() @@ -274,8 +289,7 @@ def test_data_diff_decimals_on_numeric(): def test_grain_check(sushi_context_fixed_date): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name memory.sushi.grain_items, kind full, grain("key_1", KEY_2)); SELECT key_1 as "key_1", @@ -291,8 +305,7 @@ def test_grain_check(sushi_context_fixed_date): (4, NULL, 3), (2, 3, 2), ) AS t (key_1,KEY_2, value) - """ - ) + """) model_s = load_sql_based_model(expressions, dialect="snowflake") sushi_context_fixed_date.upsert_model(model_s) sushi_context_fixed_date.plan( @@ -391,11 +404,15 @@ def test_generated_sql(sushi_context_fixed_date: Context, mocker: MockerFixture) summary_query_sql = 'SELECT SUM("s_exists") AS "s_count", SUM("t_exists") AS "t_count", SUM("row_joined") AS "join_count", SUM("null_grain") AS "null_grain_count", SUM("row_full_match") AS "full_match_count", SUM("key_matches") AS "key_matches", SUM("value_matches") AS "value_matches", COUNT(DISTINCT ("s____sqlmesh_join_key")) AS "distinct_count_s", COUNT(DISTINCT ("t____sqlmesh_join_key")) AS "distinct_count_t" FROM "memory"."sqlmesh_temp_test"."__temp_diff_abcdefgh"' compare_sql = 'SELECT ROUND(100 * (CAST(SUM("key_matches") AS DECIMAL) / COUNT("key_matches")), 9) AS "key_matches", ROUND(100 * (CAST(SUM("value_matches") AS DECIMAL) / COUNT("value_matches")), 9) AS "value_matches" FROM "memory"."sqlmesh_temp_test"."__temp_diff_abcdefgh" WHERE "row_joined" = 1' sample_query_sql = 'WITH "source_only" AS (SELECT \'source_only\' AS "__sqlmesh_sample_type", "s__key", "s__value", "s____sqlmesh_join_key", "t__key", "t__value", "t____sqlmesh_join_key" FROM "memory"."sqlmesh_temp_test"."__temp_diff_abcdefgh" WHERE "s_exists" = 1 AND "row_joined" = 0 ORDER BY "s__key" NULLS FIRST LIMIT 20), "target_only" AS (SELECT \'target_only\' AS "__sqlmesh_sample_type", "s__key", "s__value", "s____sqlmesh_join_key", "t__key", "t__value", "t____sqlmesh_join_key" FROM "memory"."sqlmesh_temp_test"."__temp_diff_abcdefgh" WHERE "t_exists" = 1 AND "row_joined" = 0 ORDER BY "t__key" NULLS FIRST LIMIT 20), "common_rows" AS (SELECT \'common_rows\' AS "__sqlmesh_sample_type", "s__key", "s__value", "s____sqlmesh_join_key", "t__key", "t__value", "t____sqlmesh_join_key" FROM "memory"."sqlmesh_temp_test"."__temp_diff_abcdefgh" WHERE "row_joined" = 1 AND "row_full_match" = 0 ORDER BY "s__key" NULLS FIRST, "t__key" NULLS FIRST LIMIT 20) SELECT "__sqlmesh_sample_type", "s__key", "s__value", "s____sqlmesh_join_key", "t__key", "t__value", "t____sqlmesh_join_key" FROM "source_only" UNION ALL SELECT "__sqlmesh_sample_type", "s__key", "s__value", "s____sqlmesh_join_key", "t__key", "t__value", "t____sqlmesh_join_key" FROM "target_only" UNION ALL SELECT "__sqlmesh_sample_type", "s__key", "s__value", "s____sqlmesh_join_key", "t__key", "t__value", "t____sqlmesh_join_key" FROM "common_rows"' - drop_sql = 'DROP TABLE IF EXISTS "memory"."sqlmesh_temp_test"."__temp_diff_abcdefgh"' + drop_sql = ( + 'DROP TABLE IF EXISTS "memory"."sqlmesh_temp_test"."__temp_diff_abcdefgh"' + ) # make with_settings() return the current instance of engine_adapter so we can still spy on _execute mocker.patch.object( - engine_adapter, "with_settings", new_callable=lambda: lambda **kwargs: engine_adapter + engine_adapter, + "with_settings", + new_callable=lambda: lambda **kwargs: engine_adapter, ) assert engine_adapter.with_settings() == engine_adapter @@ -432,7 +449,8 @@ def test_generated_sql(sushi_context_fixed_date: Context, mocker: MockerFixture) def test_tables_and_grain_inferred_from_model(sushi_context_fixed_date: Context): - (sushi_context_fixed_date.path / "models" / "waiter_revenue_by_day.sql").write_text(""" + (sushi_context_fixed_date.path / "models" / "waiter_revenue_by_day.sql").write_text( + """ MODEL ( name sushi.waiter_revenue_by_day, kind incremental_by_time_range ( @@ -461,13 +479,16 @@ def test_tables_and_grain_inferred_from_model(sushi_context_fixed_date: Context) GROUP BY o.waiter_id, o.event_date -""") +""" + ) # this creates a dev preview of "sushi.waiter_revenue_by_day" sushi_context_fixed_date.refresh() sushi_context_fixed_date.auto_categorize_changes = CategorizerConfig( sql=AutoCategorizationMode.FULL ) - sushi_context_fixed_date.plan(environment="unit_test", auto_apply=True, include_unmodified=True) + sushi_context_fixed_date.plan( + environment="unit_test", auto_apply=True, include_unmodified=True + ) table_diff = sushi_context_fixed_date.table_diff( source="unit_test", target="prod", select_models={"sushi.waiter_revenue_by_day"} @@ -488,7 +509,11 @@ def test_data_diff_array_dict(sushi_context_fixed_date): pd.DataFrame( { "key": [1, 2, 3], - "value": [np.array([51.2, 4.5678]), np.array([2.31, 12.2]), np.array([5.0])], + "value": [ + np.array([51.2, 4.5678]), + np.array([2.31, 12.2]), + np.array([5.0]), + ], "dict": [{"key1": 10, "key2": 20, "key3": 30}, {"key1": 10}, {}], } ), @@ -565,7 +590,10 @@ def test_data_diff_array_dict(sushi_context_fixed_date): def test_data_diff_array_struct_query(): engine_adapter = DuckDBConnectionConfig().create_engine_adapter() - columns_to_types = {"key": exp.DataType.build("int"), "value": exp.DataType.build("int")} + columns_to_types = { + "key": exp.DataType.build("int"), + "value": exp.DataType.build("int"), + } engine_adapter.create_table("table_diff_source", columns_to_types) engine_adapter.create_table("table_diff_target", columns_to_types) @@ -597,9 +625,7 @@ def test_data_diff_array_struct_query(): output = capture_console_output("show_row_diff", row_diff=diff) - assert ( - strip_ansi_codes(output) - == """Row Counts: + assert strip_ansi_codes(output) == """Row Counts: └── PARTIAL MATCH: 1 rows (100.0%) COMMON ROWS column comparison stats: @@ -622,13 +648,15 @@ def test_data_diff_array_struct_query(): │ 1 │ {'k': 10, 'v': 11} │ {'k': 11, 'v': 10} │ └─────┴────────────────────┴────────────────────┘ """.strip() - ) def test_data_diff_nullable_booleans(): engine_adapter = DuckDBConnectionConfig().create_engine_adapter() - columns_to_types = {"key": exp.DataType.build("int"), "value": exp.DataType.build("boolean")} + columns_to_types = { + "key": exp.DataType.build("int"), + "value": exp.DataType.build("boolean"), + } engine_adapter.create_table("table_diff_source", columns_to_types) engine_adapter.create_table("table_diff_target", columns_to_types) @@ -678,8 +706,7 @@ def test_data_diff_nullable_booleans(): def test_data_diff_multiple_models(sushi_context_fixed_date, capsys, caplog): # Create first analytics model - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name memory.sushi.analytics_1, kind full, grain(key), tags (finance),); SELECT key, @@ -689,22 +716,19 @@ def test_data_diff_multiple_models(sushi_context_fixed_date, capsys, caplog): (1, 3), (2, 4), ) AS t (key, value) - """ - ) + """) model_s = load_sql_based_model(expressions, dialect="snowflake") sushi_context_fixed_date.upsert_model(model_s) # Create second analytics model from analytics_1 - expressions_2 = d.parse( - """ + expressions_2 = d.parse(""" MODEL (name memory.sushi.analytics_2, kind full, grain(key), tags (finance),); SELECT key, value as amount, FROM memory.sushi.analytics_1 - """ - ) + """) model_s2 = load_sql_based_model(expressions_2, dialect="snowflake") sushi_context_fixed_date.upsert_model(model_s2) @@ -734,7 +758,9 @@ def test_data_diff_multiple_models(sushi_context_fixed_date, capsys, caplog): modified_model2["query"] = ( exp.select("*") .from_(model2.query.subquery()) - .union("SELECT key, amount FROM (VALUES (5, 150.2),(6,250.2),) AS t (key, amount)") + .union( + "SELECT key, amount FROM (VALUES (5, 150.2),(6,250.2),) AS t (key, amount)" + ) ) modified_sqlmodel2 = SqlModel(**modified_model2) sushi_context_fixed_date.upsert_model(modified_sqlmodel2) @@ -809,8 +835,7 @@ def test_data_diff_multiple_models(sushi_context_fixed_date, capsys, caplog): def test_data_diff_forward_only(sushi_context_fixed_date, capsys, caplog): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name memory.sushi.full_1, kind full, grain(key),); SELECT key, @@ -820,22 +845,19 @@ def test_data_diff_forward_only(sushi_context_fixed_date, capsys, caplog): (1, 3), (2, 4), ) AS t (key, value) - """ - ) + """) model_s = load_sql_based_model(expressions, dialect="snowflake") sushi_context_fixed_date.upsert_model(model_s) # Create second analytics model sourcing from first - expressions_2 = d.parse( - """ + expressions_2 = d.parse(""" MODEL (name memory.sushi.full_2, kind full, grain(key),); SELECT key, value as amount, FROM memory.sushi.full_1 - """ - ) + """) model_s2 = load_sql_based_model(expressions_2, dialect="snowflake") sushi_context_fixed_date.upsert_model(model_s2) @@ -850,7 +872,9 @@ def test_data_diff_forward_only(sushi_context_fixed_date, capsys, caplog): model = sushi_context_fixed_date.models['"MEMORY"."SUSHI"."FULL_1"'] modified_model = model.dict() - modified_model["query"] = exp.select("*").from_("(VALUES (12, 6),(5,3),) AS t (key, value)") + modified_model["query"] = exp.select("*").from_( + "(VALUES (12, 6),(5,3),) AS t (key, value)" + ) modified_sqlmodel = SqlModel(**modified_model) sushi_context_fixed_date.upsert_model(modified_sqlmodel) @@ -937,15 +961,17 @@ def test_data_diff_empty_tables(): output = capture_console_output("show_row_diff", row_diff=row_diff) assert ( - strip_ansi_codes(output) == "Neither the source nor the target table contained any records" + strip_ansi_codes(output) + == "Neither the source nor the target table contained any records" ) @pytest.mark.slow -def test_data_diff_multiple_models_lacking_grain(sushi_context_fixed_date, capsys, caplog): +def test_data_diff_multiple_models_lacking_grain( + sushi_context_fixed_date, capsys, caplog +): # Create first model with grain - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name memory.sushi.grain_model, kind full, grain(key),); SELECT key, @@ -955,22 +981,19 @@ def test_data_diff_multiple_models_lacking_grain(sushi_context_fixed_date, capsy (1, 3), (2, 4), ) AS t (key, value) - """ - ) + """) model_s = load_sql_based_model(expressions, dialect="snowflake") sushi_context_fixed_date.upsert_model(model_s) # Create second model without grain - expressions_2 = d.parse( - """ + expressions_2 = d.parse(""" MODEL (name memory.sushi.no_grain_model, kind full,); SELECT key, value as amount, FROM memory.sushi.grain_model - """ - ) + """) model_s2 = load_sql_based_model(expressions_2, dialect="snowflake") sushi_context_fixed_date.upsert_model(model_s2) @@ -1000,7 +1023,9 @@ def test_data_diff_multiple_models_lacking_grain(sushi_context_fixed_date, capsy modified_model2["query"] = ( exp.select("*") .from_(model2.query.subquery()) - .union("SELECT key, amount FROM (VALUES (5, 150.2),(6,250.2),) AS t (key, amount)") + .union( + "SELECT key, amount FROM (VALUES (5, 150.2),(6,250.2),) AS t (key, amount)" + ) ) modified_sqlmodel2 = SqlModel(**modified_model2) sushi_context_fixed_date.upsert_model(modified_sqlmodel2) @@ -1063,8 +1088,14 @@ def test_data_diff_multiple_models_lacking_grain(sushi_context_fixed_date, capsy def test_schema_diff_ignore_case(): # no changes - table_a = {"COL_A": exp.DataType.build("varchar"), "cOl_b": exp.DataType.build("int")} - table_b = {"col_a": exp.DataType.build("varchar"), "COL_b": exp.DataType.build("int")} + table_a = { + "COL_A": exp.DataType.build("varchar"), + "cOl_b": exp.DataType.build("int"), + } + table_b = { + "col_a": exp.DataType.build("varchar"), + "COL_b": exp.DataType.build("int"), + } diff = SchemaDiff( source="table_a", @@ -1077,7 +1108,10 @@ def test_schema_diff_ignore_case(): assert not diff.has_changes # added in target - table_a = {"COL_A": exp.DataType.build("varchar"), "cOl_b": exp.DataType.build("int")} + table_a = { + "COL_A": exp.DataType.build("varchar"), + "cOl_b": exp.DataType.build("int"), + } table_b = { "col_a": exp.DataType.build("varchar"), "COL_b": exp.DataType.build("int"), @@ -1107,7 +1141,10 @@ def test_schema_diff_ignore_case(): "COL_A": exp.DataType.build("varchar"), "cOl_b": exp.DataType.build("int"), } - table_b = {"col_a": exp.DataType.build("varchar"), "COL_b": exp.DataType.build("int")} + table_b = { + "col_a": exp.DataType.build("varchar"), + "COL_b": exp.DataType.build("int"), + } diff = SchemaDiff( source="table_a", @@ -1127,7 +1164,10 @@ def test_schema_diff_ignore_case(): assert not diff.modified # column type change - table_a = {"CoL_A": exp.DataType.build("varchar"), "cOl_b": exp.DataType.build("int")} + table_a = { + "CoL_A": exp.DataType.build("varchar"), + "cOl_b": exp.DataType.build("int"), + } table_b = {"col_a": exp.DataType.build("date"), "COL_b": exp.DataType.build("int")} diff = SchemaDiff( @@ -1152,7 +1192,10 @@ def test_schema_diff_ignore_case(): def test_data_diff_sample_limit(): engine_adapter = DuckDBConnectionConfig().create_engine_adapter() - columns_to_types = {"id": exp.DataType.build("int"), "name": exp.DataType.build("varchar")} + columns_to_types = { + "id": exp.DataType.build("int"), + "name": exp.DataType.build("varchar"), + } engine_adapter.create_table("src", columns_to_types) engine_adapter.create_table("target", columns_to_types) @@ -1175,8 +1218,12 @@ def test_data_diff_sample_limit(): target_records[2] = "modified_target_2" target_records[7] = "modified_target_7" - src_df = pd.DataFrame.from_records([{"id": k, "name": v} for k, v in src_records.items()]) - target_df = pd.DataFrame.from_records([{"id": k, "name": v} for k, v in target_records.items()]) + src_df = pd.DataFrame.from_records( + [{"id": k, "name": v} for k, v in src_records.items()] + ) + target_df = pd.DataFrame.from_records( + [{"id": k, "name": v} for k, v in target_records.items()] + ) engine_adapter.insert_append("src", src_df) engine_adapter.insert_append("target", target_df) @@ -1229,7 +1276,10 @@ def test_data_diff_nulls_in_some_grain_columns(): engine_adapter.insert_append("target", target_df) table_diff = TableDiff( - adapter=engine_adapter, source="src", target="target", on=["key1", "key2", "key3"] + adapter=engine_adapter, + source="src", + target="target", + on=["key1", "key2", "key3"], ) diff = table_diff.row_diff() diff --git a/tests/core/test_test.py b/tests/core/test_test.py index d679f09393..d1abbb7727 100644 --- a/tests/core/test_test.py +++ b/tests/core/test_test.py @@ -1,43 +1,38 @@ from __future__ import annotations import datetime -import typing as t import io -from pathlib import Path +import typing as t import unittest -from unittest.mock import call, patch +from pathlib import Path from shutil import rmtree +from unittest.mock import call, patch import pandas as pd # noqa: TID253 import pytest +from IPython.utils.capture import capture_output from pytest_mock.plugin import MockerFixture from sqlglot import exp -from IPython.utils.capture import capture_output from sqlmesh.cli.project_init import init_example_project from sqlmesh.core import constants as c -from sqlmesh.core.config import ( - Config, - DuckDBConnectionConfig, - SparkConnectionConfig, - GatewayConfig, - ModelDefaultsConfig, -) -from sqlmesh.core.context import Context, ExecutionContext +from sqlmesh.core.config import (Config, DuckDBConnectionConfig, GatewayConfig, + ModelDefaultsConfig, SparkConnectionConfig) from sqlmesh.core.console import get_console +from sqlmesh.core.context import Context, ExecutionContext from sqlmesh.core.dialect import parse from sqlmesh.core.engine_adapter import EngineAdapter from sqlmesh.core.macros import MacroEvaluator, macro from sqlmesh.core.model import Model, SqlModel, load_sql_based_model, model from sqlmesh.core.model.common import ParsableSql -from sqlmesh.core.test.definition import ModelTest, PythonModelTest, SqlModelTest -from sqlmesh.core.test.result import ModelTextTestResult from sqlmesh.core.test.context import TestExecutionContext +from sqlmesh.core.test.definition import (ModelTest, PythonModelTest, + SqlModelTest) +from sqlmesh.core.test.result import ModelTextTestResult from sqlmesh.utils import Verbosity from sqlmesh.utils.errors import ConfigError, SQLMeshError, TestError from sqlmesh.utils.yaml import dump as dump_yaml from sqlmesh.utils.yaml import load as load_yaml - from tests.utils.test_helpers import use_terminal_console if t.TYPE_CHECKING: @@ -77,7 +72,10 @@ def _create_model( return t.cast( SqlModel, load_sql_based_model( - parsed_definition, dialect=dialect, default_catalog=default_catalog, **kwargs + parsed_definition, + dialect=dialect, + default_catalog=default_catalog, + **kwargs, ), ) @@ -133,8 +131,7 @@ def full_model_with_two_ctes(request) -> SqlModel: def test_ctes(sushi_context: Context, full_model_with_two_ctes: SqlModel) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -151,8 +148,7 @@ def test_ctes(sushi_context: Context, full_model_with_two_ctes: SqlModel) -> Non vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_with_two_ctes), context=sushi_context, @@ -163,8 +159,7 @@ def test_ctes(sushi_context: Context, full_model_with_two_ctes: SqlModel) -> Non def test_ctes_only(sushi_context: Context, full_model_with_two_ctes: SqlModel) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -179,8 +174,7 @@ def test_ctes_only(sushi_context: Context, full_model_with_two_ctes: SqlModel) - vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_with_two_ctes), context=sushi_context, @@ -191,8 +185,7 @@ def test_ctes_only(sushi_context: Context, full_model_with_two_ctes: SqlModel) - def test_query_only(sushi_context: Context, full_model_with_two_ctes: SqlModel) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -204,8 +197,7 @@ def test_query_only(sushi_context: Context, full_model_with_two_ctes: SqlModel) vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_with_two_ctes), context=sushi_context, @@ -213,11 +205,12 @@ def test_query_only(sushi_context: Context, full_model_with_two_ctes: SqlModel) ) -def test_with_rows(sushi_context: Context, full_model_with_single_cte: SqlModel) -> None: +def test_with_rows( + sushi_context: Context, full_model_with_single_cte: SqlModel +) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -235,8 +228,7 @@ def test_with_rows(sushi_context: Context, full_model_with_single_cte: SqlModel) vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_with_single_cte), context=sushi_context, @@ -244,11 +236,12 @@ def test_with_rows(sushi_context: Context, full_model_with_single_cte: SqlModel) ) -def test_without_rows(sushi_context: Context, full_model_with_single_cte: SqlModel) -> None: +def test_without_rows( + sushi_context: Context, full_model_with_single_cte: SqlModel +) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -263,8 +256,7 @@ def test_without_rows(sushi_context: Context, full_model_with_single_cte: SqlMod vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_with_single_cte), context=sushi_context, @@ -272,11 +264,12 @@ def test_without_rows(sushi_context: Context, full_model_with_single_cte: SqlMod ) -def test_column_order(sushi_context: Context, full_model_without_ctes: SqlModel) -> None: +def test_column_order( + sushi_context: Context, full_model_without_ctes: SqlModel +) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -292,8 +285,7 @@ def test_column_order(sushi_context: Context, full_model_without_ctes: SqlModel) vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_without_ctes), context=sushi_context, @@ -303,8 +295,7 @@ def test_column_order(sushi_context: Context, full_model_without_ctes: SqlModel) def test_row_order(sushi_context: Context, full_model_without_ctes: SqlModel) -> None: # Input and output rows are in different orders - body = load_yaml( - """ + body = load_yaml(""" test_foo: model: sushi.foo inputs: @@ -326,8 +317,7 @@ def test_row_order(sushi_context: Context, full_model_without_ctes: SqlModel) -> vars: start: 2022-01-01 end: 2022-01-01 - """ - ) + """) # Model query without ORDER BY should pass unit test _check_successful_or_raise( @@ -373,8 +363,7 @@ def test_row_order(sushi_context: Context, full_model_without_ctes: SqlModel) -> _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_array_order: model: test inputs: @@ -388,11 +377,12 @@ def test_row_order(sushi_context: Context, full_model_without_ctes: SqlModel) -> - aggregated_duplicates: - c - b - """ - ), + """), test_name="test_array_order", model=_create_model(model_sql), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( """AssertionError: Data mismatch (exp: expected, act: actual)\n\n""" @@ -425,8 +415,7 @@ def test_row_order(sushi_context: Context, full_model_without_ctes: SqlModel) -> def test_partial_data(sushi_context: Context, waiter_names_input: str) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - f""" + body=load_yaml(f""" test_foo: model: sushi.foo inputs: @@ -447,8 +436,7 @@ def test_partial_data(sushi_context: Context, waiter_names_input: str) -> None: - id: 3 name: 'bob' str: nan - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( @@ -492,8 +480,7 @@ def test_partial_data(sushi_context: Context, waiter_names_input: str) -> None: def test_format_inline(sushi_context: Context, waiter_names_input: str) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - f""" + body=load_yaml(f""" test_foo: model: sushi.foo inputs: @@ -504,8 +491,7 @@ def test_format_inline(sushi_context: Context, waiter_names_input: str) -> None: name: alice - id: 2 name: 'bob' - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( @@ -555,15 +541,18 @@ def test_format_inline(sushi_context: Context, waiter_names_input: str) -> None: ], ) def test_format_path( - sushi_context: Context, tmp_path: Path, input_data: str, filename: str, file_data: str + sushi_context: Context, + tmp_path: Path, + input_data: str, + filename: str, + file_data: str, ) -> None: test_csv_file = tmp_path / filename test_csv_file.write_text(file_data) _check_successful_or_raise( _create_test( - body=load_yaml( - f""" + body=load_yaml(f""" test_foo: model: sushi.foo inputs: @@ -574,8 +563,7 @@ def test_format_path( name: alice - id: 2 name: 'bob' - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( @@ -596,8 +584,7 @@ def test_unsupported_format_failure( match="Unsupported data format 'xml' for 'sushi.waiter_names'", ): _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo description: XML format isn't supported to load data (fails intentionally) @@ -609,8 +596,7 @@ def test_unsupported_format_failure( query: - id: 1 value: null - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_without_ctes), context=sushi_context, @@ -621,8 +607,7 @@ def test_unsupported_format_failure( match="Unsupported data format 'xml' for 'sushi.waiter_names'", ): _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo description: XML without path doesn't raise error @@ -646,8 +631,7 @@ def test_unsupported_format_failure( name: alice - id: 2 name: 'bob' - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( @@ -662,8 +646,7 @@ def test_unsupported_format_failure( def test_partial_output_columns() -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -688,18 +671,20 @@ def test_partial_output_columns() -> None: b: 2 - a: 5 b: 6 - """ - ), + """), test_name="test_foo", - model=_create_model("WITH t AS (SELECT a, b, c, d FROM raw) SELECT a, b, c, d FROM t"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + model=_create_model( + "WITH t AS (SELECT a, b, c, d FROM raw) SELECT a, b, c, d FROM t" + ), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -727,11 +712,14 @@ def test_partial_output_columns() -> None: - a: 5 b: 6 c: 7 - """ - ), + """), test_name="test_foo", - model=_create_model("WITH t AS (SELECT a, b, c, d FROM raw) SELECT a, b, c, d FROM t"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + model=_create_model( + "WITH t AS (SELECT a, b, c, d FROM raw) SELECT a, b, c, d FROM t" + ), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) @@ -739,8 +727,7 @@ def test_partial_output_columns() -> None: def test_partial_data_column_order(sushi_context: Context) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -757,8 +744,7 @@ def test_partial_data_column_order(sushi_context: Context) -> None: - id: 9876 name: hello event_date: 2020-01-02 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( @@ -774,8 +760,7 @@ def test_partial_data_column_order(sushi_context: Context) -> None: # - output partial must be true _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -793,8 +778,7 @@ def test_partial_data_column_order(sushi_context: Context) -> None: - event_date: 2020-01-02 id: 1234 name: hello - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( @@ -810,8 +794,7 @@ def test_partial_data_column_order(sushi_context: Context) -> None: def test_partial_data_missing_schemas(sushi_context: Context) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -824,8 +807,7 @@ def test_partial_data_missing_schemas(sushi_context: Context) -> None: - a: 1 b: bla - b: baz - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(_create_model("SELECT * FROM unknown")), context=sushi_context, @@ -833,8 +815,7 @@ def test_partial_data_missing_schemas(sushi_context: Context) -> None: ) _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -853,8 +834,7 @@ def test_partial_data_missing_schemas(sushi_context: Context) -> None: date: 2023-02-10 month: 2023-02-01 null_date: - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( @@ -866,7 +846,9 @@ def test_partial_data_missing_schemas(sushi_context: Context) -> None: ) -def test_partially_inferred_schemas(sushi_context: Context, mocker: MockerFixture) -> None: +def test_partially_inferred_schemas( + sushi_context: Context, mocker: MockerFixture +) -> None: parent = _create_model( "SELECT a, b, s::ROW(d DATE) AS s FROM sushi.unknown", meta="MODEL (name sushi.parent, kind FULL, dialect trino)", @@ -881,8 +863,7 @@ def test_partially_inferred_schemas(sushi_context: Context, mocker: MockerFixtur ) child = t.cast(SqlModel, sushi_context.upsert_model(child)) - body = load_yaml( - """ + body = load_yaml(""" test_child: model: sushi.child inputs: @@ -896,8 +877,7 @@ def test_partially_inferred_schemas(sushi_context: Context, mocker: MockerFixtur - a: 1 b: bla d: 2020-01-01 - """ - ) + """) mocker.patch("sqlmesh.core.test.definition.random_id", return_value="jzngz56a") test = _create_test(body, "test_child", child, sushi_context) @@ -919,8 +899,7 @@ def test_partially_inferred_schemas(sushi_context: Context, mocker: MockerFixtur def test_uninferrable_schema() -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -929,11 +908,12 @@ def test_uninferrable_schema() -> None: outputs: query: - value: null - """ - ), + """), test_name="test_foo", model=_create_model("SELECT value FROM raw"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( """Failed to infer the data type of column 'value' for '"raw"'. This issue can """ @@ -944,11 +924,12 @@ def test_uninferrable_schema() -> None: ) -def test_missing_column_failure(sushi_context: Context, full_model_without_ctes: SqlModel) -> None: +def test_missing_column_failure( + sushi_context: Context, full_model_without_ctes: SqlModel +) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo description: sushi.foo's output has a missing column (fails intentionally) @@ -961,8 +942,7 @@ def test_missing_column_failure(sushi_context: Context, full_model_without_ctes: query: - id: 1 value: null - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_without_ctes), context=sushi_context, @@ -979,8 +959,7 @@ def test_missing_column_failure(sushi_context: Context, full_model_without_ctes: def test_row_difference_failure() -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -990,11 +969,12 @@ def test_row_difference_failure() -> None: query: - value: 1 - value: 2 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT value FROM raw"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( "AssertionError: Data mismatch (rows are different)\n\n" @@ -1005,8 +985,7 @@ def test_row_difference_failure() -> None: ) _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1015,13 +994,14 @@ def test_row_difference_failure() -> None: outputs: query: - value: 1 - """ - ), + """), test_name="test_foo", model=_create_model( "SELECT value FROM raw UNION ALL SELECT value + 1 AS value FROM raw" ), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( "AssertionError: Data mismatch (rows are different)\n\n" @@ -1032,8 +1012,7 @@ def test_row_difference_failure() -> None: ) _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1044,13 +1023,14 @@ def test_row_difference_failure() -> None: - value: 1 - value: 3 - value: 4 - """ - ), + """), test_name="test_foo", model=_create_model( "SELECT value FROM raw UNION ALL SELECT value + 1 AS value FROM raw" ), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( "AssertionError: Data mismatch (rows are different)\n\n" @@ -1069,8 +1049,7 @@ def test_index_preservation_with_later_rows() -> None: # Test comparison with differences in later rows _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1093,11 +1072,12 @@ def test_index_preservation_with_later_rows() -> None: value: 999 - id: 4 value: 888 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT id, value FROM raw"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( "AssertionError: Data mismatch (exp: expected, act: actual)\n\n" @@ -1111,8 +1091,7 @@ def test_index_preservation_with_later_rows() -> None: # Test with null values in later rows _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1139,11 +1118,12 @@ def test_index_preservation_with_later_rows() -> None: value: null - id: 5 value: 500 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT id, value FROM raw"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( "AssertionError: Data mismatch (exp: expected, act: actual)\n\n" @@ -1159,8 +1139,7 @@ def test_index_preservation_with_later_rows() -> None: def test_unknown_column_error() -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1170,11 +1149,12 @@ def test_unknown_column_error() -> None: outputs: query: - foo: 1 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT id, value FROM raw"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run(), expected_msg=( "sqlmesh.utils.errors.TestError: Failed to run test:\n" @@ -1186,10 +1166,11 @@ def test_unknown_column_error() -> None: def test_invalid_outputs_error() -> None: - with pytest.raises(TestError, match="Incomplete test, outputs must contain 'query' or 'ctes'"): + with pytest.raises( + TestError, match="Incomplete test, outputs must contain 'query' or 'ctes'" + ): _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1198,31 +1179,31 @@ def test_invalid_outputs_error() -> None: outputs: rows: - id: 1 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT id FROM raw"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ) def test_empty_rows(sushi_context: Context) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: sushi.items: [] outputs: query: [] - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model( _create_model( - "SELECT id FROM sushi.items", default_catalog=sushi_context.default_catalog + "SELECT id FROM sushi.items", + default_catalog=sushi_context.default_catalog, ) ), context=sushi_context, @@ -1231,8 +1212,7 @@ def test_empty_rows(sushi_context: Context) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_a: model: a inputs: @@ -1242,13 +1222,14 @@ def test_empty_rows(sushi_context: Context) -> None: rows: [] outputs: query: [] - """ - ), + """), test_name="test_a", model=sushi_context.upsert_model( _create_model("SELECT x FROM b", default_catalog="memory") ), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) @@ -1256,8 +1237,7 @@ def test_empty_rows(sushi_context: Context) -> None: @pytest.mark.parametrize("full_model_without_ctes", ["snowflake"], indirect=True) def test_normalization(full_model_without_ctes: SqlModel) -> None: normalized_body = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1275,18 +1255,22 @@ def test_normalization(full_model_without_ctes: SqlModel) -> None: vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=full_model_without_ctes, - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="snowflake"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="snowflake")) + ), ).body expected_body = { "model": '"MEMORY"."SUSHI"."FOO"', "inputs": {'"RAW"': {"rows": [{"ID": 1}]}}, "outputs": { - "ctes": {'"SOURCE"': {"rows": [{"ID": 1}]}, '"RENAMED"': {"rows": [{"FID": 1}]}}, + "ctes": { + '"SOURCE"': {"rows": [{"ID": 1}]}, + '"RENAMED"': {"rows": [{"FID": 1}]}, + }, "query": {"rows": [{"FID": 1}]}, }, "vars": {"start": datetime.date(2022, 1, 1), "end": datetime.date(2022, 1, 1)}, @@ -1298,8 +1282,7 @@ def test_normalization(full_model_without_ctes: SqlModel) -> None: def test_source_func() -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: xyz outputs: @@ -1307,16 +1290,15 @@ def test_source_func() -> None: - month: 2023-01-01 - month: 2023-02-01 - month: 2023-03-01 - """ - ), + """), test_name="test_foo", - model=_create_model( - """ + model=_create_model(""" SELECT range::DATE AS month FROM RANGE(DATE '2023-01-01', DATE '2023-04-01', INTERVAL 1 MONTH) AS r - """ + """), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) ), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), ).run() ) @@ -1324,8 +1306,7 @@ def test_source_func() -> None: def test_nested_data_types(sushi_context: Context) -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1349,13 +1330,14 @@ def test_nested_data_types(sushi_context: Context) -> None: array2: [{'k': 'hello', 'v': {'v_str': 'there', 'v_int': 10, 'v_int_arr': [1, 2]}}] struct: {'x': [1, 2, 3], 'y': 'foo', 'z': 1, 'w': {'a': 5}} - array1: [2, 3] - """ - ), + """), test_name="test_foo", model=_create_model( "SELECT array1, array2, struct FROM sushi.raw", default_catalog="memory" ), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) @@ -1363,8 +1345,7 @@ def test_nested_data_types(sushi_context: Context) -> None: def test_freeze_time(mocker: MockerFixture) -> None: mocker.patch("sqlmesh.core.test.definition.random_id", return_value="jzngz56a") test = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: xyz outputs: @@ -1374,13 +1355,14 @@ def test_freeze_time(mocker: MockerFixture) -> None: cur_timestamp: "2023-01-01 12:05:03" vars: execution_time: "2023-01-01 12:05:03+00:00" - """ - ), + """), test_name="test_foo", model=_create_model( "SELECT CURRENT_DATE AS cur_date, CURRENT_TIME AS cur_time, CURRENT_TIMESTAMP AS cur_timestamp" ), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ) spy_execute = mocker.spy(test.engine_adapter, "_execute") @@ -1396,13 +1378,14 @@ def test_freeze_time(mocker: MockerFixture) -> None: '''CAST('2023-01-01 12:05:03+00:00' AS TIMESTAMP) AS "cur_timestamp"''', False, ), - call('DROP SCHEMA IF EXISTS "memory"."sqlmesh_test_jzngz56a" CASCADE', False), + call( + 'DROP SCHEMA IF EXISTS "memory"."sqlmesh_test_jzngz56a" CASCADE', False + ), ] ) test = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: xyz outputs: @@ -1410,11 +1393,12 @@ def test_freeze_time(mocker: MockerFixture) -> None: - cur_timestamp: "2023-01-01 12:05:03+00:00" vars: execution_time: "2023-01-01 12:05:03+00:00" - """ - ), + """), test_name="test_foo", model=_create_model("SELECT CURRENT_TIMESTAMP AS cur_timestamp"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="bigquery"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="bigquery")) + ), ) spy_execute = mocker.spy(test.engine_adapter, "_execute") @@ -1440,8 +1424,7 @@ def execute(context, start, end, execution_time, **kwargs): _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_py_model: model: py_model outputs: @@ -1450,11 +1433,14 @@ def execute(context, start, end, execution_time, **kwargs): ts2: "2023-01-01 10:05:03" vars: execution_time: "2023-01-01 12:05:03+02:00" - """ - ), + """), test_name="test_py_model", - model=model.get_registry()["py_model"].model(module_path=Path("."), path=Path(".")), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + model=model.get_registry()["py_model"].model( + module_path=Path("."), path=Path(".") + ), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) @@ -1467,11 +1453,12 @@ def test_successes(sushi_context: Context) -> None: assert "test_customer_revenue_by_day" in successful_tests -def test_create_external_model_fixture(sushi_context: Context, mocker: MockerFixture) -> None: +def test_create_external_model_fixture( + sushi_context: Context, mocker: MockerFixture +) -> None: mocker.patch("sqlmesh.core.test.definition.random_id", return_value="jzngz56a") test = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -1480,8 +1467,7 @@ def test_create_external_model_fixture(sushi_context: Context, mocker: MockerFix outputs: query: - x: 1 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT x FROM c.db.external"), context=sushi_context, @@ -1497,22 +1483,25 @@ def test_create_external_model_fixture(sushi_context: Context, mocker: MockerFix def test_runtime_stage() -> None: @macro() def test_macro(evaluator: MacroEvaluator) -> t.List[bool]: - return [evaluator.runtime_stage == "testing", evaluator.runtime_stage == "loading"] + return [ + evaluator.runtime_stage == "testing", + evaluator.runtime_stage == "loading", + ] _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: foo outputs: query: - c: [true, false] - """ - ), + """), test_name="test_foo", model=_create_model("SELECT [@test_macro()] AS c"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) @@ -1524,8 +1513,12 @@ def test_gateway(copy_to_temp_path: t.Callable, mocker: MockerFixture) -> None: config = Config( gateways={ - "main": GatewayConfig(connection=DuckDBConnectionConfig(database=db_db_path)), - "test": GatewayConfig(test_connection=DuckDBConnectionConfig(database=test_db_path)), + "main": GatewayConfig( + connection=DuckDBConnectionConfig(database=db_db_path) + ), + "test": GatewayConfig( + test_connection=DuckDBConnectionConfig(database=test_db_path) + ), }, model_defaults=ModelDefaultsConfig(dialect="duckdb"), ) @@ -1578,8 +1571,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> context = Context(paths=tmp_path, config=config) _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_example_full_model_alt: model: sqlmesh_example.full_model inputs: @@ -1605,8 +1597,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> num_orders: 1 - item_id: 4 num_orders: 0 - """ - ), + """), test_name="test_example_full_model_alt", model=context.get_model("sqlmesh_example.full_model"), context=context, @@ -1615,8 +1606,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_example_full_model_partial: model: sqlmesh_example.full_model inputs: @@ -1631,8 +1621,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> rows: - item_id: null num_orders: 2 - """ - ), + """), test_name="test_example_full_model_partial", model=context.get_model("sqlmesh_example.full_model"), context=context, @@ -1641,8 +1630,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_example_full_model_partial: model: sqlmesh_example.full_model inputs: @@ -1658,8 +1646,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> query: partial: true query: "SELECT 2 AS num_orders UNION ALL SELECT 1 AS num_orders" - """ - ), + """), test_name="test_example_full_model_partial", model=context.get_model("sqlmesh_example.full_model"), context=context, @@ -1668,8 +1655,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> mocker.patch("sqlmesh.core.test.definition.random_id", return_value="jzngz56a") test = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: xyz inputs: @@ -1678,11 +1664,12 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> outputs: query: - struct_value: {'x': 1, 'n': {'y': 2}} - """ - ), + """), test_name="test_foo", model=_create_model("SELECT struct_value FROM foo"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ) spy_execute = mocker.spy(test.engine_adapter, "_execute") _check_successful_or_raise(test.run()) @@ -1698,8 +1685,7 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> match="Invalid test, cannot set both 'query' and 'rows' for 'foo'", ): _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: xyz inputs: @@ -1710,11 +1696,12 @@ def test_generate_input_data_using_sql(mocker: MockerFixture, tmp_path: Path) -> outputs: query: - struct_value: {'x': 1, 'n': {'y': 2}} - """ - ), + """), test_name="test_foo", model=_create_model("SELECT struct_value FROM foo"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ) @@ -1740,15 +1727,13 @@ def execute(context, start, end, execution_time, **kwargs): _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_pyspark_model: model: pyspark_model outputs: query: - col: 1 - """ - ), + """), test_name="test_pyspark_model", model=model.get_registry()["pyspark_model"].model( module_path=Path("."), path=Path(".") @@ -1821,7 +1806,9 @@ def init_context_and_validate_results(config: Config, **kwargs): # Case 2: Test gateway variables config = Config( gateways={ - "main": GatewayConfig(connection=DuckDBConnectionConfig(), variables=variables), + "main": GatewayConfig( + connection=DuckDBConnectionConfig(), variables=variables + ), }, model_defaults=ModelDefaultsConfig(dialect="duckdb"), ) @@ -1830,7 +1817,9 @@ def init_context_and_validate_results(config: Config, **kwargs): # Case 3: Test gateway variables overriding root variables config = Config( gateways={ - "main": GatewayConfig(connection=DuckDBConnectionConfig(), variables=variables), + "main": GatewayConfig( + connection=DuckDBConnectionConfig(), variables=variables + ), }, model_defaults=ModelDefaultsConfig(dialect="duckdb"), variables=incorrect_variables, @@ -1843,7 +1832,9 @@ def init_context_and_validate_results(config: Config, **kwargs): "main": GatewayConfig( connection=DuckDBConnectionConfig(), variables=incorrect_variables ), - "secondary": GatewayConfig(connection=DuckDBConnectionConfig(), variables=variables), + "secondary": GatewayConfig( + connection=DuckDBConnectionConfig(), variables=variables + ), }, model_defaults=ModelDefaultsConfig(dialect="duckdb"), ) @@ -1857,7 +1848,9 @@ def init_context_and_validate_results(config: Config, **kwargs): "main": GatewayConfig( connection=DuckDBConnectionConfig(), variables=incorrect_variables ), - "secon\tdary": GatewayConfig(connection=DuckDBConnectionConfig(), variables=variables), + "secon\tdary": GatewayConfig( + connection=DuckDBConnectionConfig(), variables=variables + ), }, model_defaults=ModelDefaultsConfig(dialect="duckdb"), ) @@ -1868,19 +1861,19 @@ def init_context_and_validate_results(config: Config, **kwargs): def test_custom_testing_schema(mocker: MockerFixture) -> None: test = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: xyz schema: my_schema outputs: query: - a: 1 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT 1 AS a"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ) spy_execute = mocker.spy(test.engine_adapter, "_execute") @@ -1897,19 +1890,19 @@ def test_custom_testing_schema(mocker: MockerFixture) -> None: def test_pretty_query(mocker: MockerFixture) -> None: test = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: xyz schema: my_schema outputs: query: - a: 1 - """ - ), + """), test_name="test_foo", model=_create_model("SELECT 1 AS a"), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ) test.engine_adapter._pretty_sql = True spy_execute = mocker.spy(test.engine_adapter, "_execute") @@ -1972,8 +1965,7 @@ def test_complicated_recursive_cte() -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_recursive_ctes: model: test inputs: @@ -2008,24 +2000,23 @@ def test_complicated_recursive_cte() -> None: rows: - aggregated_duplicates: [a, b, c, d, e, f, g] - aggregated_duplicates: [x, y] - """ - ), + """), test_name="test_recursive_ctes", model=_create_model(model_sql), - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) def test_unknown_model_warns(mocker: MockerFixture) -> None: - body = load_yaml( - """ + body = load_yaml(""" model: unknown outputs: query: - c: 1 - """ - ) + """) with patch.object(get_console(), "log_warning") as mock_logger: ModelTest.create_test( @@ -2096,7 +2087,10 @@ def test_test_generation(tmp_path: Path) -> None: assert len(test) == 1 assert "test_full_model" in test assert "vars" in test["test_full_model"] - assert test["test_full_model"]["vars"] == {"start": "2020-01-01", "end": "2024-01-01"} + assert test["test_full_model"]["vars"] == { + "start": "2020-01-01", + "end": "2024-01-01", + } assert "ctes" in test["test_full_model"]["outputs"] assert "cte" in test["test_full_model"]["outputs"]["ctes"] @@ -2168,7 +2162,9 @@ def test_test_generation(tmp_path: Path) -> None: ), ], ) -def test_test_generation_with_data_structures(tmp_path: Path, column: str, expected: str) -> None: +def test_test_generation_with_data_structures( + tmp_path: Path, column: str, expected: str +) -> None: def create_test(context: Context, query: str): context.create_test( "sqlmesh_example.foo", @@ -2188,9 +2184,13 @@ def create_test(context: Context, query: str): "MODEL (name sqlmesh_example.foo); SELECT col FROM sqlmesh_example.bar;" ) bar_sql_file = tmp_path / "models" / "bar.sql" - bar_sql_file.write_text("MODEL (name sqlmesh_example.bar); SELECT col FROM external_table;") + bar_sql_file.write_text( + "MODEL (name sqlmesh_example.bar); SELECT col FROM external_table;" + ) - test = create_test(Context(paths=tmp_path, config=config), f"SELECT {column} AS col") + test = create_test( + Context(paths=tmp_path, config=config), f"SELECT {column} AS col" + ) assert test["test_foo"]["inputs"] == {'"memory"."sqlmesh_example"."bar"': expected} assert test["test_foo"]["outputs"] == {"query": expected} @@ -2207,14 +2207,18 @@ def test_test_generation_with_timestamp(tmp_path: Path) -> None: "MODEL (name sqlmesh_example.foo); SELECT ts_col FROM sqlmesh_example.bar;" ) bar_sql_file = tmp_path / "models" / "bar.sql" - bar_sql_file.write_text("MODEL (name sqlmesh_example.bar); SELECT ts_col FROM external_table;") + bar_sql_file.write_text( + "MODEL (name sqlmesh_example.bar); SELECT ts_col FROM external_table;" + ) context = Context(paths=tmp_path, config=config) input_queries = { "sqlmesh_example.bar": "SELECT TIMESTAMP '2024-09-20 11:30:00.123456789' AS ts_col" } - context.create_test("sqlmesh_example.foo", input_queries=input_queries, overwrite=True) + context.create_test( + "sqlmesh_example.foo", input_queries=input_queries, overwrite=True + ) test = load_yaml(context.path / c.TESTS / "test_foo.yaml") @@ -2246,7 +2250,9 @@ def test_test_generation_with_recursive_ctes(tmp_path: Path) -> None: context = Context(paths=tmp_path, config=config) context.plan(auto_apply=True) - context.create_test("sqlmesh_example.foo", input_queries={}, overwrite=True, include_ctes=True) + context.create_test( + "sqlmesh_example.foo", input_queries={}, overwrite=True, include_ctes=True + ) test = load_yaml(context.path / c.TESTS / "test_foo.yaml") assert len(test) == 1 @@ -2262,7 +2268,9 @@ def test_test_generation_with_recursive_ctes(tmp_path: Path) -> None: _check_successful_or_raise(context.test()) -def test_test_with_gateway_specific_model(tmp_path: Path, mocker: MockerFixture) -> None: +def test_test_with_gateway_specific_model( + tmp_path: Path, mocker: MockerFixture +) -> None: init_example_project(tmp_path, engine_type="duckdb") config = Config( @@ -2289,12 +2297,15 @@ def test_test_with_gateway_specific_model(tmp_path: Path, mocker: MockerFixture) assert context.engine_adapter == context.engine_adapters["main"] with pytest.raises( - SQLMeshError, match=r"Gateway 'wrong' not found in the available engine adapters." + SQLMeshError, + match=r"Gateway 'wrong' not found in the available engine adapters.", ): context._get_engine_adapter("wrong") # Create test should use the gateway specific engine adapter - context.create_test("sqlmesh_example.gw_model", input_queries=input_queries, overwrite=True) + context.create_test( + "sqlmesh_example.gw_model", input_queries=input_queries, overwrite=True + ) assert context._get_engine_adapter("second") == context.engine_adapters["second"] assert len(context.engine_adapters) == 2 @@ -2316,8 +2327,7 @@ def test_test_with_resolve_template_macro(tmp_path: Path): models_dir = tmp_path / "models" models_dir.mkdir() - (models_dir / "foo.sql").write_text( - """ + (models_dir / "foo.sql").write_text(""" MODEL ( name test.foo, kind full, @@ -2328,13 +2338,11 @@ def test_test_with_resolve_template_macro(tmp_path: Path): SELECT t.a + 1 as a FROM @resolve_template('@{schema_name}.dev_@{table_name}', mode := 'table') as t - """ - ) + """) tests_dir = tmp_path / "tests" tests_dir.mkdir() - (tests_dir / "test_foo.yaml").write_text( - """ + (tests_dir / "test_foo.yaml").write_text(""" test_resolve_template_macro: model: test.foo inputs: @@ -2343,8 +2351,7 @@ def test_test_with_resolve_template_macro(tmp_path: Path): outputs: query: - a: 2 - """ - ) + """) context = Context(paths=tmp_path, config=config) _check_successful_or_raise(context.test()) @@ -2364,8 +2371,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: original_test_file = tmp_path / "tests" / "test_full_model.yaml" new_test_file = tmp_path / "tests" / "test_full_model_error.yaml" - new_test_file.write_text( - """ + new_test_file.write_text(""" test_example_full_model: model: sqlmesh_example.full_model description: This is a test @@ -2385,8 +2391,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: num_orders: 2 - item_id: 4 num_orders: 3 - """ - ) + """) config = Config( default_connection=DuckDBConnectionConfig(), @@ -2403,8 +2408,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: # Order may change due to concurrent execution assert "F." in output or ".F" in output - assert ( - f"""This is a test + assert f"""This is a test ---------------------------------------------------------------------- Data mismatch ┏━━━━━┳━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━┓ @@ -2414,9 +2418,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: │ 1 │ 4.0 │ 2.0 │ 3.0 │ 1.0 │ └─────┴─────────────────┴─────────────────┴─────────────────┴──────────────────┘ -----------------------------------------------------------------------""" - in output - ) +----------------------------------------------------------------------""" in output assert "Ran 2 tests" in output assert "Failed tests (1):" in output @@ -2427,8 +2429,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: output = captured_output.stdout - assert ( - f"""This is a test + assert f"""This is a test ---------------------------------------------------------------------- Column 'item_id' mismatch ┏━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━┓ @@ -2444,13 +2445,13 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: │ 1 │ 3.0 │ 1.0 │ └─────────────┴────────────────────────┴───────────────────┘ -----------------------------------------------------------------------""" - in output - ) +----------------------------------------------------------------------""" in output # Case 3: Assert that concurrent execution is working properly for i in range(50): - copy_test_file(original_test_file, tmp_path / "tests" / f"test_success_{i}.yaml", i) + copy_test_file( + original_test_file, tmp_path / "tests" / f"test_success_{i}.yaml", i + ) copy_test_file(new_test_file, tmp_path / "tests" / f"test_failure_{i}.yaml", i) # Re-initialize context to pick up the new test files @@ -2467,9 +2468,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: # Case 4: Test that wide tables are split into even chunks for default verbosity rmtree(tmp_path / "tests") - wide_model_query = ( - "SELECT 1 AS col_1, 2 AS col_2, 3 AS col_3, 4 AS col_4, 5 AS col_5, 6 AS col_6, 7 AS col_7" - ) + wide_model_query = "SELECT 1 AS col_1, 2 AS col_2, 3 AS col_3, 4 AS col_4, 5 AS col_5, 6 AS col_6, 7 AS col_7" wide_model = _create_model( meta="MODEL(name test.test_wide_model)", @@ -2531,8 +2530,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: tests_dir.mkdir() null_test_file = tmp_path / "tests" / "test_null_in_third_row.yaml" - null_test_file.write_text( - """ + null_test_file.write_text(""" test_null_third_row: model: sqlmesh_example.full_model description: Test null value in third row @@ -2556,8 +2554,7 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: num_orders: 1 - item_id: 3 num_orders: null - """ - ) + """) # Re-initialize context to pick up the modified test file context = Context(paths=tmp_path, config=config) @@ -2568,15 +2565,12 @@ def copy_test_file(test_file: Path, new_test_file: Path, index: int) -> None: output = captured_output.stdout # Check for null value difference in the 3rd row (index 2) - assert ( - """ + assert """ ┏━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━┓ ┃ Row ┃ num_orders: Expected ┃ num_orders: Actual ┃ ┡━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━┩ │ 2 │ nan │ 1.0 │ -└──────┴───────────────────────────┴───────────────────────┘""" - in output - ) +└──────┴───────────────────────────┴───────────────────────┘""" in output @use_terminal_console @@ -2584,8 +2578,7 @@ def test_test_output_with_invalid_model_name(tmp_path: Path) -> None: init_example_project(tmp_path, engine_type="duckdb") wrong_test_file = tmp_path / "tests" / "test_incorrect_model_name.yaml" - wrong_test_file.write_text( - """ + wrong_test_file.write_text(""" test_example_full_model: model: invalid_model description: This is an invalid test @@ -2605,8 +2598,7 @@ def test_test_output_with_invalid_model_name(tmp_path: Path) -> None: num_orders: 2 - item_id: 2 num_orders: 2 - """ - ) + """) config = Config( default_connection=DuckDBConnectionConfig(), @@ -2630,8 +2622,7 @@ def test_number_of_tests_found(tmp_path: Path) -> None: # Example project contains 1 test and we add a new file with 2 tests test_file = tmp_path / "tests" / "test_new.yaml" - test_file.write_text( - """ + test_file.write_text(""" test_example_full_model1: model: sqlmesh_example.full_model inputs: @@ -2669,8 +2660,7 @@ def test_number_of_tests_found(tmp_path: Path) -> None: num_orders: 2 - item_id: 2 num_orders: 1 - """ - ) + """) context = Context(paths=tmp_path) @@ -2695,8 +2685,7 @@ def test_freeze_time_concurrent(tmp_path: Path) -> None: macros_dir.mkdir() macro_file = macros_dir / "test_datetime_now.py" - macro_file.write_text( - """ + macro_file.write_text(""" from sqlglot import exp import datetime from sqlmesh.core.macros import macro @@ -2708,24 +2697,20 @@ def test_datetime_now(evaluator): @macro() def test_sqlglot_expr(evaluator): return exp.CurrentDate().sql(evaluator.dialect) - """ - ) + """) models_dir = tmp_path / "models" models_dir.mkdir() sql_model1 = models_dir / "sql_model1.sql" - sql_model1.write_text( - """ + sql_model1.write_text(""" MODEL(NAME sql_model1); SELECT @test_datetime_now() AS col_exec_ds_time, @test_sqlglot_expr() AS col_current_date; - """ - ) + """) for model_name in ["sql_model1", "sql_model2", "py_model"]: for i in range(5): test_2019 = tmp_path / "tests" / f"test_2019_{model_name}_{i}.yaml" - test_2019.write_text( - f""" + test_2019.write_text(f""" test_2019_{model_name}_{i}: model: {model_name} vars: @@ -2735,12 +2720,10 @@ def test_sqlglot_expr(evaluator): rows: - col_exec_ds_time: '2019-12-01' col_current_date: '2019-12-01' - """ - ) + """) test_2025 = tmp_path / "tests" / f"test_2025_{model_name}_{i}.yaml" - test_2025.write_text( - f""" + test_2025.write_text(f""" test_2025_{model_name}_{i}: model: {model_name} vars: @@ -2750,17 +2733,21 @@ def test_sqlglot_expr(evaluator): rows: - col_exec_ds_time: '2025-12-01' col_current_date: '2025-12-01' - """ - ) + """) ctx = Context( paths=tmp_path, - config=Config(default_test_connection=DuckDBConnectionConfig(concurrent_tasks=8)), + config=Config( + default_test_connection=DuckDBConnectionConfig(concurrent_tasks=8) + ), ) @model( "py_model", - columns={"col_exec_ds_time": "timestamp_ntz", "col_current_date": "timestamp_ntz"}, + columns={ + "col_exec_ds_time": "timestamp_ntz", + "col_current_date": "timestamp_ntz", + }, ) def execute(context, start, end, execution_time, **kwargs): datetime_now_utc = datetime.datetime.now(tz=datetime.timezone.utc) @@ -2772,7 +2759,9 @@ def execute(context, start, end, execution_time, **kwargs): [{"col_exec_ds_time": datetime_now_utc, "col_current_date": current_date}] ) - python_model = model.get_registry()["py_model"].model(module_path=Path("."), path=Path(".")) + python_model = model.get_registry()["py_model"].model( + module_path=Path("."), path=Path(".") + ) ctx.upsert_model(python_model) ctx.upsert_model( @@ -2803,7 +2792,9 @@ def upstream_table_python(context, **kwargs): path=Path("."), ) - context = ExecutionContext(sushi_context.engine_adapter, sushi_context.snapshots, None, None) + context = ExecutionContext( + sushi_context.engine_adapter, sushi_context.snapshots, None, None + ) df = list(python_model.render(context=context))[0] # Verify the actual model output matches the expected actual external table's values @@ -2843,8 +2834,7 @@ def test_model_test_text_result_reporting_no_traceback( sushi_context: Context, full_model_with_two_ctes: SqlModel, is_error: bool ) -> None: test = _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: sushi.foo inputs: @@ -2859,8 +2849,7 @@ def test_model_test_text_result_reporting_no_traceback( vars: start: 2022-01-01 end: 2022-01-01 - """ - ), + """), test_name="test_foo", model=sushi_context.upsert_model(full_model_with_two_ctes), context=sushi_context, @@ -2908,8 +2897,7 @@ def test_timestamp_normalization() -> None: _check_successful_or_raise( _create_test( - body=load_yaml( - """ + body=load_yaml(""" test_foo: model: temp_agg_model_with_timestamp inputs: @@ -2922,17 +2910,20 @@ def test_timestamp_normalization() -> None: rows: - id: id1 agg_timestamp_col: ["2024-01-02T15:00:00.000000"] - """ - ), + """), test_name="test_foo", model=model, - context=Context(config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb"))), + context=Context( + config=Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) + ), ).run() ) @use_terminal_console -def test_disable_test_logging_if_no_tests_found(mocker: MockerFixture, tmp_path: Path) -> None: +def test_disable_test_logging_if_no_tests_found( + mocker: MockerFixture, tmp_path: Path +) -> None: init_example_project(tmp_path, engine_type="duckdb") config = Config( @@ -2963,24 +2954,26 @@ def test_test_generation_with_timestamp_nat(tmp_path: Path) -> None: "MODEL (name sqlmesh_example.foo); SELECT ts_col FROM sqlmesh_example.bar;" ) bar_sql_file = tmp_path / "models" / "bar.sql" - bar_sql_file.write_text("MODEL (name sqlmesh_example.bar); SELECT ts_col FROM external_table;") + bar_sql_file.write_text( + "MODEL (name sqlmesh_example.bar); SELECT ts_col FROM external_table;" + ) context = Context(paths=tmp_path, config=config) # This simulates the scenario where upstream models have NULL timestamp values - input_queries = { - "sqlmesh_example.bar": """ + input_queries = {"sqlmesh_example.bar": """ SELECT ts_col FROM ( VALUES (TIMESTAMP '2024-09-20 11:30:00.123456789'), (CAST(NULL AS TIMESTAMP)), (TIMESTAMP '2024-09-21 15:45:00.987654321') ) AS t(ts_col) - """ - } + """} # This should not raise an exception even with NULL timestamp values - context.create_test("sqlmesh_example.foo", input_queries=input_queries, overwrite=True) + context.create_test( + "sqlmesh_example.foo", input_queries=input_queries, overwrite=True + ) test = load_yaml(context.path / c.TESTS / "test_foo.yaml") assert len(test) == 1 @@ -3007,9 +3000,13 @@ def test_test_generation_with_timestamp_nat(tmp_path: Path) -> None: # Verify that the output matches the input (since the model just selects from bar) query_output = outputs["query"] assert len(query_output) == 3 - assert query_output[0]["ts_col"] == datetime.datetime(2024, 9, 20, 11, 30, 0, 123456) + assert query_output[0]["ts_col"] == datetime.datetime( + 2024, 9, 20, 11, 30, 0, 123456 + ) assert query_output[1]["ts_col"] is None - assert query_output[2]["ts_col"] == datetime.datetime(2024, 9, 21, 15, 45, 0, 987654) + assert query_output[2]["ts_col"] == datetime.datetime( + 2024, 9, 21, 15, 45, 0, 987654 + ) def test_parameterized_name_sql_model() -> None: @@ -3043,7 +3040,8 @@ def test_parameterized_name_sql_model() -> None: model=model, context=Context( config=Config( - model_defaults=ModelDefaultsConfig(dialect="snowflake"), variables=variables + model_defaults=ModelDefaultsConfig(dialect="snowflake"), + variables=variables, ) ), ) @@ -3092,7 +3090,8 @@ def execute( model=python_model, context=Context( config=Config( - model_defaults=ModelDefaultsConfig(dialect="snowflake"), variables=variables + model_defaults=ModelDefaultsConfig(dialect="snowflake"), + variables=variables, ) ), ) @@ -3142,7 +3141,8 @@ def test_parameterized_name_self_referential_model(): model=model, context=Context( config=Config( - model_defaults=ModelDefaultsConfig(dialect="snowflake"), variables=variables + model_defaults=ModelDefaultsConfig(dialect="snowflake"), + variables=variables, ) ), ) @@ -3151,7 +3151,9 @@ def test_parameterized_name_self_referential_model(): test1_model_query = test1._render_model_query().sql(dialect="snowflake") assert '"GOLD"."SUSHI"."FOO"' not in test1_model_query assert ( - test1._test_fixture_table('"GOLD"."SUSHI"."FOO"').sql(dialect="snowflake", identify=True) + test1._test_fixture_table('"GOLD"."SUSHI"."FOO"').sql( + dialect="snowflake", identify=True + ) in test1_model_query ) @@ -3174,7 +3176,8 @@ def test_parameterized_name_self_referential_model(): model=model, context=Context( config=Config( - model_defaults=ModelDefaultsConfig(dialect="snowflake"), variables=variables + model_defaults=ModelDefaultsConfig(dialect="snowflake"), + variables=variables, ) ), ) @@ -3183,7 +3186,9 @@ def test_parameterized_name_self_referential_model(): test2_model_query = test2._render_model_query().sql(dialect="snowflake") assert '"GOLD"."SUSHI"."FOO"' not in test2_model_query assert ( - test2._test_fixture_table('"GOLD"."SUSHI"."FOO"').sql(dialect="snowflake", identify=True) + test2._test_fixture_table('"GOLD"."SUSHI"."FOO"').sql( + dialect="snowflake", identify=True + ) in test2_model_query ) @@ -3206,9 +3211,13 @@ def execute( context: ExecutionContext, **kwargs: t.Any, ) -> pd.DataFrame: - current_table = context.resolve_table(f"{context.var('table_catalog')}.sushi.foo") + current_table = context.resolve_table( + f"{context.var('table_catalog')}.sushi.foo" + ) current_df = context.fetchdf(f"select id from {current_table}") - upstream_table = context.resolve_table(f"{context.var('table_catalog')}.sushi.bar") + upstream_table = context.resolve_table( + f"{context.var('table_catalog')}.sushi.bar" + ) upstream_df = context.fetchdf(f"select id from {upstream_table}") return pd.DataFrame([{"ID": upstream_df["ID"].sum() + current_df["ID"].sum()}]) @@ -3237,7 +3246,9 @@ def execute( assert model_bar.fqn == '"GOLD"."SUSHI"."BAR"' ctx = Context( - config=Config(model_defaults=ModelDefaultsConfig(dialect="snowflake"), variables=variables) + config=Config( + model_defaults=ModelDefaultsConfig(dialect="snowflake"), variables=variables + ) ) ctx.upsert_model(model_foo) ctx.upsert_model(model_bar) @@ -3283,8 +3294,7 @@ def execute( def test_python_model_test_variables_override(tmp_path: Path) -> None: py_model = tmp_path / "models" / "test_var_model.py" py_model.parent.mkdir(parents=True, exist_ok=True) - py_model.write_text( - """ + py_model.write_text(""" import pandas as pd # noqa: TID253 from sqlmesh import model, ExecutionContext import typing as t @@ -3301,8 +3311,7 @@ def execute(context: ExecutionContext, **kwargs: t.Any) -> pd.DataFrame: "id": 1 if my_flag else 2, "flag_value": my_flag, "var_value": other_var, - }])""" - ) + }])""") config = Config( model_defaults=ModelDefaultsConfig(dialect="duckdb"), @@ -3383,8 +3392,7 @@ def execute(context: ExecutionContext, **kwargs: t.Any) -> pd.DataFrame: def test_python_model_sorting(tmp_path: Path) -> None: py_model = tmp_path / "models" / "test_sort_model.py" py_model.parent.mkdir(parents=True, exist_ok=True) - py_model.write_text( - """ + py_model.write_text(""" import pandas as pd # noqa: TID253 from sqlmesh import model, ExecutionContext import typing as t @@ -3400,8 +3408,7 @@ def execute(context: ExecutionContext, **kwargs: t.Any) -> pd.DataFrame: {"id": 3, "value": "c"}, {"id": 1, "value": "a"}, {"id": 2, "value": "b"}, - ])""" - ) + ])""") config = Config(model_defaults=ModelDefaultsConfig(dialect="duckdb")) context = Context(config=config, paths=tmp_path) @@ -3434,8 +3441,7 @@ def execute(context: ExecutionContext, **kwargs: t.Any) -> pd.DataFrame: def test_cte_failure(tmp_path: Path) -> None: models_dir = tmp_path / "models" models_dir.mkdir() - (models_dir / "foo.sql").write_text( - """ + (models_dir / "foo.sql").write_text(""" MODEL ( name test.foo, kind full @@ -3445,8 +3451,7 @@ def test_cte_failure(tmp_path: Path) -> None: SELECT 1 AS id ) SELECT id FROM model_cte - """ - ) + """) config = Config( default_connection=DuckDBConnectionConfig(), @@ -3471,8 +3476,7 @@ def test_cte_failure(tmp_path: Path) -> None: # Case 1: Ensure that a single CTE failure is reported correctly tests_dir = tmp_path / "tests" tests_dir.mkdir() - (tests_dir / "test_foo.yaml").write_text( - """ + (tests_dir / "test_foo.yaml").write_text(""" test_foo: model: test.foo outputs: @@ -3482,8 +3486,7 @@ def test_cte_failure(tmp_path: Path) -> None: - id: 2 query: - id: 1 - """ - ) + """) # Re-initialize context to pick up the new test file context = Context(paths=tmp_path, config=config) @@ -3500,8 +3503,7 @@ def test_cte_failure(tmp_path: Path) -> None: assert "Failed tests (1)" in output # Case 2: Ensure that both CTE and query failures are reported correctly - (tests_dir / "test_foo.yaml").write_text( - """ + (tests_dir / "test_foo.yaml").write_text(""" test_foo: model: test.foo outputs: @@ -3511,8 +3513,7 @@ def test_cte_failure(tmp_path: Path) -> None: - id: 2 query: - id: 2 - """ - ) + """) # Re-initialize context to pick up the modified test file context = Context(paths=tmp_path, config=config) diff --git a/tests/dbt/cli/conftest.py b/tests/dbt/cli/conftest.py index 26757bf3ab..2d24b205c9 100644 --- a/tests/dbt/cli/conftest.py +++ b/tests/dbt/cli/conftest.py @@ -1,7 +1,8 @@ -import typing as t import functools -from click.testing import CliRunner, Result +import typing as t + import pytest +from click.testing import CliRunner, Result @pytest.fixture diff --git a/tests/dbt/cli/test_global_flags.py b/tests/dbt/cli/test_global_flags.py index 7e2262bd80..e4c7c0c62d 100644 --- a/tests/dbt/cli/test_global_flags.py +++ b/tests/dbt/cli/test_global_flags.py @@ -1,17 +1,21 @@ +import logging import typing as t from pathlib import Path + import pytest -import logging -from pytest_mock import MockerFixture from click.testing import Result -from sqlmesh.utils.errors import SQLMeshError +from pytest_mock import MockerFixture from sqlglot.errors import SqlglotError + +from sqlmesh.utils.errors import SQLMeshError from tests.dbt.conftest import EmptyProjectCreator pytestmark = pytest.mark.slow -def test_profile_and_target(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result]): +def test_profile_and_target( + jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result] +): # profile doesnt exist - error result = invoke_cli(["--profile", "nonexist"]) assert result.exit_code == 1 @@ -97,7 +101,9 @@ def test_run_error_handler( assert "Traceback" not in result.output -def test_log_level(invoke_cli: t.Callable[..., Result], create_empty_project: EmptyProjectCreator): +def test_log_level( + invoke_cli: t.Callable[..., Result], create_empty_project: EmptyProjectCreator +): create_empty_project() result = invoke_cli(["--log-level", "info", "list"]) @@ -110,7 +116,9 @@ def test_log_level(invoke_cli: t.Callable[..., Result], create_empty_project: Em def test_profiles_dir( - invoke_cli: t.Callable[..., Result], create_empty_project: EmptyProjectCreator, tmp_path: Path + invoke_cli: t.Callable[..., Result], + create_empty_project: EmptyProjectCreator, + tmp_path: Path, ): project_dir, _ = create_empty_project(project_name="test_profiles_dir") @@ -129,7 +137,10 @@ def test_profiles_dir( assert result.exit_code > 0, result.output # alternative ~/.dbt/profiles.yml might exist but doesn't contain the profile - assert "profiles.yml not found" in result.output or "not found in profiles" in result.output + assert ( + "profiles.yml not found" in result.output + or "not found in profiles" in result.output + ) # should pass if we specify --profiles-dir result = invoke_cli(["--profiles-dir", str(new_profiles_yml.parent), "list"]) @@ -162,7 +173,10 @@ def test_project_dir( assert result.exit_code != 0, result.output # profiles.yml might exist but doesn't contain the profile - assert "profiles.yml not found" in result.output or "not found in profiles" in result.output + assert ( + "profiles.yml not found" in result.output + or "not found in profiles" in result.output + ) # should pass if it can find both files, either because we specified --profiles-dir explicitly or the profiles.yml was found in --project-dir result = invoke_cli( diff --git a/tests/dbt/cli/test_list.py b/tests/dbt/cli/test_list.py index 3e6a55125c..3e907f6f49 100644 --- a/tests/dbt/cli/test_list.py +++ b/tests/dbt/cli/test_list.py @@ -1,6 +1,7 @@ import typing as t -import pytest from pathlib import Path + +import pytest from click.testing import Result pytestmark = pytest.mark.slow @@ -32,7 +33,9 @@ def test_list_select(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Resul assert "─ jaffle_shop.raw_orders" not in result.output -def test_list_select_exclude(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result]): +def test_list_select_exclude( + jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result] +): # single exclude result = invoke_cli(["list", "--select", "raw_customers+", "--exclude", "orders"]) @@ -63,22 +66,19 @@ def test_list_select_exclude(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[.. def test_list_with_vars(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result]): - ( - jaffle_shop_duckdb / "models" / "vars_model.sql" - ).write_text(""" + (jaffle_shop_duckdb / "models" / "vars_model.sql").write_text( + """ select * from {{ ref('custom' + var('foo')) }} - """) + """ + ) result = invoke_cli(["list", "--vars", "foo: ers"]) assert result.exit_code == 0 assert not result.exception - assert ( - """├── jaffle_shop.vars_model -│ └── depends_on: jaffle_shop.customers""" - in result.output - ) + assert """├── jaffle_shop.vars_model +│ └── depends_on: jaffle_shop.customers""" in result.output def test_list_models_mutually_exclusive( @@ -90,7 +90,9 @@ def test_list_models_mutually_exclusive( result = invoke_cli(["list", "--resource-type", "test", "--models", "bar"]) assert result.exit_code != 0 - assert '"models" and "resource_type" are mutually exclusive arguments' in result.output + assert ( + '"models" and "resource_type" are mutually exclusive arguments' in result.output + ) def test_list_models(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result]): diff --git a/tests/dbt/cli/test_operations.py b/tests/dbt/cli/test_operations.py index 4aa508e21f..6feb056751 100644 --- a/tests/dbt/cli/test_operations.py +++ b/tests/dbt/cli/test_operations.py @@ -1,15 +1,17 @@ +import logging import typing as t from pathlib import Path + import pytest -from sqlmesh_dbt.operations import create -from sqlmesh_dbt.console import DbtCliConsole -from sqlmesh.utils import yaml -from sqlmesh.utils.errors import SQLMeshError import time_machine -from sqlmesh.core.plan import PlanBuilder + from sqlmesh.core.config.common import VirtualEnvironmentMode +from sqlmesh.core.plan import PlanBuilder +from sqlmesh.utils import yaml +from sqlmesh.utils.errors import SQLMeshError +from sqlmesh_dbt.console import DbtCliConsole +from sqlmesh_dbt.operations import create from tests.dbt.conftest import EmptyProjectCreator -import logging pytestmark = pytest.mark.slow @@ -35,7 +37,7 @@ def plan( def test_create_sets_and_persists_default_start_date(jaffle_shop_duckdb: Path): with time_machine.travel("2020-01-02 00:00:00 UTC"): - from sqlmesh.utils.date import yesterday_ds, to_ds + from sqlmesh.utils.date import to_ds, yesterday_ds assert yesterday_ds() == "2020-01-01" @@ -50,7 +52,7 @@ def test_create_sets_and_persists_default_start_date(jaffle_shop_duckdb: Path): ) # check that the date set on the first invocation persists to future invocations - from sqlmesh.utils.date import yesterday_ds, to_ds + from sqlmesh.utils.date import to_ds, yesterday_ds assert yesterday_ds() != "2020-01-01" @@ -128,7 +130,9 @@ def test_run_option_mapping(jaffle_shop_duckdb: Path): operations.context.console = console plan = operations.run() - standalone_audit_name = "relationships_orders_customer_id__customer_id__ref_customers_" + standalone_audit_name = ( + "relationships_orders_customer_id__customer_id__ref_customers_" + ) assert plan.environment.name == "prod" assert console.no_prompts is True assert console.no_diff is True @@ -180,9 +184,9 @@ def test_run_option_mapping(jaffle_shop_duckdb: Path): assert plan.end_bounded is False assert plan.ignore_cron is True assert plan.skip_backfill is False - assert plan.selected_models_to_backfill == {k for k in operations.context.snapshots} - { - '"jaffle_shop"."main"."customers"' - } - {standalone_audit_name} + assert plan.selected_models_to_backfill == { + k for k in operations.context.snapshots + } - {'"jaffle_shop"."main"."customers"'} - {standalone_audit_name} assert {s.name for s in plan.snapshots} == ( plan.selected_models_to_backfill | {standalone_audit_name} ) @@ -271,7 +275,9 @@ def test_run_option_mapping_dev(jaffle_shop_duckdb: Path): ], ) def test_run_option_full_refresh( - create_empty_project: EmptyProjectCreator, env_name: str, vde_mode: VirtualEnvironmentMode + create_empty_project: EmptyProjectCreator, + env_name: str, + vde_mode: VirtualEnvironmentMode, ): # create config file prior to load project_path, models_path = create_empty_project(project_name="test") @@ -304,7 +310,9 @@ def test_run_option_full_refresh( assert plan.requires_backfill assert not plan.empty_backfill assert not plan.skip_backfill - assert plan.models_to_backfill == set(['"test"."main"."model_a"', '"test"."main"."model_b"']) + assert plan.models_to_backfill == set( + ['"test"."main"."model_a"', '"test"."main"."model_b"'] + ) if vde_mode == VirtualEnvironmentMode.DEV_ONLY: # We do not clear intervals across all model versions in the default DEV_ONLY mode, even when targeting prod, @@ -336,7 +344,9 @@ def test_run_option_full_refresh_with_selector(jaffle_shop_duckdb: Path): assert plan.models_to_backfill == set(['"jaffle_shop"."main"."stg_customers"']) -def test_create_sets_concurrent_tasks_based_on_threads(create_empty_project: EmptyProjectCreator): +def test_create_sets_concurrent_tasks_based_on_threads( + create_empty_project: EmptyProjectCreator, +): project_dir, _ = create_empty_project(project_name="test") # add a postgres target because duckdb overrides to concurrent_tasks=1 regardless of what gets specified diff --git a/tests/dbt/cli/test_options.py b/tests/dbt/cli/test_options.py index 962ff0beb3..d85964b10c 100644 --- a/tests/dbt/cli/test_options.py +++ b/tests/dbt/cli/test_options.py @@ -1,8 +1,10 @@ import typing as t + import pytest -from sqlmesh_dbt.options import YamlParamType from click.exceptions import BadParameter +from sqlmesh_dbt.options import YamlParamType + @pytest.mark.parametrize( "input,expected", @@ -15,7 +17,9 @@ ("{key: value, date: 20180101}", {"key": "value", "date": 20180101}), ], ) -def test_yaml_param_type(input: str, expected: t.Union[BadParameter, t.Dict[str, t.Any]]): +def test_yaml_param_type( + input: str, expected: t.Union[BadParameter, t.Dict[str, t.Any]] +): if isinstance(expected, BadParameter): with pytest.raises(BadParameter, match=expected.message): YamlParamType().convert(input, None, None) diff --git a/tests/dbt/cli/test_run.py b/tests/dbt/cli/test_run.py index c640950a27..20dba13f1b 100644 --- a/tests/dbt/cli/test_run.py +++ b/tests/dbt/cli/test_run.py @@ -1,9 +1,11 @@ +import shutil import typing as t -import pytest from pathlib import Path -import shutil -from click.testing import Result + +import pytest import time_machine +from click.testing import Result + from sqlmesh_dbt.operations import create from tests.cli.test_cli import FREEZE_TIME from tests.dbt.conftest import EmptyProjectCreator @@ -20,7 +22,9 @@ def test_run(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result]): assert "Model batches executed" in result.output -def test_run_with_selectors(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result]): +def test_run_with_selectors( + jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result] +): with time_machine.travel(FREEZE_TIME): # do an initial run to create the objects # otherwise the selected subset may depend on something that hasnt been created @@ -49,7 +53,9 @@ def test_run_with_changes_and_full_refresh( project_path, models_path = create_empty_project(project_name="test") engine_adapter = create(project_path).context.engine_adapter - engine_adapter.execute("create table external_table as select 'foo' as a, 'bar' as b") + engine_adapter.execute( + "create table external_table as select 'foo' as a, 'bar' as b" + ) (models_path / "model_a.sql").write_text("select a, b from external_table") (models_path / "model_b.sql").write_text("select a, b from {{ ref('model_a') }}") @@ -83,14 +89,19 @@ def test_run_with_changes_and_full_refresh( assert result.exit_code == 0 assert not result.exception - assert engine_adapter.fetchall("select a, b from model_a") == [("foo", "bar"), ("baz", "bing")] + assert engine_adapter.fetchall("select a, b from model_a") == [ + ("foo", "bar"), + ("baz", "bing"), + ] assert engine_adapter.fetchall("select a, b, c from model_b") == [ ("foo", "bar", "changed"), ("baz", "bing", "changed"), ] -def test_run_with_threads(jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result]): +def test_run_with_threads( + jaffle_shop_duckdb: Path, invoke_cli: t.Callable[..., Result] +): result = invoke_cli(["run", "--threads", "4"]) assert result.exit_code == 0 assert not result.exception diff --git a/tests/dbt/cli/test_selectors.py b/tests/dbt/cli/test_selectors.py index 17f0195f58..1ac8ff7cec 100644 --- a/tests/dbt/cli/test_selectors.py +++ b/tests/dbt/cli/test_selectors.py @@ -1,9 +1,11 @@ import typing as t +from pathlib import Path + import pytest -from sqlmesh_dbt import selectors -from sqlmesh.core.selector import DbtSelector + from sqlmesh.core.context import Context -from pathlib import Path +from sqlmesh.core.selector import DbtSelector +from sqlmesh_dbt import selectors @pytest.mark.parametrize( @@ -62,7 +64,9 @@ def test_exclusion(dbt_exclude: t.List[str], expected: t.Optional[str]): def test_selection_and_exclusion( dbt_select: t.List[str], dbt_exclude: t.List[str], expected: t.Optional[str] ): - assert selectors.to_sqlmesh(dbt_select=dbt_select, dbt_exclude=dbt_exclude) == expected + assert ( + selectors.to_sqlmesh(dbt_select=dbt_select, dbt_exclude=dbt_exclude) == expected + ) @pytest.mark.parametrize( @@ -191,7 +195,10 @@ def test_select_by_dbt_names( '"jaffle_shop"."main"."raw_payments"', }, ), - (["+customers"], {'"jaffle_shop"."main"."orders"', '"jaffle_shop"."main"."agg_orders"'}), + ( + ["+customers"], + {'"jaffle_shop"."main"."orders"', '"jaffle_shop"."main"."agg_orders"'}, + ), ( ["+tag:agg"], { @@ -267,7 +274,9 @@ def test_selection_and_exclusion_by_dbt_names( selector = ctx._new_selector() assert isinstance(selector, DbtSelector) - sqlmesh_selector = selectors.to_sqlmesh(dbt_select=dbt_select, dbt_exclude=dbt_exclude) + sqlmesh_selector = selectors.to_sqlmesh( + dbt_select=dbt_select, dbt_exclude=dbt_exclude + ) assert sqlmesh_selector assert selector.expand_model_selections([sqlmesh_selector]) == expected @@ -296,8 +305,12 @@ def test_selection_and_exclusion_by_dbt_names( ), ], ) -def test_consolidate(input_args: t.Dict[str, t.Any], expected: t.Union[t.Tuple[str, str], str]): - all_input_args: t.Dict[str, t.Any] = dict(select=[], exclude=[], models=[], resource_type=None) +def test_consolidate( + input_args: t.Dict[str, t.Any], expected: t.Union[t.Tuple[str, str], str] +): + all_input_args: t.Dict[str, t.Any] = dict( + select=[], exclude=[], models=[], resource_type=None + ) all_input_args.update(input_args) @@ -318,7 +331,9 @@ def test_models_by_dbt_names(jaffle_shop_duckdb_context: Context): assert isinstance(selector, DbtSelector) selector_expr = selectors.to_sqlmesh( - *selectors.consolidate(select=[], exclude=[], models=["jaffle_shop"], resource_type=None) + *selectors.consolidate( + select=[], exclude=[], models=["jaffle_shop"], resource_type=None + ) ) assert selector_expr diff --git a/tests/dbt/conftest.py b/tests/dbt/conftest.py index 5e6444c8e6..a9e5ca7bc6 100644 --- a/tests/dbt/conftest.py +++ b/tests/dbt/conftest.py @@ -1,7 +1,8 @@ from __future__ import annotations -import typing as t import os +import typing as t +import uuid from pathlib import Path import pytest @@ -12,7 +13,6 @@ from sqlmesh.dbt.project import Project from sqlmesh.dbt.target import PostgresConfig from sqlmesh_dbt.operations import init_project_if_required -import uuid class EmptyProjectCreator(t.Protocol): @@ -82,8 +82,12 @@ def _create_empty_project( @pytest.fixture -def jaffle_shop_duckdb(copy_to_temp_path: t.Callable[..., t.List[Path]]) -> t.Iterable[Path]: - fixture_path = Path(__file__).parent.parent / "fixtures" / "dbt" / "jaffle_shop_duckdb" +def jaffle_shop_duckdb( + copy_to_temp_path: t.Callable[..., t.List[Path]], +) -> t.Iterable[Path]: + fixture_path = ( + Path(__file__).parent.parent / "fixtures" / "dbt" / "jaffle_shop_duckdb" + ) assert fixture_path.exists() current_path = os.getcwd() @@ -106,7 +110,9 @@ def jaffle_shop_duckdb_context(jaffle_shop_duckdb: Path) -> Context: @pytest.fixture() def runtime_renderer() -> t.Callable: def create_renderer(context: DbtContext, **kwargs: t.Any) -> t.Callable: - environment = context.jinja_macros.build_environment(**{**context.jinja_globals, **kwargs}) + environment = context.jinja_macros.build_environment( + **{**context.jinja_globals, **kwargs} + ) def render(value: str) -> str: return environment.from_string(value).render() diff --git a/tests/dbt/test_adapter.py b/tests/dbt/test_adapter.py index 5570212668..0598144b45 100644 --- a/tests/dbt/test_adapter.py +++ b/tests/dbt/test_adapter.py @@ -11,6 +11,8 @@ from sqlglot import exp, parse_one from sqlmesh.core.dialect import schema_ +from sqlmesh.core.schema_diff import (SchemaDiffer, + TableAlterChangeColumnTypeOperation) from sqlmesh.core.snapshot import SnapshotId from sqlmesh.dbt.adapter import ParsetimeAdapter from sqlmesh.dbt.project import Project @@ -18,7 +20,6 @@ from sqlmesh.dbt.target import BigQueryConfig, SnowflakeConfig from sqlmesh.utils.errors import ConfigError from sqlmesh.utils.jinja import JinjaMacroRegistry -from sqlmesh.core.schema_diff import SchemaDiffer, TableAlterChangeColumnTypeOperation pytestmark = pytest.mark.dbt @@ -36,22 +37,29 @@ def test_adapter_relation(sushi_test_project: Project, runtime_renderer: t.Calla table_name="foo.bar", target_columns_to_types={"baz": exp.DataType.build("int")} ) engine_adapter.create_table( - table_name="foo.another", target_columns_to_types={"col": exp.DataType.build("int")} + table_name="foo.another", + target_columns_to_types={"col": exp.DataType.build("int")}, ) engine_adapter.create_view( - view_name="foo.bar_view", query_or_df=t.cast(exp.Query, parse_one("select * from foo.bar")) + view_name="foo.bar_view", + query_or_df=t.cast(exp.Query, parse_one("select * from foo.bar")), ) engine_adapter.create_table( - table_name="ignored.ignore", target_columns_to_types={"col": exp.DataType.build("int")} + table_name="ignored.ignore", + target_columns_to_types={"col": exp.DataType.build("int")}, ) assert ( - renderer("{{ adapter.get_relation(database=None, schema='foo', identifier='bar') }}") + renderer( + "{{ adapter.get_relation(database=None, schema='foo', identifier='bar') }}" + ) == '"memory"."foo"."bar"' ) assert ( - renderer("{{ adapter.get_relation(database=None, schema='foo', identifier='bar').type }}") + renderer( + "{{ adapter.get_relation(database=None, schema='foo', identifier='bar').type }}" + ) == "table" ) @@ -66,15 +74,16 @@ def test_adapter_relation(sushi_test_project: Project, runtime_renderer: t.Calla "{%- set relation = adapter.get_relation(database=None, schema='foo', identifier='bar') -%} {{ adapter.get_columns_in_relation(relation) }}" ) == str([Column.from_description(name="baz", raw_data_type="INT")]) - assert renderer("{{ adapter.list_relations(database=None, schema='foo')|length }}") == "3" + assert ( + renderer("{{ adapter.list_relations(database=None, schema='foo')|length }}") + == "3" + ) - assert renderer( - """ + assert renderer(""" {%- set from = adapter.get_relation(database=None, schema='foo', identifier='bar') -%} {%- set to = adapter.get_relation(database=None, schema='foo', identifier='another') -%} {{ adapter.get_missing_columns(from, to) -}} - """ - ) == str([Column.from_description(name="baz", raw_data_type="INT")]) + """) == str([Column.from_description(name="baz", raw_data_type="INT")]) assert ( renderer( @@ -129,7 +138,9 @@ def test_bigquery_get_columns_in_relation( SchemaField(name="created_at", field_type="TIMESTAMP", mode="NULLABLE"), ] adapter_mock.get_bq_schema.return_value = table_schema - renderer = runtime_renderer(context, engine_adapter=adapter_mock, dialect="bigquery") + renderer = runtime_renderer( + context, engine_adapter=adapter_mock, dialect="bigquery" + ) assert renderer( "{%- set relation = api.Relation.create(database='test', schema='test', identifier='test_table') -%}" "{{ adapter.get_columns_in_relation(relation) }}" @@ -145,7 +156,9 @@ def test_normalization( context = sushi_test_project.context assert context.target - data_object = DataObject(catalog="test", schema="bla", name="bob", type=DataObjectType.TABLE) + data_object = DataObject( + catalog="test", schema="bla", name="bob", type=DataObjectType.TABLE + ) # bla and bob will be normalized to lowercase since the target is duckdb adapter_mock = mocker.MagicMock() @@ -157,7 +170,9 @@ def test_normalization( schema_bla = schema_("bla", "test", quoted=True) relation_bla_bob = exp.table_("bob", db="bla", catalog="test", quoted=True) - duckdb_renderer("{{ adapter.get_relation(database=None, schema='bla', identifier='bob') }}") + duckdb_renderer( + "{{ adapter.get_relation(database=None, schema='bla', identifier='bob') }}" + ) adapter_mock.get_data_object.assert_has_calls([call(relation_bla_bob)]) # bla and bob will be normalized to uppercase since the target is Snowflake, even though the default dialect is duckdb @@ -178,10 +193,14 @@ def test_normalization( schema_bla = schema_("bla", "test", quoted=True) relation_bla_bob = exp.table_("bob", db="bla", catalog="test", quoted=True) - renderer("{{ adapter.get_relation(database=None, schema='bla', identifier='bob') }}") + renderer( + "{{ adapter.get_relation(database=None, schema='bla', identifier='bob') }}" + ) adapter_mock.get_data_object.assert_has_calls([call(relation_bla_bob)]) - renderer("{{ adapter.get_relation(database='custom_db', schema='bla', identifier='bob') }}") + renderer( + "{{ adapter.get_relation(database='custom_db', schema='bla', identifier='bob') }}" + ) adapter_mock.get_data_object.assert_has_calls( [call(exp.table_("bob", db="bla", catalog="custom_db", quoted=True))] ) @@ -212,10 +231,14 @@ def test_normalization( # raise in adapter.execute right before returning from the method with pytest.raises(AssertionError): renderer("{{ run_query('SELECT * FROM t') }}") - adapter_mock.fetchdf.assert_has_calls([call(expected_star_query, quote_identifiers=False)]) + adapter_mock.fetchdf.assert_has_calls( + [call(expected_star_query, quote_identifiers=False)] + ) renderer("{% call statement('something') %} {{ 'SELECT * FROM t' }} {% endcall %}") - adapter_mock.execute.assert_has_calls([call(expected_star_query, quote_identifiers=False)]) + adapter_mock.execute.assert_has_calls( + [call(expected_star_query, quote_identifiers=False)] + ) # Enforce case-sensitivity for database object names setattr( @@ -239,10 +262,15 @@ def test_normalization( def test_adapter_dispatch(sushi_test_project: Project, runtime_renderer: t.Callable): context = sushi_test_project.context renderer = runtime_renderer(context) - assert renderer("{{ adapter.dispatch('current_engine', 'customers')() }}") == "duckdb" + assert ( + renderer("{{ adapter.dispatch('current_engine', 'customers')() }}") == "duckdb" + ) assert renderer("{{ adapter.dispatch('current_timestamp')() }}") == "now()" assert renderer("{{ adapter.dispatch('current_timestamp', 'dbt')() }}") == "now()" - assert renderer("{{ adapter.dispatch('select_distinct', 'customers')() }}") == "distinct" + assert ( + renderer("{{ adapter.dispatch('select_distinct', 'customers')() }}") + == "distinct" + ) # test with keyword arguments assert ( @@ -251,22 +279,30 @@ def test_adapter_dispatch(sushi_test_project: Project, runtime_renderer: t.Calla ) == "duckdb" ) - assert renderer("{{ adapter.dispatch(macro_name='current_timestamp')() }}") == "now()" assert ( - renderer("{{ adapter.dispatch(macro_name='current_timestamp', macro_namespace='dbt')() }}") + renderer("{{ adapter.dispatch(macro_name='current_timestamp')() }}") == "now()" + ) + assert ( + renderer( + "{{ adapter.dispatch(macro_name='current_timestamp', macro_namespace='dbt')() }}" + ) == "now()" ) # mixing positional and keyword arguments assert ( - renderer("{{ adapter.dispatch('current_engine', macro_namespace='customers')() }}") + renderer( + "{{ adapter.dispatch('current_engine', macro_namespace='customers')() }}" + ) == "duckdb" ) assert ( - renderer("{{ adapter.dispatch('current_timestamp', macro_namespace=None)() }}") == "now()" + renderer("{{ adapter.dispatch('current_timestamp', macro_namespace=None)() }}") + == "now()" ) assert ( - renderer("{{ adapter.dispatch('current_timestamp', macro_namespace='dbt')() }}") == "now()" + renderer("{{ adapter.dispatch('current_timestamp', macro_namespace='dbt')() }}") + == "now()" ) with pytest.raises(ConfigError, match=r"Macro 'current_engine'.*was not found."): @@ -316,9 +352,9 @@ def test_adapter_map_snapshot_tables( table_name="foo.bar", target_columns_to_types={"col": exp.DataType.build("int")} ) - expected_test_model_table_name = parse_one('"memory"."sqlmesh"."test_db__test_model"').sql( - dialect=project_dialect - ) + expected_test_model_table_name = parse_one( + '"memory"."sqlmesh"."test_db__test_model"' + ).sql(dialect=project_dialect) assert ( renderer( @@ -329,15 +365,22 @@ def test_adapter_map_snapshot_tables( assert "baz" in renderer("{{ run_query('SELECT * FROM test_db.test_model') }}") - expected_foo_bar_table_name = parse_one('"memory"."foo"."bar"').sql(dialect=project_dialect) + expected_foo_bar_table_name = parse_one('"memory"."foo"."bar"').sql( + dialect=project_dialect + ) assert ( - renderer("{{ adapter.get_relation(database=none, schema='foo', identifier='bar') }}") + renderer( + "{{ adapter.get_relation(database=none, schema='foo', identifier='bar') }}" + ) == expected_foo_bar_table_name ) assert renderer("{{ adapter.resolve_schema(test_model) }}") == "sqlmesh" - assert renderer("{{ adapter.resolve_identifier(test_model) }}") == "test_db__test_model" + assert ( + renderer("{{ adapter.resolve_identifier(test_model) }}") + == "test_db__test_model" + ) assert renderer("{{ adapter.resolve_schema(foo_bar) }}") == "foo" assert renderer("{{ adapter.resolve_identifier(foo_bar) }}") == "bar" @@ -369,15 +412,20 @@ def test_adapter_get_relation_normalization( assert context.target engine_adapter = context.target.to_sqlmesh().create_engine_adapter() engine_adapter._default_catalog = '"memory"' - renderer = runtime_renderer(context, engine_adapter=engine_adapter, dialect="snowflake") + renderer = runtime_renderer( + context, engine_adapter=engine_adapter, dialect="snowflake" + ) engine_adapter.create_schema('"FOO"') engine_adapter.create_table( - table_name='"FOO"."BAR"', target_columns_to_types={"baz": exp.DataType.build("int")} + table_name='"FOO"."BAR"', + target_columns_to_types={"baz": exp.DataType.build("int")}, ) assert ( - renderer("{{ adapter.get_relation(database=None, schema='foo', identifier='bar') }}") + renderer( + "{{ adapter.get_relation(database=None, schema='foo', identifier='bar') }}" + ) == '"memory"."FOO"."BAR"' ) @@ -402,9 +450,15 @@ def test_adapter_expand_target_column_types( from_columns = { "int_col": exp.DataType.build("int"), "same_text_col": exp.DataType.build("varchar(1)"), # varchar(1) -> varchar(1) - "unexpandable_text_col": exp.DataType.build("varchar(2)"), # varchar(4) -> varchar(2) - "expandable_text_col1": exp.DataType.build("varchar(16)"), # varchar(8) -> varchar(16) - "expandable_text_col2": exp.DataType.build("varchar(64)"), # varchar(32) -> varchar(64) + "unexpandable_text_col": exp.DataType.build( + "varchar(2)" + ), # varchar(4) -> varchar(2) + "expandable_text_col1": exp.DataType.build( + "varchar(16)" + ), # varchar(8) -> varchar(16) + "expandable_text_col2": exp.DataType.build( + "varchar(64)" + ), # varchar(32) -> varchar(64) } to_columns = { "int_col": exp.DataType.build("int"), diff --git a/tests/dbt/test_config.py b/tests/dbt/test_config.py index 01f4cf6128..9e1641d8d1 100644 --- a/tests/dbt/test_config.py +++ b/tests/dbt/test_config.py @@ -6,7 +6,6 @@ import pytest from dbt.adapters.base import BaseRelation, Column from pytest_mock import MockerFixture - from sqlglot import exp from sqlmesh import Context @@ -14,40 +13,29 @@ from sqlmesh.core.config import Config, ModelDefaultsConfig from sqlmesh.core.dialect import jinja_query from sqlmesh.core.model import SqlModel -from sqlmesh.core.model.kind import OnDestructiveChange, OnAdditiveChange +from sqlmesh.core.model.kind import OnAdditiveChange, OnDestructiveChange from sqlmesh.core.state_sync import CachingStateSync, EngineAdapterStateSync from sqlmesh.dbt.builtin import Api from sqlmesh.dbt.column import ColumnConfig from sqlmesh.dbt.common import Dependencies from sqlmesh.dbt.context import DbtContext from sqlmesh.dbt.loader import sqlmesh_config -from sqlmesh.dbt.model import ( - IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind, - Materialization, - ModelConfig, -) +from sqlmesh.dbt.model import (IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, Materialization, + ModelConfig) from sqlmesh.dbt.project import Project from sqlmesh.dbt.relation import Policy from sqlmesh.dbt.source import SourceConfig -from sqlmesh.dbt.target import ( - TARGET_TYPE_TO_CONFIG_CLASS, - BigQueryConfig, - DatabricksConfig, - DuckDbConfig, - MSSQLConfig, - PostgresConfig, - RedshiftConfig, - SnowflakeConfig, - TargetConfig, - TrinoConfig, - AthenaConfig, - ClickhouseConfig, - SCHEMA_DIFFER_OVERRIDES, -) +from sqlmesh.dbt.target import (SCHEMA_DIFFER_OVERRIDES, + TARGET_TYPE_TO_CONFIG_CLASS, AthenaConfig, + BigQueryConfig, ClickhouseConfig, + DatabricksConfig, DuckDbConfig, MSSQLConfig, + PostgresConfig, RedshiftConfig, + SnowflakeConfig, TargetConfig, TrinoConfig) from sqlmesh.dbt.test import TestConfig from sqlmesh.utils.errors import ConfigError -from sqlmesh.utils.yaml import load as yaml_load, dump as yaml_dump +from sqlmesh.utils.yaml import dump as yaml_dump +from sqlmesh.utils.yaml import load as yaml_load from tests.dbt.conftest import EmptyProjectCreator pytestmark = pytest.mark.dbt @@ -86,7 +74,9 @@ ({}, {"uknown": "value"}, {"uknown": "value"}), ], ) -def test_update(current: t.Dict[str, t.Any], new: t.Dict[str, t.Any], expected: t.Dict[str, t.Any]): +def test_update( + current: t.Dict[str, t.Any], new: t.Dict[str, t.Any], expected: t.Dict[str, t.Any] +): config = ModelConfig(**current).update_with(new) assert {k: v for k, v in config.dict().items() if k in expected} == expected @@ -401,7 +391,10 @@ def test_variables(assert_exp_eq, sushi_test_project): "invalid_var": "{{ ref('ref_without_closing_paren' }}", } assert sushi_test_project.packages["sushi"].variables == expected_sushi_variables - assert sushi_test_project.packages["customers"].variables == expected_customer_variables + assert ( + sushi_test_project.packages["customers"].variables + == expected_customer_variables + ) @pytest.mark.slow @@ -410,8 +403,14 @@ def test_variables_override(init_and_plan_context: t.Callable): "tests/fixtures/dbt/sushi_test", config="test_config_with_var_override" ) dbt_project = context._loaders[0]._load_projects()[0] # type: ignore - assert dbt_project.packages["sushi"].variables["some_var"] == "overridden_from_config_py" - assert dbt_project.packages["customers"].variables["some_var"] == "overridden_from_config_py" + assert ( + dbt_project.packages["sushi"].variables["some_var"] + == "overridden_from_config_py" + ) + assert ( + dbt_project.packages["customers"].variables["some_var"] + == "overridden_from_config_py" + ) @pytest.mark.slow @@ -427,9 +426,13 @@ def test_nested_variables(sushi_test_project): dependencies=Dependencies(variables=["nested_vars"]), ) context = sushi_test_project.context.copy() - context.set_and_render_variables(sushi_test_project.packages["sushi"].variables, "sushi") + context.set_and_render_variables( + sushi_test_project.packages["sushi"].variables, "sushi" + ) sqlmesh_model = model_config.to_sqlmesh(context) - assert sqlmesh_model.jinja_macros.global_objs["vars"]["nested_vars"] == {"some_nested_var": 2} + assert sqlmesh_model.jinja_macros.global_objs["vars"]["nested_vars"] == { + "some_nested_var": 2 + } @pytest.mark.slow @@ -448,12 +451,15 @@ def test_source_config(sushi_test_project: Project): "identifier": "order_items", } actual_config = { - k: getattr(source_configs["streaming.order_items"], k) for k, v in expected_config.items() + k: getattr(source_configs["streaming.order_items"], k) + for k, v in expected_config.items() } assert actual_config == expected_config assert ( - source_configs["streaming.order_items"].canonical_name(sushi_test_project.context) + source_configs["streaming.order_items"].canonical_name( + sushi_test_project.context + ) == "raw.order_items" ) @@ -481,7 +487,10 @@ def test_seed_config(sushi_test_project: Project, mocker: MockerFixture): assert raw_items_seed.to_sqlmesh(context).name == "sushi.waiter_names" raw_items_seed.dialect_ = "snowflake" - assert raw_items_seed.to_sqlmesh(sushi_test_project.context).name == "sushi.waiter_names" + assert ( + raw_items_seed.to_sqlmesh(sushi_test_project.context).name + == "sushi.waiter_names" + ) assert ( raw_items_seed.to_sqlmesh(sushi_test_project.context).fqn == '"MEMORY"."SUSHI"."WAITER_NAMES"' @@ -490,21 +499,33 @@ def test_seed_config(sushi_test_project: Project, mocker: MockerFixture): waiter_revenue_semicolon_seed = seed_configs["waiter_revenue_semicolon"] expected_config_semicolon = { - "path": Path(sushi_test_project.context.project_root, "seeds/waiter_revenue_semicolon.csv"), + "path": Path( + sushi_test_project.context.project_root, + "seeds/waiter_revenue_semicolon.csv", + ), "schema_": "sushi", "delimiter": ";", } actual_config_semicolon = { - k: getattr(waiter_revenue_semicolon_seed, k) for k, v in expected_config_semicolon.items() + k: getattr(waiter_revenue_semicolon_seed, k) + for k, v in expected_config_semicolon.items() } assert actual_config_semicolon == expected_config_semicolon - assert waiter_revenue_semicolon_seed.canonical_name(context) == "sushi.waiter_revenue_semicolon" assert ( - waiter_revenue_semicolon_seed.to_sqlmesh(context).name == "sushi.waiter_revenue_semicolon" + waiter_revenue_semicolon_seed.canonical_name(context) + == "sushi.waiter_revenue_semicolon" + ) + assert ( + waiter_revenue_semicolon_seed.to_sqlmesh(context).name + == "sushi.waiter_revenue_semicolon" ) assert waiter_revenue_semicolon_seed.delimiter == ";" - assert set(waiter_revenue_semicolon_seed.columns.keys()) == {"waiter_id", "revenue", "quarter"} + assert set(waiter_revenue_semicolon_seed.columns.keys()) == { + "waiter_id", + "revenue", + "quarter", + } def test_quoting(): @@ -548,7 +569,9 @@ def mock_source_macro(source_name, table_name): source_name="my_source", identifier="FILENAME.CSV", ) - assert source_dot.canonical_name(mock_context) == 'RAW_DEV.raw_schema."FILENAME.CSV"' + assert ( + source_dot.canonical_name(mock_context) == 'RAW_DEV.raw_schema."FILENAME.CSV"' + ) # 2. Identifier with a space source_space = SourceConfig( @@ -556,7 +579,10 @@ def mock_source_macro(source_name, table_name): source_name="my_source", identifier="my table space", ) - assert source_space.canonical_name(mock_context) == 'RAW_DEV.raw_schema."my table space"' + assert ( + source_space.canonical_name(mock_context) + == 'RAW_DEV.raw_schema."my table space"' + ) # 3. Standard identifier (without dots or spaces) should not be quoted source_std = SourceConfig( @@ -582,8 +608,14 @@ def mock_source_macro(source_name, table_name): identifier="my_table_std", ) - assert source_dot_target.canonical_name(mock_context_target_db) == 'raw_schema."FILENAME.CSV"' - assert source_std_target.canonical_name(mock_context_target_db) == "raw_schema.my_table_std" + assert ( + source_dot_target.canonical_name(mock_context_target_db) + == 'raw_schema."FILENAME.CSV"' + ) + assert ( + source_std_target.canonical_name(mock_context_target_db) + == "raw_schema.my_table_std" + ) def _test_warehouse_config( @@ -610,8 +642,7 @@ def test_duckdb_threads(tmp_path): copytree(dbt_project_dir, temp_dir, symlinks=True) with open(temp_dir / "profiles.yml", "w", encoding="utf-8") as f: - f.write( - """ + f.write(""" sushi: outputs: in_memory: @@ -619,8 +650,7 @@ def test_duckdb_threads(tmp_path): schema: sushi threads: 4 target: in_memory - """ - ) + """) config = sqlmesh_config(temp_dir) assert config.gateways["in_memory"].connection.concurrent_tasks == 1 @@ -652,7 +682,8 @@ def test_snowflake_config(): sqlmesh_config = config.to_sqlmesh() assert sqlmesh_config.application == "Tobiko_SQLMesh" assert ( - sqlmesh_config.schema_differ_overrides == SCHEMA_DIFFER_OVERRIDES["schema_differ_overrides"] + sqlmesh_config.schema_differ_overrides + == SCHEMA_DIFFER_OVERRIDES["schema_differ_overrides"] ) @@ -883,7 +914,10 @@ def test_databricks_config_oauth(): assert as_sqlmesh.auth_type == "databricks-oauth" assert as_sqlmesh.oauth_client_id == "client-id" assert as_sqlmesh.oauth_client_secret == "client-secret" - assert as_sqlmesh.schema_differ_overrides == SCHEMA_DIFFER_OVERRIDES["schema_differ_overrides"] + assert ( + as_sqlmesh.schema_differ_overrides + == SCHEMA_DIFFER_OVERRIDES["schema_differ_overrides"] + ) def test_bigquery_config(): @@ -1086,16 +1120,22 @@ def test_db_type_to_relation_class(): from dbt.adapters.snowflake import SnowflakeRelation assert (TARGET_TYPE_TO_CONFIG_CLASS["bigquery"].relation_class) == BigQueryRelation - assert (TARGET_TYPE_TO_CONFIG_CLASS["databricks"].relation_class) == DatabricksRelation + assert ( + TARGET_TYPE_TO_CONFIG_CLASS["databricks"].relation_class + ) == DatabricksRelation assert (TARGET_TYPE_TO_CONFIG_CLASS["duckdb"].relation_class) == DuckDBRelation assert (TARGET_TYPE_TO_CONFIG_CLASS["redshift"].relation_class) == RedshiftRelation - assert (TARGET_TYPE_TO_CONFIG_CLASS["snowflake"].relation_class) == SnowflakeRelation + assert ( + TARGET_TYPE_TO_CONFIG_CLASS["snowflake"].relation_class + ) == SnowflakeRelation + from dbt.adapters.athena.relation import AthenaRelation from dbt.adapters.clickhouse.relation import ClickHouseRelation from dbt.adapters.trino.relation import TrinoRelation - from dbt.adapters.athena.relation import AthenaRelation - assert (TARGET_TYPE_TO_CONFIG_CLASS["clickhouse"].relation_class) == ClickHouseRelation + assert ( + TARGET_TYPE_TO_CONFIG_CLASS["clickhouse"].relation_class + ) == ClickHouseRelation assert (TARGET_TYPE_TO_CONFIG_CLASS["trino"].relation_class) == TrinoRelation assert (TARGET_TYPE_TO_CONFIG_CLASS["athena"].relation_class) == AthenaRelation @@ -1112,9 +1152,9 @@ def test_db_type_to_column_class(): assert (TARGET_TYPE_TO_CONFIG_CLASS["duckdb"].column_class) == Column assert (TARGET_TYPE_TO_CONFIG_CLASS["snowflake"].column_class) == SnowflakeColumn + from dbt.adapters.athena.column import AthenaColumn from dbt.adapters.clickhouse.column import ClickHouseColumn from dbt.adapters.trino.column import TrinoColumn - from dbt.adapters.athena.column import AthenaColumn assert (TARGET_TYPE_TO_CONFIG_CLASS["clickhouse"].column_class) == ClickHouseColumn assert (TARGET_TYPE_TO_CONFIG_CLASS["trino"].column_class) == TrinoColumn @@ -1130,7 +1170,9 @@ def test_variable_override(): project = Project.load( DbtContext( project_root=Path(project_root), - sqlmesh_config=Config(model_defaults=ModelDefaultsConfig(start="2021-01-01")), + sqlmesh_config=Config( + model_defaults=ModelDefaultsConfig(start="2021-01-01") + ), ), variables={"yet_another_var": 2}, ) @@ -1151,7 +1193,9 @@ def test_depends_on(assert_exp_eq, sushi_test_project): sqlmesh_model = model_config.to_sqlmesh(context) assert sqlmesh_model.depends_on_ == {'"memory"."sushi"."waiter_revenue_by_day_v2"'} assert sqlmesh_model.depends_on == {'"memory"."sushi"."waiter_revenue_by_day_v2"'} - assert sqlmesh_model.full_depends_on == {'"memory"."sushi"."waiter_revenue_by_day_v2"'} + assert sqlmesh_model.full_depends_on == { + '"memory"."sushi"."waiter_revenue_by_day_v2"' + } # Make sure the query wasn't rendered assert not sqlmesh_model._query_renderer._cache @@ -1262,9 +1306,9 @@ def test_empty_vars_config(tmp_path): model.write_text("SELECT 1 as id") # Load the project + from sqlmesh.core.config import Config from sqlmesh.dbt.context import DbtContext from sqlmesh.dbt.project import Project - from sqlmesh.core.config import Config context = DbtContext(project_root=dbt_project_dir, sqlmesh_config=Config()) diff --git a/tests/dbt/test_custom_materializations.py b/tests/dbt/test_custom_materializations.py index c1625d0251..2c7e7cf44e 100644 --- a/tests/dbt/test_custom_materializations.py +++ b/tests/dbt/test_custom_materializations.py @@ -9,11 +9,11 @@ from sqlmesh.core.config import ModelDefaultsConfig from sqlmesh.core.engine_adapter import DuckDBEngineAdapter from sqlmesh.core.model.kind import DbtCustomKind +from sqlmesh.dbt.basemodel import Materialization from sqlmesh.dbt.context import DbtContext from sqlmesh.dbt.manifest import ManifestHelper from sqlmesh.dbt.model import ModelConfig from sqlmesh.dbt.profile import Profile -from sqlmesh.dbt.basemodel import Materialization pytestmark = pytest.mark.dbt @@ -39,7 +39,9 @@ def test_custom_materialization_manifest_loading(): assert custom_incremental.adapter == "default" assert "make_temp_relation(new_relation)" in custom_incremental.definition assert "run_hooks(pre_hooks)" in custom_incremental.definition - assert " {{ return({'relations': [new_relation]}) }}" in custom_incremental.definition + assert ( + " {{ return({'relations': [new_relation]}) }}" in custom_incremental.definition + ) @pytest.mark.xdist_group("dbt_manifest") @@ -131,7 +133,9 @@ def test_custom_materialization_model_kind(): assert isinstance(custom_incremental.kind, DbtCustomKind) assert custom_incremental.kind.materialization == "custom_incremental" - custom_with_filter = sqlmesh_context.get_model("sushi.custom_incremental_with_filter") + custom_with_filter = sqlmesh_context.get_model( + "sushi.custom_incremental_with_filter" + ) assert isinstance(custom_with_filter.kind, DbtCustomKind) assert custom_with_filter.kind.materialization == "custom_incremental" @@ -294,7 +298,9 @@ def test_adapter_specific_materialization_override(copy_to_temp_path: t.Callable sushi_context.apply(plan) # check that the table was created with the correct adapter type - result = sushi_context.engine_adapter.fetchdf("SELECT * FROM sushi.test_adapter_specific") + result = sushi_context.engine_adapter.fetchdf( + "SELECT * FROM sushi.test_adapter_specific" + ) assert len(result) == 1 assert "adapter_type" in result.columns assert result["adapter_type"][0] == "duckdb_adapter" @@ -656,11 +662,15 @@ def test_custom_materialization_lineage_tracking(copy_to_temp_path: t.Callable): ) assert waiter_names_result["count"][0] > 0 - simple_a_result = context.engine_adapter.fetchdf("SELECT a FROM sushi.simple_model_a") + simple_a_result = context.engine_adapter.fetchdf( + "SELECT a FROM sushi.simple_model_a" + ) assert len(simple_a_result) > 0 assert simple_a_result["a"][0] == 1 - simple_b_result = context.engine_adapter.fetchdf("SELECT a FROM sushi.simple_model_b") + simple_b_result = context.engine_adapter.fetchdf( + "SELECT a FROM sushi.simple_model_b" + ) assert len(simple_b_result) > 0 assert simple_b_result["a"][0] == 1 @@ -696,7 +706,9 @@ def test_custom_materialization_lineage_tracking(copy_to_temp_path: t.Callable): assert len(analytics_summary_result) > 0 assert analytics_summary_result["model_type"][0] == "downstream_lineage_test" - assert all(cat in ["High", "Medium", "Low"] for cat in analytics_summary_result["category"]) + assert all( + cat in ["High", "Medium", "Low"] for cat in analytics_summary_result["category"] + ) assert all(val >= 0 for val in analytics_summary_result["final_computation"]) # Test that lineage information is preserved in dev environments @@ -719,7 +731,10 @@ def test_custom_materialization_lineage_tracking(copy_to_temp_path: t.Callable): # Dev and prod should have the same data as they share physical data assert dev_analytics_result["count"][0] == prod_analytics_result["count"][0] - assert dev_analytics_result["unique_waiters"][0] == prod_analytics_result["unique_waiters"][0] + assert ( + dev_analytics_result["unique_waiters"][0] + == prod_analytics_result["unique_waiters"][0] + ) @pytest.mark.xdist_group("dbt_manifest") @@ -748,14 +763,18 @@ def test_custom_materialization_grants(copy_to_temp_path: t.Callable, mocker): (models_dir / "test_grants_model.sql").write_text(grants_model_content) mocker.patch.object(DuckDBEngineAdapter, "SUPPORTS_GRANTS", True) - mocker.patch.object(DuckDBEngineAdapter, "_get_current_grants_config", return_value={}) + mocker.patch.object( + DuckDBEngineAdapter, "_get_current_grants_config", return_value={} + ) sync_grants_calls = [] def mock_sync_grants(*args, **kwargs): sync_grants_calls.append((args, kwargs)) - mocker.patch.object(DuckDBEngineAdapter, "sync_grants_config", side_effect=mock_sync_grants) + mocker.patch.object( + DuckDBEngineAdapter, "sync_grants_config", side_effect=mock_sync_grants + ) context = Context(paths=path) diff --git a/tests/dbt/test_docs.py b/tests/dbt/test_docs.py index 7c21edb970..b9a57d0f59 100644 --- a/tests/dbt/test_docs.py +++ b/tests/dbt/test_docs.py @@ -1,4 +1,5 @@ from pathlib import Path + import pytest from sqlmesh.core.config.model import ModelDefaultsConfig @@ -6,7 +7,6 @@ from sqlmesh.dbt.manifest import ManifestHelper from sqlmesh.dbt.profile import Profile - pytestmark = pytest.mark.dbt diff --git a/tests/dbt/test_integration.py b/tests/dbt/test_integration.py index ab22bf7826..6cf31c09ce 100644 --- a/tests/dbt/test_integration.py +++ b/tests/dbt/test_integration.py @@ -19,7 +19,9 @@ from sqlmesh.core.config.connection import DuckDBConnectionConfig from sqlmesh.core.engine_adapter import DuckDBEngineAdapter from sqlmesh.utils.pandas import columns_to_types_from_df -from sqlmesh.utils.yaml import YAML, load as yaml_load, dump as yaml_dump +from sqlmesh.utils.yaml import YAML +from sqlmesh.utils.yaml import dump as yaml_dump +from sqlmesh.utils.yaml import load as yaml_load from sqlmesh_dbt.operations import init_project_if_required from tests.utils.pandas import compare_dataframes, create_df @@ -71,7 +73,11 @@ def is_check(self) -> bool: class TestSCDType2: - source_schema = {"customer_id": "int32", "status": "object", "updated_at": "datetime64[us]"} + source_schema = { + "customer_id": "int32", + "status": "object", + "updated_at": "datetime64[us]", + } target_schema = { "customer_id": "int32", "status": "object", @@ -134,7 +140,9 @@ def _make_function( dbt_data_file = dbt_data_dir / "local.db" dbt_profile_config = { "test": { - "outputs": {"duckdb": {"type": "duckdb", "path": str(dbt_data_file)}}, + "outputs": { + "duckdb": {"type": "duckdb", "path": str(dbt_data_file)} + }, "target": "duckdb", } } @@ -144,16 +152,14 @@ def _make_function( if include_dbt_adapter_support: sqlmesh_config_file = dbt_project_dir / "config.py" with open(sqlmesh_config_file, "w", encoding="utf-8") as f: - f.write( - """from pathlib import Path + f.write("""from pathlib import Path from sqlmesh.dbt.loader import sqlmesh_config config = sqlmesh_config(Path(__file__).parent) -test_config = config""" - ) +test_config = config""") return dbt_project_dir, dbt_data_file return _make_function @@ -186,7 +192,9 @@ def _make_function( updated_at::TIMESTAMP AS updated_at FROM sushi.raw_marketing""" - with open(project_root / "models" / "marketing.sql", "w", encoding="utf-8") as f: + with open( + project_root / "models" / "marketing.sql", "w", encoding="utf-8" + ) as f: f.write(snapshot_def) data_dir = tmp_path / "sqlm_data" data_dir.mkdir() @@ -208,7 +216,9 @@ def _replace_source_table( "sushi.raw_marketing", df, target_columns_to_types=columns_to_types ) else: - adapter.create_table("sushi.raw_marketing", target_columns_to_types=columns_to_types) + adapter.create_table( + "sushi.raw_marketing", target_columns_to_types=columns_to_types + ) def _normalize_dbt_dataframe( self, @@ -228,7 +238,9 @@ def update_now_column(col_timestamp: pd.Timestamp) -> datetime.datetime: col_datetime = pd.to_datetime(now_time).tz_localize(None) return col_datetime - df = df.rename(columns={"dbt_valid_from": "valid_from", "dbt_valid_to": "valid_to"}) + df = df.rename( + columns={"dbt_valid_from": "valid_from", "dbt_valid_to": "valid_to"} + ) if test_type.is_dbt_runtime: df = df.drop(columns=["dbt_updated_at", "dbt_scd_id"]) for col in ["valid_from", "valid_to"]: @@ -265,7 +277,8 @@ def _init_test( ): project_dir, data_file = ( create_scd_type_2_sqlmesh_project( - test_strategy=test_strategy, invalidate_hard_deletes=invalidate_hard_deletes + test_strategy=test_strategy, + invalidate_hard_deletes=invalidate_hard_deletes, ) if test_type.is_sqlmesh else create_scd_type_2_dbt_project( @@ -314,7 +327,9 @@ def test_scd_type_2_by_time( invalidate_hard_deletes: bool, ): if test_type.is_dbt_runtime and DBT_VERSION < (1, 5, 0): - pytest.skip("The dbt version being tested doesn't support the dbtRunner so skipping.") + pytest.skip( + "The dbt version being tested doesn't support the dbtRunner so skipping." + ) run, adapter, context = self._init_test( create_scd_type_2_dbt_project, @@ -327,7 +342,8 @@ def test_scd_type_2_by_time( time_expected_mapping: t.Dict[ str, t.Tuple[ - t.List[t.Tuple[int, str, str]], t.List[t.Tuple[int, str, str, str, t.Optional[str]]] + t.List[t.Tuple[int, str, str]], + t.List[t.Tuple[int, str, str, str, t.Optional[str]]], ], ] = { "2020-01-01 00:00:00 UTC": ( @@ -354,7 +370,13 @@ def test_scd_type_2_by_time( (4, "d", "2020-01-02 00:00:00"), ], [ - (1, "a", "2020-01-01 00:00:00", "2020-01-01 00:00:00", "2020-01-02 00:00:00"), + ( + 1, + "a", + "2020-01-01 00:00:00", + "2020-01-01 00:00:00", + "2020-01-02 00:00:00", + ), (1, "x", "2020-01-02 00:00:00", "2020-01-02 00:00:00", None), (2, "b", "2020-01-01 00:00:00", "2020-01-01 00:00:00", None), ( @@ -381,8 +403,20 @@ def test_scd_type_2_by_time( (5, "e", "2020-01-03 00:00:00"), ], [ - (1, "a", "2020-01-01 00:00:00", "2020-01-01 00:00:00", "2020-01-02 00:00:00"), - (1, "x", "2020-01-02 00:00:00", "2020-01-02 00:00:00", "2020-01-03 00:00:00"), + ( + 1, + "a", + "2020-01-01 00:00:00", + "2020-01-01 00:00:00", + "2020-01-02 00:00:00", + ), + ( + 1, + "x", + "2020-01-02 00:00:00", + "2020-01-02 00:00:00", + "2020-01-03 00:00:00", + ), (1, "y", "2020-01-03 00:00:00", "2020-01-03 00:00:00", None), ( 2, @@ -396,7 +430,11 @@ def test_scd_type_2_by_time( "c", "2020-01-01 00:00:00", "2020-01-01 00:00:00", - "2020-01-02 00:00:00" if invalidate_hard_deletes else "2020-02-01 00:00:00", + ( + "2020-01-02 00:00:00" + if invalidate_hard_deletes + else "2020-02-01 00:00:00" + ), ), # Since 3 was deleted and came back and the updated at time when it came back # is greater than the execution time when it was deleted, we have the valid_from @@ -411,7 +449,10 @@ def test_scd_type_2_by_time( ), } time_start_end_mapping = {} - for time, (starting_source_data, expected_table_data) in time_expected_mapping.items(): + for time, ( + starting_source_data, + expected_table_data, + ) in time_expected_mapping.items(): self._replace_source_table(adapter, starting_source_data) # Tick when running dbt runtime because it hangs during execution for unknown reasons. with time_machine.travel(time, tick=test_type.is_dbt_runtime): @@ -420,7 +461,9 @@ def test_scd_type_2_by_time( end_time = self._get_duckdb_now(adapter) time_start_end_mapping[time] = (start_time, end_time) df_actual = self._get_current_df( - adapter, test_type=test_type, time_start_end_mapping=time_start_end_mapping + adapter, + test_type=test_type, + time_start_end_mapping=time_start_end_mapping, ) df_expected = create_df(expected_table_data, self.target_schema) compare_dataframes(df_actual, df_expected, msg=f"Failed on time {time}") @@ -449,7 +492,8 @@ def test_scd_type_2_by_column( time_expected_mapping: t.Dict[ str, t.Tuple[ - t.List[t.Tuple[int, str, str]], t.List[t.Tuple[int, str, str, str, t.Optional[str]]] + t.List[t.Tuple[int, str, str]], + t.List[t.Tuple[int, str, str, str, t.Optional[str]]], ], ] = { "2020-01-01 00:00:00 UTC": ( @@ -476,7 +520,13 @@ def test_scd_type_2_by_column( (4, "d", "2020-01-02 00:00:00"), ], [ - (1, "a", "2020-01-01 00:00:00", "2020-01-01 00:00:00", "2020-01-02 00:00:00"), + ( + 1, + "a", + "2020-01-01 00:00:00", + "2020-01-01 00:00:00", + "2020-01-02 00:00:00", + ), (1, "x", "2020-01-02 00:00:00", "2020-01-02 00:00:00", None), (2, "b", "2020-01-01 00:00:00", "2020-01-01 00:00:00", None), ( @@ -503,8 +553,20 @@ def test_scd_type_2_by_column( (5, "e", "2020-01-03 00:00:00"), ], [ - (1, "a", "2020-01-01 00:00:00", "2020-01-01 00:00:00", "2020-01-02 00:00:00"), - (1, "x", "2020-01-02 00:00:00", "2020-01-02 00:00:00", "2020-01-04 00:00:00"), + ( + 1, + "a", + "2020-01-01 00:00:00", + "2020-01-01 00:00:00", + "2020-01-02 00:00:00", + ), + ( + 1, + "x", + "2020-01-02 00:00:00", + "2020-01-02 00:00:00", + "2020-01-04 00:00:00", + ), (1, "y", "2020-01-03 00:00:00", "2020-01-04 00:00:00", None), ( 2, @@ -518,7 +580,11 @@ def test_scd_type_2_by_column( "c", "2020-01-01 00:00:00", "2020-01-01 00:00:00", - "2020-01-02 00:00:00" if invalidate_hard_deletes else "2020-01-04 00:00:00", + ( + "2020-01-02 00:00:00" + if invalidate_hard_deletes + else "2020-01-04 00:00:00" + ), ), # Since 3 was deleted and came back then the valid_from is set to the execution_time when it # came back. @@ -529,7 +595,10 @@ def test_scd_type_2_by_column( ), } time_start_end_mapping = {} - for time, (starting_source_data, expected_table_data) in time_expected_mapping.items(): + for time, ( + starting_source_data, + expected_table_data, + ) in time_expected_mapping.items(): self._replace_source_table(adapter, starting_source_data) with time_machine.travel(time, tick=False): start_time = self._get_duckdb_now(adapter) @@ -537,7 +606,9 @@ def test_scd_type_2_by_column( end_time = self._get_duckdb_now(adapter) time_start_end_mapping[time] = (start_time, end_time) df_actual = self._get_current_df( - adapter, test_type=test_type, time_start_end_mapping=time_start_end_mapping + adapter, + test_type=test_type, + time_start_end_mapping=time_start_end_mapping, ) df_expected = create_df(expected_table_data, self.target_schema) compare_dataframes(df_actual, df_expected, msg=f"Failed on time {time}") @@ -624,14 +695,18 @@ def test_state_schema_isolation_per_target(jaffle_shop_duckdb: Path): init_project_if_required(jaffle_shop_duckdb) # start off with the prod target - prod_ctx = Context(paths=[jaffle_shop_duckdb], config_loader_kwargs={"target": "prod"}) + prod_ctx = Context( + paths=[jaffle_shop_duckdb], config_loader_kwargs={"target": "prod"} + ) assert prod_ctx.config.get_state_schema() == "sqlmesh_state_jaffle_shop_prod_schema" assert all("prod_schema" in fqn for fqn in prod_ctx.models) assert prod_ctx.plan(auto_apply=True).has_changes assert not prod_ctx.plan(auto_apply=True).has_changes # dev target should have changes - new state separate from prod - dev_ctx = Context(paths=[jaffle_shop_duckdb], config_loader_kwargs={"target": "dev"}) + dev_ctx = Context( + paths=[jaffle_shop_duckdb], config_loader_kwargs={"target": "dev"} + ) assert dev_ctx.config.get_state_schema() == "sqlmesh_state_jaffle_shop_dev_schema" assert all("dev_schema" in fqn for fqn in dev_ctx.models) assert dev_ctx.plan(auto_apply=True).has_changes @@ -640,7 +715,9 @@ def test_state_schema_isolation_per_target(jaffle_shop_duckdb: Path): # no explicitly specified target should use dev because that's what's set for the default in the profiles.yml assert profiles_yml["jaffle_shop"]["target"] == "dev" default_ctx = Context(paths=[jaffle_shop_duckdb]) - assert default_ctx.config.get_state_schema() == "sqlmesh_state_jaffle_shop_dev_schema" + assert ( + default_ctx.config.get_state_schema() == "sqlmesh_state_jaffle_shop_dev_schema" + ) assert all("dev_schema" in fqn for fqn in default_ctx.models) assert not default_ctx.plan(auto_apply=True).has_changes diff --git a/tests/dbt/test_manifest.py b/tests/dbt/test_manifest.py index 2ecf8b8980..dc675d420e 100644 --- a/tests/dbt/test_manifest.py +++ b/tests/dbt/test_manifest.py @@ -6,11 +6,11 @@ from sqlmesh.core.config import ModelDefaultsConfig from sqlmesh.dbt.basemodel import Dependencies +from sqlmesh.dbt.builtin import Api, _relation_info_to_relation from sqlmesh.dbt.common import ModelAttrs from sqlmesh.dbt.context import DbtContext from sqlmesh.dbt.manifest import ManifestHelper, _convert_jinja_test_to_macro from sqlmesh.dbt.profile import Profile -from sqlmesh.dbt.builtin import Api, _relation_info_to_relation from sqlmesh.dbt.util import DBT_VERSION from sqlmesh.utils.jinja import MacroReference @@ -45,7 +45,10 @@ def test_manifest_helper(caplog): assert models["top_waiters"].dialect_ == "postgres" assert models["waiters"].dependencies == Dependencies( - macros={MacroReference(name="incremental_by_time"), MacroReference(name="source")}, + macros={ + MacroReference(name="incremental_by_time"), + MacroReference(name="source"), + }, sources={"streaming.orders"}, ) assert models["waiters"].materialized == "ephemeral" @@ -81,7 +84,10 @@ def test_manifest_helper(caplog): macros=[MacroReference(name="ref")], ) assert waiter_as_customer_by_day_config.materialized == "incremental" - assert waiter_as_customer_by_day_config.incremental_strategy == "incremental_by_time_range" + assert ( + waiter_as_customer_by_day_config.incremental_strategy + == "incremental_by_time_range" + ) assert waiter_as_customer_by_day_config.cluster_by == ["ds"] assert waiter_as_customer_by_day_config.time_column == "ds" @@ -103,7 +109,9 @@ def test_manifest_helper(caplog): has_dynamic_var_names=True, ) assert waiter_revenue_by_day_config.materialized == "incremental" - assert waiter_revenue_by_day_config.incremental_strategy == "incremental_by_time_range" + assert ( + waiter_revenue_by_day_config.incremental_strategy == "incremental_by_time_range" + ) assert waiter_revenue_by_day_config.cluster_by == ["ds"] assert waiter_revenue_by_day_config.time_column == "ds" assert waiter_revenue_by_day_config.dialect_ == "bigquery" @@ -132,8 +140,14 @@ def test_manifest_helper(caplog): assert all(s.quoting["identifier"] is False for s in sources.values()) assert sources["streaming.order_items"].freshness == { - "warn_after": {"count": 10 if DBT_VERSION < (1, 9, 5) else 12, "period": "hour"}, - "error_after": {"count": 11 if DBT_VERSION < (1, 9, 5) else 13, "period": "hour"}, + "warn_after": { + "count": 10 if DBT_VERSION < (1, 9, 5) else 12, + "period": "hour", + }, + "error_after": { + "count": 11 if DBT_VERSION < (1, 9, 5) else 13, + "period": "hour", + }, "filter": None, } @@ -218,7 +232,8 @@ def test_source_meta_external_location(): "external_location": "read_parquet('path/to/external/{name}.parquet')" } assert ( - parquet_orders.relation_info.external == "read_parquet('path/to/external/orders.parquet')" + parquet_orders.relation_info.external + == "read_parquet('path/to/external/orders.parquet')" ) api = Api("duckdb") diff --git a/tests/dbt/test_model.py b/tests/dbt/test_model.py index a954f98f41..a4a0827194 100644 --- a/tests/dbt/test_model.py +++ b/tests/dbt/test_model.py @@ -1,27 +1,27 @@ import datetime import logging - -import pytest - +import typing as t from pathlib import Path +import pytest from sqlglot import exp from sqlglot.errors import SchemaError + from sqlmesh import Context -from sqlmesh.core.console import NoopConsole, get_console -from sqlmesh.core.model import TimeColumn, IncrementalByTimeRangeKind -from sqlmesh.core.model.kind import OnDestructiveChange, OnAdditiveChange, SCDType2ByColumnKind -from sqlmesh.core.state_sync.db.snapshot import _snapshot_to_json from sqlmesh.core.config.common import VirtualEnvironmentMode +from sqlmesh.core.console import NoopConsole, get_console +from sqlmesh.core.model import IncrementalByTimeRangeKind, TimeColumn +from sqlmesh.core.model.kind import (OnAdditiveChange, OnDestructiveChange, + SCDType2ByColumnKind) from sqlmesh.core.model.meta import GrantsTargetLayer +from sqlmesh.core.state_sync.db.snapshot import _snapshot_to_json from sqlmesh.dbt.common import Dependencies from sqlmesh.dbt.context import DbtContext from sqlmesh.dbt.model import ModelConfig from sqlmesh.dbt.target import BigQueryConfig, DuckDbConfig, PostgresConfig from sqlmesh.dbt.test import TestConfig -from sqlmesh.utils.yaml import YAML from sqlmesh.utils.date import to_ds -import typing as t +from sqlmesh.utils.yaml import YAML pytestmark = pytest.mark.dbt @@ -79,11 +79,15 @@ def test_test_to_sqlmesh_creates_correct_audit_type( dbt_dummy_postgres_config: PostgresConfig, ) -> None: """Test that TestConfig.to_sqlmesh creates the correct audit type based on is_standalone""" - from sqlmesh.core.audit.definition import StandaloneAudit, ModelAudit + from sqlmesh.core.audit.definition import ModelAudit, StandaloneAudit # Set up models in context my_model = ModelConfig( - name="my_model", sql="SELECT 1", schema="test_schema", database="test_db", alias="my_model" + name="my_model", + sql="SELECT 1", + schema="test_schema", + database="test_db", + alias="my_model", ) other_model = ModelConfig( name="other_model", @@ -225,7 +229,10 @@ def test_manifest_filters_standalone_tests_from_models( @pytest.mark.slow def test_load_invalid_ref_audit_constraints( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: yaml = YAML() project_dir, model_dir = create_empty_project(project_name="local") @@ -295,7 +302,10 @@ def test_load_invalid_ref_audit_constraints( @pytest.mark.slow def test_load_microbatch_all_defined( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: project_dir, model_dir = create_empty_project(project_name="local") # add `tests` to model config since this is loaded by dbt and ignored and we shouldn't error when loading it @@ -336,7 +346,10 @@ def test_load_microbatch_all_defined( @pytest.mark.slow def test_load_microbatch_all_defined_diff_values( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: project_dir, model_dir = create_empty_project(project_name="local") # add `tests` to model config since this is loaded by dbt and ignored and we shouldn't error when loading it @@ -378,7 +391,10 @@ def test_load_microbatch_all_defined_diff_values( @pytest.mark.slow def test_load_microbatch_required_only( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: project_dir, model_dir = create_empty_project(project_name="local") # add `tests` to model config since this is loaded by dbt and ignored and we shouldn't error when loading it @@ -417,9 +433,14 @@ def test_load_microbatch_required_only( @pytest.mark.slow def test_load_incremental_time_range_strategy_required_only( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: - project_dir, model_dir = create_empty_project(project_name="local", start="2025-01-01") + project_dir, model_dir = create_empty_project( + project_name="local", start="2025-01-01" + ) # add `tests` to model config since this is loaded by dbt and ignored and we shouldn't error when loading it incremental_time_range_contents = """ {{ @@ -459,9 +480,14 @@ def test_load_incremental_time_range_strategy_required_only( @pytest.mark.slow def test_load_incremental_time_range_strategy_all_defined( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: - project_dir, model_dir = create_empty_project(project_name="local", start="2025-01-01") + project_dir, model_dir = create_empty_project( + project_name="local", start="2025-01-01" + ) # add `tests` to model config since this is loaded by dbt and ignored and we shouldn't error when loading it incremental_time_range_contents = """ {{ @@ -522,9 +548,14 @@ def test_load_incremental_time_range_strategy_all_defined( @pytest.mark.slow def test_load_deprecated_incremental_time_column( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: - project_dir, model_dir = create_empty_project(project_name="local", start="2025-01-01") + project_dir, model_dir = create_empty_project( + project_name="local", start="2025-01-01" + ) # add `tests` to model config since this is loaded by dbt and ignored and we shouldn't error when loading it incremental_time_range_contents = """ {{ @@ -570,7 +601,10 @@ def test_load_deprecated_incremental_time_column( @pytest.mark.slow def test_load_microbatch_with_ref( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: yaml = YAML() project_dir, model_dir = create_empty_project(project_name="local") @@ -625,18 +659,25 @@ def test_load_microbatch_with_ref( microbatch_two_snapshot_fqn = '"local"."main"."microbatch_two"' context = Context(paths=project_dir) assert ( - context.render(microbatch_snapshot_fqn, start="2025-01-01", end="2025-01-10").sql() + context.render( + microbatch_snapshot_fqn, start="2025-01-01", end="2025-01-10" + ).sql() == 'SELECT "cola" AS "cola", "ds_source" AS "ds" FROM (SELECT * FROM "local"."my_source"."my_table" AS "my_table" WHERE "ds_source" >= \'2025-01-01 00:00:00+00:00\' AND "ds_source" < \'2025-01-11 00:00:00+00:00\') AS "_0"' ) assert ( - context.render(microbatch_two_snapshot_fqn, start="2025-01-01", end="2025-01-10").sql() + context.render( + microbatch_two_snapshot_fqn, start="2025-01-01", end="2025-01-10" + ).sql() == 'SELECT "_0"."cola" AS "cola", "_0"."ds" AS "ds" FROM (SELECT "microbatch"."cola" AS "cola", "microbatch"."ds" AS "ds" FROM "local"."main"."microbatch" AS "microbatch" WHERE "microbatch"."ds" < \'2025-01-11 00:00:00+00:00\' AND "microbatch"."ds" >= \'2025-01-01 00:00:00+00:00\') AS "_0"' ) @pytest.mark.slow def test_load_microbatch_with_ref_no_filter( - tmp_path: Path, caplog, dbt_dummy_postgres_config: PostgresConfig, create_empty_project + tmp_path: Path, + caplog, + dbt_dummy_postgres_config: PostgresConfig, + create_empty_project, ) -> None: yaml = YAML() project_dir, model_dir = create_empty_project(project_name="local") @@ -691,17 +732,23 @@ def test_load_microbatch_with_ref_no_filter( microbatch_two_snapshot_fqn = '"local"."main"."microbatch_two"' context = Context(paths=project_dir) assert ( - context.render(microbatch_snapshot_fqn, start="2025-01-01", end="2025-01-10").sql() + context.render( + microbatch_snapshot_fqn, start="2025-01-01", end="2025-01-10" + ).sql() == 'SELECT "cola" AS "cola", "ds" AS "ds" FROM "local"."my_source"."my_table" AS "my_table"' ) assert ( - context.render(microbatch_two_snapshot_fqn, start="2025-01-01", end="2025-01-10").sql() + context.render( + microbatch_two_snapshot_fqn, start="2025-01-01", end="2025-01-10" + ).sql() == 'SELECT "microbatch"."cola" AS "cola", "microbatch"."ds" AS "ds" FROM "local"."main"."microbatch" AS "microbatch"' ) @pytest.mark.slow -def test_load_multiple_snapshots_defined_in_same_file(sushi_test_dbt_context: Context) -> None: +def test_load_multiple_snapshots_defined_in_same_file( + sushi_test_dbt_context: Context, +) -> None: context = sushi_test_dbt_context assert context.get_model("snapshots.items_snapshot") assert context.get_model("snapshots.items_check_snapshot") @@ -713,7 +760,9 @@ def test_load_multiple_snapshots_defined_in_same_file(sushi_test_dbt_context: Co @pytest.mark.slow -def test_dbt_snapshot_with_check_cols_expressions(sushi_test_dbt_context: Context) -> None: +def test_dbt_snapshot_with_check_cols_expressions( + sushi_test_dbt_context: Context, +) -> None: context = sushi_test_dbt_context model = context.get_model("snapshots.items_check_with_cast_snapshot") assert model is not None @@ -770,11 +819,16 @@ def test_dbt_jinja_macro_undefined_variable_error(create_empty_project): error_message = str(exc_info.value) assert "Failed to update model schemas" in error_message assert "Could not render jinja for" in error_message - assert "Undefined macro/variable: 'columns' in macro: 'select_columns'" in error_message + assert ( + "Undefined macro/variable: 'columns' in macro: 'select_columns'" + in error_message + ) @pytest.mark.slow -def test_node_name_populated_for_dbt_models(dbt_dummy_postgres_config: PostgresConfig) -> None: +def test_node_name_populated_for_dbt_models( + dbt_dummy_postgres_config: PostgresConfig, +) -> None: model_config = ModelConfig( unique_id="model.test_package.test_model", fqn=["test_package", "test_model"], @@ -874,7 +928,9 @@ def test_jinja_config_no_query(create_empty_project): context = Context(paths=project_dir) # loads without error and contains empty query (which will error at runtime) - assert not context.snapshots['"local"."main"."comment_config_model"'].model.render_query() + assert not context.snapshots[ + '"local"."main"."comment_config_model"' + ].model.render_query() @pytest.mark.slow @@ -1058,7 +1114,9 @@ def test_ephemeral_model_ignores_grants() -> None: ) assert sqlmesh_model.kind.is_embedded - assert sqlmesh_model.grants is None # grants config is skipped for ephemeral / embedded models + assert ( + sqlmesh_model.grants is None + ) # grants config is skipped for ephemeral / embedded models def test_conditional_ref_in_unexecuted_branch(copy_to_temp_path: t.Callable): @@ -1108,9 +1166,15 @@ def test_conditional_ref_in_unexecuted_branch(copy_to_temp_path: t.Callable): ) # And run plan with this conditional model for good measure - plan = sushi_context.plan(select_models=["sushi.conditional_ref_model", "sushi.simple_model_a"]) + plan = sushi_context.plan( + select_models=["sushi.conditional_ref_model", "sushi.simple_model_a"] + ) sushi_context.apply(plan) - upstream_ref = sushi_context.engine_adapter.fetchone("SELECT * FROM sushi.simple_model_a") + upstream_ref = sushi_context.engine_adapter.fetchone( + "SELECT * FROM sushi.simple_model_a" + ) assert upstream_ref == (1,) - result = sushi_context.engine_adapter.fetchone("SELECT * FROM sushi.conditional_ref_model") + result = sushi_context.engine_adapter.fetchone( + "SELECT * FROM sushi.conditional_ref_model" + ) assert result == (1,) diff --git a/tests/dbt/test_test.py b/tests/dbt/test_test.py index fb33220c0c..23d6f38633 100644 --- a/tests/dbt/test_test.py +++ b/tests/dbt/test_test.py @@ -16,8 +16,8 @@ def test_multiline_test_kwarg() -> None: @pytest.mark.xdist_group("dbt_manifest") def test_tests_get_unique_names(tmp_path: Path, create_empty_project) -> None: - from sqlmesh.utils.yaml import YAML from sqlmesh.core.context import Context + from sqlmesh.utils.yaml import YAML yaml = YAML() project_dir, model_dir = create_empty_project(project_name="local") @@ -66,7 +66,11 @@ def test_tests_get_unique_names(tmp_path: Path, create_empty_project) -> None: "name": "status", "data_tests": [ {"accepted_values": {"values": ["value1", "value2"]}}, - {"accepted_values": {"values": ["value1", "value2", "value3"]}}, + { + "accepted_values": { + "values": ["value1", "value2", "value3"] + } + }, { "accepted_values": { "name": "custom_accepted_values_name", @@ -123,7 +127,9 @@ def test_tests_get_unique_names(tmp_path: Path, create_empty_project) -> None: context = Context(paths=project_dir) - all_audit_names = list(context._audits.keys()) + list(context._standalone_audits.keys()) + all_audit_names = list(context._audits.keys()) + list( + context._standalone_audits.keys() + ) assert sorted(all_audit_names) == [ "local.accepted_values_my_model_status__value1__value2", "local.accepted_values_my_model_status__value1__value2__value3", diff --git a/tests/dbt/test_transformation.py b/tests/dbt/test_transformation.py index fe6073dfad..0b1474af0e 100644 --- a/tests/dbt/test_transformation.py +++ b/tests/dbt/test_transformation.py @@ -1,69 +1,53 @@ -import agate -from datetime import datetime, timedelta import json import logging import typing as t +from datetime import datetime, timedelta from pathlib import Path from unittest.mock import patch -from sqlmesh.dbt.util import DBT_VERSION - +import agate import pytest from dbt.adapters.base import BaseRelation from jinja2 import Template +from sqlmesh.dbt.util import DBT_VERSION + if DBT_VERSION >= (1, 4, 0): from dbt.exceptions import CompilationError else: from dbt.exceptions import CompilationException as CompilationError # type: ignore + import time_machine from pytest_mock.plugin import MockerFixture from sqlglot import exp, parse_one + from sqlmesh.core import dialect as d +from sqlmesh.core.audit import StandaloneAudit +from sqlmesh.core.console import get_console +from sqlmesh.core.context import Context from sqlmesh.core.environment import EnvironmentNamingInfo from sqlmesh.core.macros import RuntimeStage +from sqlmesh.core.model import (EmbeddedKind, FullKind, + IncrementalByTimeRangeKind, + IncrementalByUniqueKeyKind, + IncrementalUnmanagedKind, ManagedKind, + SqlModel, ViewKind) +from sqlmesh.core.model.kind import (OnAdditiveChange, OnDestructiveChange, + SCDType2ByColumnKind, SCDType2ByTimeKind) from sqlmesh.core.renderer import render_statements -from sqlmesh.core.audit import StandaloneAudit -from sqlmesh.core.context import Context -from sqlmesh.core.console import get_console -from sqlmesh.core.model import ( - EmbeddedKind, - FullKind, - IncrementalByTimeRangeKind, - IncrementalByUniqueKeyKind, - IncrementalUnmanagedKind, - ManagedKind, - SqlModel, - ViewKind, -) -from sqlmesh.core.model.kind import ( - SCDType2ByColumnKind, - SCDType2ByTimeKind, - OnDestructiveChange, - OnAdditiveChange, -) from sqlmesh.core.state_sync.db.snapshot import _snapshot_to_json -from sqlmesh.dbt.builtin import _relation_info_to_relation, Config +from sqlmesh.dbt.builtin import Config, _relation_info_to_relation +from sqlmesh.dbt.column import (ColumnConfig, column_descriptions_to_sqlmesh, + column_types_to_sqlmesh) from sqlmesh.dbt.common import Dependencies -from sqlmesh.dbt.builtin import _relation_info_to_relation -from sqlmesh.dbt.column import ( - ColumnConfig, - column_descriptions_to_sqlmesh, - column_types_to_sqlmesh, -) from sqlmesh.dbt.context import DbtContext from sqlmesh.dbt.model import Materialization, ModelConfig -from sqlmesh.dbt.source import SourceConfig from sqlmesh.dbt.project import Project from sqlmesh.dbt.relation import Policy from sqlmesh.dbt.seed import SeedConfig -from sqlmesh.dbt.target import ( - BigQueryConfig, - DuckDbConfig, - SnowflakeConfig, - ClickhouseConfig, - PostgresConfig, -) +from sqlmesh.dbt.source import SourceConfig +from sqlmesh.dbt.target import (BigQueryConfig, ClickhouseConfig, DuckDbConfig, + PostgresConfig, SnowflakeConfig) from sqlmesh.dbt.test import TestConfig from sqlmesh.utils.errors import ConfigError, SQLMeshError from sqlmesh.utils.jinja import MacroReference @@ -74,9 +58,14 @@ def test_model_name(dbt_dummy_postgres_config: PostgresConfig): context = DbtContext() context._target = dbt_dummy_postgres_config - assert ModelConfig(schema="foo", path="models/bar.sql").canonical_name(context) == "foo.bar" assert ( - ModelConfig(schema="foo", path="models/bar.sql", alias="baz").canonical_name(context) + ModelConfig(schema="foo", path="models/bar.sql").canonical_name(context) + == "foo.bar" + ) + assert ( + ModelConfig(schema="foo", path="models/bar.sql", alias="baz").canonical_name( + context + ) == "foo.baz" ) assert ( @@ -100,7 +89,10 @@ def test_materialization(): with patch.object(get_console(), "log_warning") as mock_logger: model_config = ModelConfig( - name="model", alias="model", schema="schema", materialized="materialized_view" + name="model", + alias="model", + schema="schema", + materialized="materialized_view", ) assert ( @@ -111,13 +103,17 @@ def test_materialization(): # clickhouse "dictionary" materialization with pytest.raises(ConfigError): - ModelConfig(name="model", alias="model", schema="schema", materialized="dictionary") + ModelConfig( + name="model", alias="model", schema="schema", materialized="dictionary" + ) def test_dbt_custom_materialization(): sushi_context = Context(paths=["tests/fixtures/dbt/sushi_test"]) - plan_builder = sushi_context.plan_builder(select_models=["sushi.custom_incremental_model"]) + plan_builder = sushi_context.plan_builder( + select_models=["sushi.custom_incremental_model"] + ) plan = plan_builder.build() assert len(plan.selected_models) == 1 selected_model = list(plan.selected_models)[0] @@ -139,7 +135,9 @@ def test_dbt_custom_materialization(): # running with execution time one day in the future to simulate an incremental insert tomorrow = datetime.now() + timedelta(days=1) - sushi_context.run(select_models=["sushi.custom_incremental_model"], execution_time=tomorrow) + sushi_context.run( + select_models=["sushi.custom_incremental_model"], execution_time=tomorrow + ) result_after_run = sushi_context.engine_adapter.fetchdf(query) assert {"created_at", "id"}.issubset(result_after_run.columns) @@ -165,7 +163,9 @@ def test_dbt_custom_materialization_with_time_filter_and_macro(): # select both custom materialiasation models with the wildcard selector = ["sushi.custom_incremental*"] - plan_builder = sushi_context.plan_builder(select_models=selector, execution_time=today) + plan_builder = sushi_context.plan_builder( + select_models=selector, execution_time=today + ) plan = plan_builder.build() assert len(plan.selected_models) == 2 @@ -178,7 +178,9 @@ def test_dbt_custom_materialization_with_time_filter_and_macro(): select_daily = "SELECT * FROM sushi.custom_incremental_model ORDER BY created_at" # this model uses `run_started_at` as a filter (which we populate with execution time) with 2 day interval - select_filter = "SELECT * FROM sushi.custom_incremental_with_filter ORDER BY created_at" + select_filter = ( + "SELECT * FROM sushi.custom_incremental_with_filter ORDER BY created_at" + ) sushi_context.apply(plan) result = sushi_context.engine_adapter.fetchdf(select_daily) @@ -228,7 +230,9 @@ def test_dbt_custom_materialization_with_time_filter_and_macro(): assert result_after_run_filter["created_at"][1].date() == two_days_later.date() # assert hooks have executed for both plan and incremental runs - hook_result = sushi_context.engine_adapter.fetchdf("SELECT * FROM hook_table ORDER BY id") + hook_result = sushi_context.engine_adapter.fetchdf( + "SELECT * FROM hook_table ORDER BY id" + ) assert len(hook_result) == 3 hook_result["id"][0] == 1 assert hook_result["id"].is_monotonic_increasing @@ -242,9 +246,17 @@ def test_model_kind(): context.project_name = "Test" context.target = DuckDbConfig(name="target", schema="foo") - assert ModelConfig(materialized=Materialization.TABLE).model_kind(context) == FullKind() - assert ModelConfig(materialized=Materialization.VIEW).model_kind(context) == ViewKind() - assert ModelConfig(materialized=Materialization.EPHEMERAL).model_kind(context) == EmbeddedKind() + assert ( + ModelConfig(materialized=Materialization.TABLE).model_kind(context) + == FullKind() + ) + assert ( + ModelConfig(materialized=Materialization.VIEW).model_kind(context) == ViewKind() + ) + assert ( + ModelConfig(materialized=Materialization.EPHEMERAL).model_kind(context) + == EmbeddedKind() + ) assert ModelConfig( materialized=Materialization.SNAPSHOT, unique_key=["id"], @@ -315,12 +327,14 @@ def test_model_kind(): assert isinstance(check_cols_multiple_expr.columns[0], exp.Cast) assert isinstance(check_cols_multiple_expr.columns[1], exp.Coalesce) - assert check_cols_multiple_expr.columns[0].sql() == 'CAST("created_at" AS TIMESTAMPTZ)' + assert ( + check_cols_multiple_expr.columns[0].sql() == 'CAST("created_at" AS TIMESTAMPTZ)' + ) assert check_cols_multiple_expr.columns[1].sql() == "COALESCE(\"status\", 'active')" - assert ModelConfig(materialized=Materialization.INCREMENTAL, time_column="foo").model_kind( - context - ) == IncrementalByTimeRangeKind( + assert ModelConfig( + materialized=Materialization.INCREMENTAL, time_column="foo" + ).model_kind(context) == IncrementalByTimeRangeKind( time_column="foo", dialect="duckdb", forward_only=True, @@ -363,7 +377,9 @@ def test_model_kind(): ) assert ModelConfig( - materialized=Materialization.INCREMENTAL, unique_key=["bar"], incremental_strategy="merge" + materialized=Materialization.INCREMENTAL, + unique_key=["bar"], + incremental_strategy="merge", ).model_kind(context) == IncrementalByUniqueKeyKind( unique_key=["bar"], dialect="duckdb", @@ -373,7 +389,9 @@ def test_model_kind(): on_additive_change=OnAdditiveChange.IGNORE, ) - dbt_incremental_predicate = "DBT_INTERNAL_DEST.session_start > dateadd(day, -7, current_date)" + dbt_incremental_predicate = ( + "DBT_INTERNAL_DEST.session_start > dateadd(day, -7, current_date)" + ) expected_sqlmesh_predicate = parse_one( "__MERGE_TARGET__.session_start > DATEADD(day, -7, CURRENT_DATE)" ) @@ -393,9 +411,9 @@ def test_model_kind(): on_additive_change=OnAdditiveChange.IGNORE, ) - assert ModelConfig(materialized=Materialization.INCREMENTAL, unique_key=["bar"]).model_kind( - context - ) == IncrementalByUniqueKeyKind( + assert ModelConfig( + materialized=Materialization.INCREMENTAL, unique_key=["bar"] + ).model_kind(context) == IncrementalByUniqueKeyKind( unique_key=["bar"], dialect="duckdb", forward_only=True, @@ -427,7 +445,9 @@ def test_model_kind(): ) assert ModelConfig( - materialized=Materialization.INCREMENTAL, unique_key=["bar"], disable_restatement=True + materialized=Materialization.INCREMENTAL, + unique_key=["bar"], + disable_restatement=True, ).model_kind(context) == IncrementalByUniqueKeyKind( unique_key=["bar"], dialect="duckdb", @@ -483,7 +503,9 @@ def test_model_kind(): ) assert ModelConfig( - materialized=Materialization.INCREMENTAL, time_column="foo", incremental_strategy="merge" + materialized=Materialization.INCREMENTAL, + time_column="foo", + incremental_strategy="merge", ).model_kind(context) == IncrementalByTimeRangeKind( time_column="foo", dialect="duckdb", @@ -586,9 +608,9 @@ def test_model_kind(): on_additive_change=OnAdditiveChange.IGNORE, ) - assert ModelConfig(materialized=Materialization.INCREMENTAL, forward_only=False).model_kind( - context - ) == IncrementalUnmanagedKind( + assert ModelConfig( + materialized=Materialization.INCREMENTAL, forward_only=False + ).model_kind(context) == IncrementalUnmanagedKind( insert_overwrite=True, disable_restatement=False, forward_only=False, @@ -605,7 +627,9 @@ def test_model_kind(): ) assert ModelConfig( - materialized=Materialization.INCREMENTAL, incremental_strategy="append", full_refresh=None + materialized=Materialization.INCREMENTAL, + incremental_strategy="append", + full_refresh=None, ).model_kind(context) == IncrementalUnmanagedKind( disable_restatement=False, on_destructive_change=OnDestructiveChange.IGNORE, @@ -673,9 +697,9 @@ def test_model_kind(): ) assert ( - ModelConfig(materialized=Materialization.DYNAMIC_TABLE, target_lag="1 hour").model_kind( - context - ) + ModelConfig( + materialized=Materialization.DYNAMIC_TABLE, target_lag="1 hour" + ).model_kind(context) == ManagedKind() ) @@ -774,7 +798,12 @@ def test_model_columns(): context = DbtContext() context.project_name = "Foo" context.target = SnowflakeConfig( - name="target", schema="test", database="test", account="foo", user="bar", password="baz" + name="target", + schema="test", + database="test", + account="foo", + user="bar", + password="baz", ) sqlmesh_model = model.to_sqlmesh(context) @@ -909,9 +938,11 @@ def test_seed_column_inference(tmp_path): context.target = DuckDbConfig(name="target", schema="test") sqlmesh_seed = seed.to_sqlmesh(context) assert sqlmesh_seed.columns_to_types == { - "int_col": exp.DataType.build("int") - if DBT_VERSION >= (1, 8, 0) - else exp.DataType.build("double"), + "int_col": ( + exp.DataType.build("int") + if DBT_VERSION >= (1, 8, 0) + else exp.DataType.build("double") + ), "double_col": exp.DataType.build("double"), "datetime_col": exp.DataType.build("datetime"), "date_col": exp.DataType.build("date"), @@ -1001,7 +1032,9 @@ def test_seed_delimiter(tmp_path): seed_csv = tmp_path / "seed_with_delimiter.csv" with open(seed_csv, "w", encoding="utf-8") as fd: - fd.writelines("\n".join(["id|name|city", "0|Ayrton|SP", "1|Max|MC", "2|Niki|VIE"])) + fd.writelines( + "\n".join(["id|name|city", "0|Ayrton|SP", "1|Max|MC", "2|Niki|VIE"]) + ) seed = SeedConfig( name="test_model_pipe", @@ -1042,7 +1075,10 @@ def test_seed_delimiter(tmp_path): sqlmesh_seed_semicolon = seed_semicolon.to_sqlmesh(context) expected_columns_semicolon = {"id", "value", "status"} - assert set(sqlmesh_seed_semicolon.columns_to_types.keys()) == expected_columns_semicolon + assert ( + set(sqlmesh_seed_semicolon.columns_to_types.keys()) + == expected_columns_semicolon + ) seed_df_semicolon = next(sqlmesh_seed_semicolon.render_seed()) assert seed_df_semicolon.iloc[0]["value"] == 100 @@ -1147,7 +1183,8 @@ def test_hooks(sushi_test_dbt_context: Context, model_fqn: str): with patch.object(logger, "debug") as mock_logger: engine_adapter.execute( waiters.render_post_statements( - engine_adapter=sushi_test_dbt_context.engine_adapter, execution_time="2023-01-01" + engine_adapter=sushi_test_dbt_context.engine_adapter, + execution_time="2023-01-01", ) ) assert "post-hook" in mock_logger.call_args[0][0] @@ -1162,7 +1199,11 @@ def test_seed_delimiter_integration(sushi_test_dbt_context: Context): assert seed_model.columns_to_types is not None # this should be loaded with semicolon delimiter otherwise it'd resylt in an one column table - assert set(seed_model.columns_to_types.keys()) == {"waiter_id", "revenue", "quarter"} + assert set(seed_model.columns_to_types.keys()) == { + "waiter_id", + "revenue", + "quarter", + } # columns_to_types values are correct types as well assert seed_model.columns_to_types == { @@ -1314,7 +1355,9 @@ def test_config_dict_syntax(): # Test nested dicts config4 = Config({}) - config4({"meta": {"owner": "data_team", "priority": 1}, "tags": ["daily", "critical"]}) + config4( + {"meta": {"owner": "data_team", "priority": 1}, "tags": ["daily", "critical"]} + ) assert config4._config["meta"]["owner"] == "data_team" assert config4._config["tags"] == ["daily", "critical"] @@ -1330,7 +1373,9 @@ def test_config_dict_syntax(): def test_config_dict_in_jinja(): # Test dict syntax directly with Config class config = Config({}) - template = Template("{{ config({'materialized': 'table', 'unique_key': 'id'}) }}done") + template = Template( + "{{ config({'materialized': 'table', 'unique_key': 'id'}) }}done" + ) result = template.render(config=config) assert result == "done" assert config._config["materialized"] == "table" @@ -1677,10 +1722,15 @@ def test_modules(sushi_test_project: Project): assert "object has no attribute 'pytz'" in str(error) # re - assert context.render("{{ modules.re.search('(?<=abc)def', 'abcdef').group(0) }}") == "def" + assert ( + context.render("{{ modules.re.search('(?<=abc)def', 'abcdef').group(0) }}") + == "def" + ) # itertools - itertools_jinja = "{% for num in modules.itertools.accumulate([5]) %}{{ num }}{% endfor %}" + itertools_jinja = ( + "{% for num in modules.itertools.accumulate([5]) %}{{ num }}{% endfor %}" + ) assert context.render(itertools_jinja) == "5" @@ -1729,12 +1779,13 @@ def test_relation(sushi_test_project: Project): def test_column(sushi_test_project: Project): context = sushi_test_project.context - assert context.render("{{ api.Column }}") == "" - - jinja = ( - "{% set col = api.Column('foo', 'integer') %}{{ col.is_integer() }} {{ col.is_string()}}" + assert ( + context.render("{{ api.Column }}") + == "" ) + jinja = "{% set col = api.Column('foo', 'integer') %}{{ col.is_integer() }} {{ col.is_string()}}" + assert context.render(jinja) == "True False" @@ -1780,7 +1831,9 @@ def test_json(sushi_test_project: Project): assert context.render("{{ tojson({'key': 'value'}) }}") == """{"key": "value"}""" assert context.render("{{ tojson(set([1])) }}") == "None" - assert context.render("""{{ fromjson('{"key": "value"}') }}""") == "{'key': 'value'}" + assert ( + context.render("""{{ fromjson('{"key": "value"}') }}""") == "{'key': 'value'}" + ) assert context.render("""{{ fromjson('invalid') }}""") == "None" @@ -1801,7 +1854,9 @@ def test_zip(sushi_test_project: Project): assert context.render("{{ zip([1, 2], ['a', 'b']) }}") == "[(1, 'a'), (2, 'b')]" assert context.render("{{ zip(12, ['a', 'b']) }}") == "None" - assert context.render("{{ zip_strict([1, 2], ['a', 'b']) }}") == "[(1, 'a'), (2, 'b')]" + assert ( + context.render("{{ zip_strict([1, 2], ['a', 'b']) }}") == "[(1, 'a'), (2, 'b')]" + ) with pytest.raises(TypeError): context.render("{{ zip_strict(12, ['a', 'b']) }}") @@ -1816,29 +1871,61 @@ def test_dbt_version(sushi_test_project: Project): @pytest.mark.xdist_group("dbt_manifest") def test_dbt_on_run_start_end(sushi_test_project: Project): # Validate perservation of dbt's order of execution - assert sushi_test_project.packages["sushi"].on_run_start["sushi-on-run-start-0"].index == 0 - assert sushi_test_project.packages["sushi"].on_run_start["sushi-on-run-start-1"].index == 1 - assert sushi_test_project.packages["sushi"].on_run_end["sushi-on-run-end-0"].index == 0 - assert sushi_test_project.packages["sushi"].on_run_end["sushi-on-run-end-1"].index == 1 assert ( - sushi_test_project.packages["customers"].on_run_start["customers-on-run-start-0"].index == 0 + sushi_test_project.packages["sushi"].on_run_start["sushi-on-run-start-0"].index + == 0 + ) + assert ( + sushi_test_project.packages["sushi"].on_run_start["sushi-on-run-start-1"].index + == 1 + ) + assert ( + sushi_test_project.packages["sushi"].on_run_end["sushi-on-run-end-0"].index == 0 + ) + assert ( + sushi_test_project.packages["sushi"].on_run_end["sushi-on-run-end-1"].index == 1 ) assert ( - sushi_test_project.packages["customers"].on_run_start["customers-on-run-start-1"].index == 1 + sushi_test_project.packages["customers"] + .on_run_start["customers-on-run-start-0"] + .index + == 0 + ) + assert ( + sushi_test_project.packages["customers"] + .on_run_start["customers-on-run-start-1"] + .index + == 1 + ) + assert ( + sushi_test_project.packages["customers"] + .on_run_end["customers-on-run-end-0"] + .index + == 0 + ) + assert ( + sushi_test_project.packages["customers"] + .on_run_end["customers-on-run-end-1"] + .index + == 1 ) - assert sushi_test_project.packages["customers"].on_run_end["customers-on-run-end-0"].index == 0 - assert sushi_test_project.packages["customers"].on_run_end["customers-on-run-end-1"].index == 1 assert ( - sushi_test_project.packages["customers"].on_run_start["customers-on-run-start-0"].sql + sushi_test_project.packages["customers"] + .on_run_start["customers-on-run-start-0"] + .sql == "CREATE TABLE IF NOT EXISTS to_be_executed_first (col VARCHAR);" ) assert ( - sushi_test_project.packages["customers"].on_run_start["customers-on-run-start-1"].sql + sushi_test_project.packages["customers"] + .on_run_start["customers-on-run-start-1"] + .sql == "CREATE TABLE IF NOT EXISTS analytic_stats_packaged_project (physical_table VARCHAR, evaluation_time VARCHAR);" ) assert ( - sushi_test_project.packages["customers"].on_run_end["customers-on-run-end-1"].sql + sushi_test_project.packages["customers"] + .on_run_end["customers-on-run-end-1"] + .sql == "{{ packaged_tables(schemas) }}" ) @@ -1900,12 +1987,20 @@ def test_partition_by(sushi_test_project: Project, caplog): assert date_trunc_expr.sql(dialect="bigquery") == "DATE_TRUNC(`ds`, MONTH)" assert date_trunc_expr.sql() == "DATE_TRUNC('MONTH', \"ds\")" - model_config.partition_by = {"field": "`ds`", "data_type": "datetime", "granularity": "day"} + model_config.partition_by = { + "field": "`ds`", + "data_type": "datetime", + "granularity": "day", + } datetime_trunc_expr = model_config.to_sqlmesh(context).partitioned_by[0] assert datetime_trunc_expr.sql(dialect="bigquery") == "datetime_trunc(`ds`, DAY)" assert datetime_trunc_expr.sql() == 'DATETIME_TRUNC("ds", DAY)' - model_config.partition_by = {"field": "ds", "data_type": "timestamp", "granularity": "day"} + model_config.partition_by = { + "field": "ds", + "data_type": "timestamp", + "granularity": "day", + } timestamp_trunc_expr = model_config.to_sqlmesh(context).partitioned_by[0] assert timestamp_trunc_expr.sql(dialect="bigquery") == "timestamp_trunc(`ds`, DAY)" assert timestamp_trunc_expr.sql() == 'TIMESTAMP_TRUNC("ds", DAY)' @@ -1920,14 +2015,25 @@ def test_partition_by(sushi_test_project: Project, caplog): == 'RANGE_BUCKET("one", GENERATE_SERIES(0, 10, 2))' ) - model_config.partition_by = {"field": "ds", "data_type": "date", "granularity": "day"} - assert model_config.to_sqlmesh(context).partitioned_by == [exp.to_column("ds", quoted=True)] + model_config.partition_by = { + "field": "ds", + "data_type": "date", + "granularity": "day", + } + assert model_config.to_sqlmesh(context).partitioned_by == [ + exp.to_column("ds", quoted=True) + ] context.target = DuckDbConfig(name="target", schema="foo") assert model_config.to_sqlmesh(context).partitioned_by == [] context.target = SnowflakeConfig( - name="target", schema="test", database="test", account="foo", user="bar", password="baz" + name="target", + schema="test", + database="test", + account="foo", + user="bar", + password="baz", ) assert model_config.to_sqlmesh(context).partitioned_by == [] assert ( @@ -1959,7 +2065,9 @@ def test_partition_by(sushi_test_project: Project, caplog): ) assert model_config.to_sqlmesh(context).partitioned_by == [] - with pytest.raises(ConfigError, match="Unexpected data_type 'string' in partition_by"): + with pytest.raises( + ConfigError, match="Unexpected data_type 'string' in partition_by" + ): ModelConfig( name="model", alias="model", @@ -2076,7 +2184,9 @@ def test_is_incremental(sushi_test_project: Project, assert_exp_eq, mocker): @pytest.mark.xdist_group("dbt_manifest") -def test_is_incremental_non_incremental_model(sushi_test_project: Project, assert_exp_eq, mocker): +def test_is_incremental_non_incremental_model( + sushi_test_project: Project, assert_exp_eq, mocker +): model_config = ModelConfig( name="model", package_name="package", @@ -2102,7 +2212,9 @@ def test_is_incremental_non_incremental_model(sushi_test_project: Project, asser @pytest.mark.xdist_group("dbt_manifest") -def test_dbt_max_partition(sushi_test_project: Project, assert_exp_eq, mocker: MockerFixture): +def test_dbt_max_partition( + sushi_test_project: Project, assert_exp_eq, mocker: MockerFixture +): model_config = ModelConfig( name="model", alias="model", @@ -2124,9 +2236,7 @@ def test_dbt_max_partition(sushi_test_project: Project, assert_exp_eq, mocker: M pre_statement = model_config.to_sqlmesh(context).pre_statements[-1] # type: ignore - assert ( - pre_statement.sql().strip() - == """ + assert pre_statement.sql().strip() == """ JINJA_STATEMENT_BEGIN; {% if is_incremental() %} DECLARE _dbt_max_partition DATETIME DEFAULT ( @@ -2134,13 +2244,14 @@ def test_dbt_max_partition(sushi_test_project: Project, assert_exp_eq, mocker: M ); {% endif %} JINJA_END;""".strip() - ) assert d.parse_one(pre_statement.sql()) == pre_statement @pytest.mark.xdist_group("dbt_manifest") -def test_bigquery_physical_properties(sushi_test_project: Project, mocker: MockerFixture): +def test_bigquery_physical_properties( + sushi_test_project: Project, mocker: MockerFixture +): context = sushi_test_project.context context.target = BigQueryConfig( name="test_target", schema="test_schema", database="test-project" @@ -2219,11 +2330,15 @@ def test_clickhouse_properties(mocker: MockerFixture): assert model_to_sqlmesh.storage_format == "MergeTree()" physical_properties = model_to_sqlmesh.physical_properties - assert [e.sql("clickhouse", identify=True) for e in physical_properties["order_by"]] == [ + assert [ + e.sql("clickhouse", identify=True) for e in physical_properties["order_by"] + ] == [ 'toStartOfWeek("ds")', '"order_col"', ] - assert [e.sql("clickhouse", identify=True) for e in physical_properties["primary_key"]] == [ + assert [ + e.sql("clickhouse", identify=True) for e in physical_properties["primary_key"] + ] == [ '"ds"', '"primary_key_col"', ] @@ -2235,7 +2350,9 @@ def test_clickhouse_properties(mocker: MockerFixture): def test_snapshot_json_payload(): sushi_context = Context(paths=["tests/fixtures/dbt/sushi_test"]) snapshot_json = json.loads( - _snapshot_to_json(sushi_context.get_snapshot("sushi.top_waiters", raise_if_missing=True)) + _snapshot_to_json( + sushi_context.get_snapshot("sushi.top_waiters", raise_if_missing=True) + ) ) assert snapshot_json["node"]["jinja_macros"]["global_objs"]["target"] == { "type": "duckdb", @@ -2275,10 +2392,17 @@ def test_dbt_vars(sushi_test_project: Project): @pytest.mark.xdist_group("dbt_manifest") -def test_snowflake_session_properties(sushi_test_project: Project, mocker: MockerFixture): +def test_snowflake_session_properties( + sushi_test_project: Project, mocker: MockerFixture +): context = sushi_test_project.context context.target = SnowflakeConfig( - name="target", schema="test", database="test", account="foo", user="bar", password="baz" + name="target", + schema="test", + database="test", + account="foo", + user="bar", + password="baz", ) base_config = ModelConfig( @@ -2298,7 +2422,9 @@ def test_snowflake_session_properties(sushi_test_project: Project, mocker: Mocke ).to_sqlmesh(context) assert model_with_warehouse.session_properties_ == exp.Tuple( - expressions=[exp.Literal.string("warehouse").eq(exp.Literal.string("test_warehouse"))] + expressions=[ + exp.Literal.string("warehouse").eq(exp.Literal.string("test_warehouse")) + ] ) assert model_with_warehouse.session_properties == {"warehouse": "test_warehouse"} @@ -2458,7 +2584,9 @@ def test_refs_in_jinja_globals(sushi_test_project: Project, mocker: MockerFixtur sqlmesh_model = t.cast( SqlModel, - sushi_test_project.packages["sushi"].models["simple_model_b"].to_sqlmesh(context), + sushi_test_project.packages["sushi"] + .models["simple_model_b"] + .to_sqlmesh(context), ) assert set(sqlmesh_model.jinja_macros.global_objs["refs"].keys()) == {"simple_model_a"} # type: ignore @@ -2525,7 +2653,10 @@ def test_grain(): materialized=Materialization.TABLE.value, grain=["id_a", "id_b"], ) - assert model.to_sqlmesh(context).grains == [exp.to_column("id_a"), exp.to_column("id_b")] + assert model.to_sqlmesh(context).grains == [ + exp.to_column("id_a"), + exp.to_column("id_b"), + ] model.grain = "id_a" assert model.to_sqlmesh(context).grains == [exp.to_column("id_a")] @@ -2605,15 +2736,30 @@ def test_on_run_start_end(): assert sorted(runtime_rendered_after_all[:-1]) == sorted(expected_statements) # Assert the models with their materialisations are present in the rendered graph_table statement - assert "'model.sushi.simple_model_a' AS unique_id, 'table' AS materialized" in graph_table_stmt - assert "'model.sushi.waiters' AS unique_id, 'ephemeral' AS materialized" in graph_table_stmt - assert "'model.sushi.simple_model_b' AS unique_id, 'table' AS materialized" in graph_table_stmt + assert ( + "'model.sushi.simple_model_a' AS unique_id, 'table' AS materialized" + in graph_table_stmt + ) + assert ( + "'model.sushi.waiters' AS unique_id, 'ephemeral' AS materialized" + in graph_table_stmt + ) + assert ( + "'model.sushi.simple_model_b' AS unique_id, 'table' AS materialized" + in graph_table_stmt + ) assert ( "'model.sushi.waiter_as_customer_by_day' AS unique_id, 'incremental' AS materialized" in graph_table_stmt ) - assert "'model.sushi.top_waiters' AS unique_id, 'view' AS materialized" in graph_table_stmt - assert "'model.customers.customers' AS unique_id, 'view' AS materialized" in graph_table_stmt + assert ( + "'model.sushi.top_waiters' AS unique_id, 'view' AS materialized" + in graph_table_stmt + ) + assert ( + "'model.customers.customers' AS unique_id, 'view' AS materialized" + in graph_table_stmt + ) assert ( "'model.customers.customer_revenue_by_day' AS unique_id, 'incremental' AS materialized" in graph_table_stmt @@ -2682,9 +2828,13 @@ def test_on_run_start_end(): @pytest.mark.xdist_group("dbt_manifest") -def test_dynamic_var_names(sushi_test_project: Project, sushi_test_dbt_context: Context): +def test_dynamic_var_names( + sushi_test_project: Project, sushi_test_dbt_context: Context +): context = sushi_test_project.context - context.set_and_render_variables(sushi_test_project.packages["sushi"].variables, "sushi") + context.set_and_render_variables( + sushi_test_project.packages["sushi"].variables, "sushi" + ) context.target = BigQueryConfig(name="production", database="main", schema="sushi") model_config = ModelConfig( name="model", @@ -2720,7 +2870,9 @@ def test_dynamic_var_names(sushi_test_project: Project, sushi_test_dbt_context: @pytest.mark.xdist_group("dbt_manifest") def test_dynamic_var_names_in_macro(sushi_test_project: Project): context = sushi_test_project.context - context.set_and_render_variables(sushi_test_project.packages["sushi"].variables, "sushi") + context.set_and_render_variables( + sushi_test_project.packages["sushi"].variables, "sushi" + ) context.target = BigQueryConfig(name="production", database="main", schema="sushi") model_config = ModelConfig( name="model", @@ -2735,7 +2887,9 @@ def test_dynamic_var_names_in_macro(sushi_test_project: Project): SELECT {{ sushi.dynamic_var_name_dependency(var_name) }} AS var """, dependencies=Dependencies( - macros=[MacroReference(package="sushi", name="dynamic_var_name_dependency")], + macros=[ + MacroReference(package="sushi", name="dynamic_var_name_dependency") + ], has_dynamic_var_names=True, ), ) @@ -2814,7 +2968,9 @@ def test_selected_resources_context_variable( ] # check the jinja macros rendering - result = context.render("{{ selected_resources }}", selected_resources=selected_resources) + result = context.render( + "{{ selected_resources }}", selected_resources=selected_resources + ) assert result == selected_resources.__repr__() result = context.render(test_jinja, selected_resources=selected_resources) @@ -2824,7 +2980,9 @@ def test_selected_resources_context_variable( assert result.strip() == "has_resources" -def test_ignore_source_depends_on_when_also_model(dbt_dummy_postgres_config: PostgresConfig): +def test_ignore_source_depends_on_when_also_model( + dbt_dummy_postgres_config: PostgresConfig, +): context = DbtContext() context._target = dbt_dummy_postgres_config @@ -2870,15 +3028,18 @@ def test_dbt_hooks_with_transaction_flag(sushi_test_dbt_context: Context): for s in pre_statements ) assert any( - "CREATE TABLE IF NOT EXISTS hook_outside_pre_table" in s.sql and s.transaction is False + "CREATE TABLE IF NOT EXISTS hook_outside_pre_table" in s.sql + and s.transaction is False for s in pre_statements ) assert any( - "CREATE TABLE IF NOT EXISTS shared_hook_table" in s.sql and s.transaction is False + "CREATE TABLE IF NOT EXISTS shared_hook_table" in s.sql + and s.transaction is False for s in pre_statements ) assert any( - "{{ insert_into_shared_hook_table('inside_pre') }}" in s.sql and s.transaction is True + "{{ insert_into_shared_hook_table('inside_pre') }}" in s.sql + and s.transaction is True for s in pre_statements ) @@ -2891,15 +3052,18 @@ def test_dbt_hooks_with_transaction_flag(sushi_test_dbt_context: Context): for s in post_statements ) assert any( - "{{ insert_into_shared_hook_table('inside_post') }}" in s.sql and s.transaction is True + "{{ insert_into_shared_hook_table('inside_post') }}" in s.sql + and s.transaction is True for s in post_statements ) assert any( - "CREATE TABLE IF NOT EXISTS hook_outside_post_table" in s.sql and s.transaction is False + "CREATE TABLE IF NOT EXISTS hook_outside_post_table" in s.sql + and s.transaction is False for s in post_statements ) assert any( - "{{ insert_into_shared_hook_table('after_commit') }}" in s.sql and s.transaction is False + "{{ insert_into_shared_hook_table('after_commit') }}" in s.sql + and s.transaction is False for s in post_statements ) @@ -2940,7 +3104,9 @@ def test_dbt_hooks_with_transaction_flag_execution(sushi_test_dbt_context: Conte model_fqn = '"memory"."sushi"."model_with_transaction_hooks"' assert model_fqn in sushi_test_dbt_context.models - plan = sushi_test_dbt_context.plan(select_models=["sushi.model_with_transaction_hooks"]) + plan = sushi_test_dbt_context.plan( + select_models=["sushi.model_with_transaction_hooks"] + ) sushi_test_dbt_context.apply(plan) result = sushi_test_dbt_context.engine_adapter.fetchdf( diff --git a/tests/dbt/test_util.py b/tests/dbt/test_util.py index ce98f48a82..69056effe2 100644 --- a/tests/dbt/test_util.py +++ b/tests/dbt/test_util.py @@ -1,6 +1,7 @@ from __future__ import annotations import pandas as pd # noqa: TID253 + from sqlmesh.dbt.util import pandas_to_agate diff --git a/tests/engines/spark/test_db_api.py b/tests/engines/spark/test_db_api.py index eab7a0c223..d170d8e2b8 100644 --- a/tests/engines/spark/test_db_api.py +++ b/tests/engines/spark/test_db_api.py @@ -17,8 +17,7 @@ def test_spark_session_cursor(spark_session: SparkSession): with pytest.raises(errors.ProgrammingError): cursor.fetchone() - cursor.execute( - """ + cursor.execute(""" SELECT * FROM VALUES ('key1', 1), @@ -28,8 +27,7 @@ def test_spark_session_cursor(spark_session: SparkSession): ('key5', 5), ('key6', 6), ('key7', 7) AS data(key, value) - """ - ) + """) assert cursor.fetchone() == ("key1", 1) assert cursor.fetchmany(size=2) == [ diff --git a/tests/fixtures/dbt/sushi_test/config.py b/tests/fixtures/dbt/sushi_test/config.py index a68b3e2333..4d7b5dfe34 100644 --- a/tests/fixtures/dbt/sushi_test/config.py +++ b/tests/fixtures/dbt/sushi_test/config.py @@ -3,9 +3,9 @@ from sqlmesh.core.config import ModelDefaultsConfig from sqlmesh.dbt.loader import sqlmesh_config - config = sqlmesh_config( - Path(__file__).parent, model_defaults=ModelDefaultsConfig(dialect="duckdb", start="Jan 1 2022") + Path(__file__).parent, + model_defaults=ModelDefaultsConfig(dialect="duckdb", start="Jan 1 2022"), ) diff --git a/tests/integrations/github/cicd/conftest.py b/tests/integrations/github/cicd/conftest.py index f869dc41ad..84f3cb6df9 100644 --- a/tests/integrations/github/cicd/conftest.py +++ b/tests/integrations/github/cicd/conftest.py @@ -1,20 +1,18 @@ import typing as t +from pathlib import Path import pytest from pytest_mock.plugin import MockerFixture -from pathlib import Path +from sqlglot.helper import ensure_list from sqlmesh.core.config import Config -from sqlmesh.core.console import set_console, get_console, MarkdownConsole +from sqlmesh.core.console import MarkdownConsole, get_console, set_console from sqlmesh.integrations.github.cicd.config import GithubCICDBotConfig -from sqlmesh.integrations.github.cicd.controller import ( - GithubController, - GithubEvent, - MergeStateStatus, - PullRequestInfo, -) +from sqlmesh.integrations.github.cicd.controller import (GithubController, + GithubEvent, + MergeStateStatus, + PullRequestInfo) from sqlmesh.utils import AttributeDict -from sqlglot.helper import ensure_list @pytest.fixture @@ -33,7 +31,9 @@ def github_client(mocker: MockerFixture): mock_pull_request = mocker.MagicMock(spec=PullRequest) mock_pull_request.base.ref = "main" - mock_pull_request.get_reviews.return_value = [mocker.MagicMock(spec=PullRequestReview)] + mock_pull_request.get_reviews.return_value = [ + mocker.MagicMock(spec=PullRequestReview) + ] mock_repository.get_pull.return_value = mock_pull_request mock_repository.get_issue.return_value = mocker.MagicMock(spec=Issue) @@ -85,7 +85,9 @@ def _make_function( ) -> GithubController: if mock_out_context: mocker.patch("sqlmesh.core.context.Context.apply", mocker.MagicMock()) - mocker.patch("sqlmesh.core.context.Context._run_plan_tests", mocker.MagicMock()) + mocker.patch( + "sqlmesh.core.context.Context._run_plan_tests", mocker.MagicMock() + ) mocker.patch("sqlmesh.core.context.Context._run_tests", mocker.MagicMock()) mocker.patch( "sqlmesh.integrations.github.cicd.controller.GithubController._get_merge_state_status", @@ -117,7 +119,9 @@ def _make_function( orig_console = get_console() try: - set_console(MarkdownConsole(warning_capture_only=True, error_capture_only=True)) + set_console( + MarkdownConsole(warning_capture_only=True, error_capture_only=True) + ) return GithubController( paths=paths, diff --git a/tests/integrations/github/cicd/test_config.py b/tests/integrations/github/cicd/test_config.py index 20b72393dc..bf3948a68f 100644 --- a/tests/integrations/github/cicd/test_config.py +++ b/tests/integrations/github/cicd/test_config.py @@ -2,14 +2,10 @@ import pytest -from sqlmesh.core.config import ( - AutoCategorizationMode, - CategorizerConfig, - Config, - load_config_from_paths, -) -from sqlmesh.utils.errors import ConfigError +from sqlmesh.core.config import (AutoCategorizationMode, CategorizerConfig, + Config, load_config_from_paths) from sqlmesh.integrations.github.cicd.config import MergeMethod +from sqlmesh.utils.errors import ConfigError from tests.utils.test_filesystem import create_temp_file pytestmark = pytest.mark.github @@ -34,7 +30,9 @@ def test_load_yaml_config_default(tmp_path): assert config.cicd_bot.invalidate_environment_after_deploy assert config.cicd_bot.merge_method is None assert config.cicd_bot.command_namespace is None - assert config.cicd_bot.auto_categorize_changes == config.plan.auto_categorize_changes + assert ( + config.cicd_bot.auto_categorize_changes == config.plan.auto_categorize_changes + ) assert config.cicd_bot.default_pr_start is None assert config.cicd_bot.default_pr_preview_start == "yesterday" assert not config.cicd_bot.enable_deploy_command @@ -121,7 +119,9 @@ def test_load_python_config_defaults(tmp_path): assert config.cicd_bot.invalidate_environment_after_deploy assert config.cicd_bot.merge_method is None assert config.cicd_bot.command_namespace is None - assert config.cicd_bot.auto_categorize_changes == config.plan.auto_categorize_changes + assert ( + config.cicd_bot.auto_categorize_changes == config.plan.auto_categorize_changes + ) assert config.cicd_bot.default_pr_start is None assert config.cicd_bot.default_pr_preview_start == "yesterday" assert not config.cicd_bot.enable_deploy_command @@ -206,7 +206,8 @@ def test_validation(tmp_path): """, ) with pytest.raises( - ConfigError, match=r".*enable_deploy_command must be set if command_namespace is set.*" + ConfigError, + match=r".*enable_deploy_command must be set if command_namespace is set.*", ): load_config_from_paths(Config, project_paths=[tmp_path / "config.yaml"]) @@ -222,7 +223,8 @@ def test_validation(tmp_path): """, ) with pytest.raises( - ConfigError, match=r".*merge_method must be set if enable_deploy_command is True.*" + ConfigError, + match=r".*merge_method must be set if enable_deploy_command is True.*", ): load_config_from_paths(Config, project_paths=[tmp_path / "config.yaml"]) @@ -301,4 +303,6 @@ def test_properties_inherit_from_project_config(tmp_path): seed=AutoCategorizationMode.FULL, ) ) - assert config.cicd_bot.pr_include_unmodified == config.plan.include_unmodified == True + assert ( + config.cicd_bot.pr_include_unmodified == config.plan.include_unmodified == True + ) diff --git a/tests/integrations/github/cicd/test_github_commands.py b/tests/integrations/github/cicd/test_github_commands.py index 01e4c9af31..7e5a8f5cd7 100644 --- a/tests/integrations/github/cicd/test_github_commands.py +++ b/tests/integrations/github/cicd/test_github_commands.py @@ -1,7 +1,7 @@ # type: ignore -import typing as t import os import pathlib +import typing as t from unittest import TestCase, mock from unittest.result import TestResult @@ -15,13 +15,13 @@ from sqlmesh.core.test.result import ModelTextTestResult from sqlmesh.core.user import User, UserRole from sqlmesh.integrations.github.cicd import command -from sqlmesh.integrations.github.cicd.config import GithubCICDBotConfig, MergeMethod -from sqlmesh.integrations.github.cicd.controller import ( - GithubController, - GithubCheckConclusion, - GithubCheckStatus, -) -from sqlmesh.utils.errors import ConflictingPlanError, PlanError, TestError, CICDBotError +from sqlmesh.integrations.github.cicd.config import (GithubCICDBotConfig, + MergeMethod) +from sqlmesh.integrations.github.cicd.controller import (GithubCheckConclusion, + GithubCheckStatus, + GithubController) +from sqlmesh.utils.errors import (CICDBotError, ConflictingPlanError, + PlanError, TestError) pytestmark = [ pytest.mark.github, @@ -62,7 +62,9 @@ def test_run_all_success_with_approvers_approved( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -78,7 +80,11 @@ def test_run_all_success_with_approvers_approved( side_effect=lambda **kwargs: (TestResult(), "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -88,7 +94,9 @@ def test_run_all_success_with_approvers_approved( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -96,7 +104,9 @@ def test_run_all_success_with_approvers_approved( assert GithubCheckConclusion(test_checks_runs[2]["conclusion"]).is_success assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress @@ -111,10 +121,14 @@ def test_run_all_success_with_approvers_approved( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -141,16 +155,14 @@ def test_run_all_success_with_approvers_approved( assert not controller._context.invalidate_environment.called assert len(created_comments) == 1 - assert created_comments[0].body.startswith( - """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body.startswith(""":robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `myoverride_2`
:ship: Prod Plan Being Applied -**`prod` environment will be initialized**""" - ) +**`prod` environment will be initialized**""") with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() assert ( @@ -192,7 +204,9 @@ def test_run_all_success_with_approvers_approved_merge_delete( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -206,7 +220,11 @@ def test_run_all_success_with_approvers_approved_merge_delete( side_effect=lambda **kwargs: (TestResult(), "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -216,7 +234,9 @@ def test_run_all_success_with_approvers_approved_merge_delete( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -231,10 +251,14 @@ def test_run_all_success_with_approvers_approved_merge_delete( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress @@ -242,7 +266,9 @@ def test_run_all_success_with_approvers_approved_merge_delete( assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -269,16 +295,14 @@ def test_run_all_success_with_approvers_approved_merge_delete( assert controller._context.invalidate_environment.called assert len(created_comments) == 1 - assert created_comments[0].body.startswith( - """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body.startswith(""":robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`
:ship: Prod Plan Being Applied -**`prod` environment will be initialized**""" - ) +**`prod` environment will be initialized**""") with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() assert ( @@ -336,7 +360,11 @@ def test_run_all_missing_approval( side_effect=lambda **kwargs: (TestResult(), "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -353,10 +381,14 @@ def test_run_all_missing_approval( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress @@ -364,14 +396,18 @@ def test_run_all_missing_approval( assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed assert GithubCheckConclusion(prod_checks_runs[1]["conclusion"]).is_skipped assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -396,12 +432,9 @@ def test_run_all_missing_approval( assert not controller._context.invalidate_environment.called assert len(created_comments) == 1 - assert ( - created_comments[0].body - == """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body == """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`""" - ) with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() assert ( @@ -443,7 +476,9 @@ def test_run_all_test_failed( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -460,7 +495,11 @@ def test_run_all_test_failed( side_effect=lambda **kwargs: (test_result, "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -471,7 +510,9 @@ def test_run_all_test_failed( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -479,7 +520,8 @@ def test_run_all_test_failed( assert GithubCheckConclusion(test_checks_runs[2]["conclusion"]).is_failure assert test_checks_runs[2]["output"]["title"] == "Tests Failed" assert ( - """sqlmesh.utils.errors.TestError: some error""" in test_checks_runs[2]["output"]["summary"] + """sqlmesh.utils.errors.TestError: some error""" + in test_checks_runs[2]["output"]["summary"] ) assert """Failed tests""" in test_checks_runs[2]["output"]["summary"] @@ -490,7 +532,9 @@ def test_run_all_test_failed( assert len(prod_plan_preview_checks_runs) == 2 assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[1]["conclusion"]).is_skipped + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[1]["conclusion"] + ).is_skipped assert ( prod_plan_preview_checks_runs[1]["output"]["title"] == "Skipped generating prod plan preview since PR was not synchronized" @@ -501,15 +545,22 @@ def test_run_all_test_failed( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 2 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[1]["conclusion"]).is_skipped - assert pr_checks_runs[1]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" + assert ( + pr_checks_runs[1]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed @@ -573,7 +624,9 @@ def test_run_all_test_exception( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -585,9 +638,15 @@ def test_run_all_test_exception( ) test_result = TestResult() test_result.addFailure(TestCase(), (None, None, None)) - controller._context._run_tests = mocker.MagicMock(side_effect=TestError("some error")) + controller._context._run_tests = mocker.MagicMock( + side_effect=TestError("some error") + ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -598,7 +657,9 @@ def test_run_all_test_exception( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -618,7 +679,9 @@ def test_run_all_test_exception( assert len(prod_plan_preview_checks_runs) == 2 assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[1]["conclusion"]).is_skipped + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[1]["conclusion"] + ).is_skipped assert ( prod_plan_preview_checks_runs[1]["output"]["title"] == "Skipped generating prod plan preview since PR was not synchronized" @@ -629,15 +692,22 @@ def test_run_all_test_exception( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 2 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[1]["conclusion"]).is_skipped - assert pr_checks_runs[1]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" + assert ( + pr_checks_runs[1]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed @@ -702,7 +772,9 @@ def test_pr_update_failure( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -716,7 +788,11 @@ def test_pr_update_failure( side_effect=lambda **kwargs: (TestResult(), "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -724,7 +800,9 @@ def raise_on_pr_plan(plan: Plan): if plan.environment.name == "hello_world_2": raise PlanError("Failed to update PR environment") - controller._context.apply = mocker.MagicMock(side_effect=lambda plan: raise_on_pr_plan(plan)) + controller._context.apply = mocker.MagicMock( + side_effect=lambda plan: raise_on_pr_plan(plan) + ) github_output_file = tmp_path / "github_output.txt" @@ -733,7 +811,9 @@ def raise_on_pr_plan(plan: Plan): command._run_all(controller) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress @@ -747,17 +827,23 @@ def raise_on_pr_plan(plan: Plan): assert len(prod_plan_preview_checks_runs) == 2 assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[1]["conclusion"]).is_skipped + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[1]["conclusion"] + ).is_skipped assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed assert GithubCheckConclusion(prod_checks_runs[1]["conclusion"]).is_skipped assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -827,7 +913,9 @@ def make_test_prod_update_failure_case( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -841,7 +929,11 @@ def make_test_prod_update_failure_case( side_effect=lambda **kwargs: (TestResult(), "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -849,7 +941,9 @@ def raise_on_prod_plan(plan: Plan): if plan.environment.name == c.PROD: raise to_raise_on_prod_plan - controller._context.apply = mocker.MagicMock(side_effect=lambda plan: raise_on_prod_plan(plan)) + controller._context.apply = mocker.MagicMock( + side_effect=lambda plan: raise_on_prod_plan(plan) + ) github_output_file = tmp_path / "github_output.txt" @@ -865,10 +959,14 @@ def raise_on_prod_plan(plan: Plan): assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress @@ -876,24 +974,31 @@ def raise_on_prod_plan(plan: Plan): assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_checks_runs[2]["conclusion"]) == expect_prod_sync_conclusion + assert ( + GithubCheckConclusion(prod_checks_runs[2]["conclusion"]) + == expect_prod_sync_conclusion + ) if expect_prod_sync_conclusion.is_action_required: - assert prod_checks_runs[2]["output"]["title"] == "Failed due to error applying plan" assert ( - prod_checks_runs[2]["output"]["summary"] - == f"""**Plan error:** + prod_checks_runs[2]["output"]["title"] + == "Failed due to error applying plan" + ) + assert prod_checks_runs[2]["output"]["summary"] == f"""**Plan error:** ``` {to_raise_on_prod_plan} ```""" - ) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -920,16 +1025,14 @@ def raise_on_prod_plan(plan: Plan): assert not controller._context.invalidate_environment.called assert len(created_comments) == 1 - assert created_comments[0].body.startswith( - """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body.startswith(""":robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`
:ship: Prod Plan Being Applied -**`prod` environment will be initialized**""" - ) +**`prod` environment will be initialized**""") with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() @@ -1070,7 +1173,9 @@ def test_comment_command_invalid( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -1084,7 +1189,11 @@ def test_comment_command_invalid( side_effect=lambda **kwargs: (TestResult(), "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() @@ -1147,13 +1256,19 @@ def test_comment_command_deploy_prod( controller = make_controller( make_event_issue_comment("created", "/deploy"), github_client, - bot_config=GithubCICDBotConfig(merge_method=MergeMethod.REBASE, enable_deploy_command=True), + bot_config=GithubCICDBotConfig( + merge_method=MergeMethod.REBASE, enable_deploy_command=True + ), ) controller._context._run_tests = mocker.MagicMock( side_effect=lambda **kwargs: (TestResult(), "") ) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] controller._context.invalidate_environment = mocker.MagicMock() assert not controller.forward_only_plan @@ -1171,10 +1286,14 @@ def test_comment_command_deploy_prod( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress @@ -1182,7 +1301,9 @@ def test_comment_command_deploy_prod( assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -1190,7 +1311,9 @@ def test_comment_command_deploy_prod( assert GithubCheckConclusion(prod_checks_runs[2]["conclusion"]).is_success assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -1212,8 +1335,7 @@ def test_comment_command_deploy_prod( assert controller._context.invalidate_environment.called assert len(created_comments) == 1 - assert created_comments[0].body.startswith( - """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body.startswith(""":robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2` - :arrow_forward: To **apply** this PR's plan to prod, comment: @@ -1222,8 +1344,7 @@ def test_comment_command_deploy_prod( :ship: Prod Plan Being Applied -**`prod` environment will be initialized**""" - ) +**`prod` environment will be initialized**""") with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() @@ -1343,7 +1464,9 @@ def test_comment_command_deploy_prod_no_deploy_detected_yet( controller = make_controller( "tests/fixtures/github/pull_request_synchronized.json", github_client, - bot_config=GithubCICDBotConfig(merge_method=MergeMethod.REBASE, enable_deploy_command=True), + bot_config=GithubCICDBotConfig( + merge_method=MergeMethod.REBASE, enable_deploy_command=True + ), ) controller._context._run_tests = mocker.MagicMock( side_effect=lambda **kwargs: (TestResult(), "") @@ -1358,7 +1481,9 @@ def test_comment_command_deploy_prod_no_deploy_detected_yet( assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed @@ -1436,7 +1561,9 @@ def test_deploy_prod_forward_only( # Prod Environment Synced step should be successful assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_in_progress assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed diff --git a/tests/integrations/github/cicd/test_github_controller.py b/tests/integrations/github/cicd/test_github_controller.py index 786341d361..d070fce5df 100644 --- a/tests/integrations/github/cicd/test_github_controller.py +++ b/tests/integrations/github/cicd/test_github_controller.py @@ -1,7 +1,7 @@ # type: ignore -import typing as t import os import pathlib +import typing as t from unittest import mock from unittest.mock import PropertyMock, call @@ -12,21 +12,20 @@ from sqlmesh.core import constants as c from sqlmesh.core.config import CategorizerConfig from sqlmesh.core.dialect import parse_one +from sqlmesh.core.linter.rule import RuleViolation from sqlmesh.core.model import SqlModel -from sqlmesh.core.user import User, UserRole from sqlmesh.core.plan.definition import Plan -from sqlmesh.core.linter.rule import RuleViolation -from sqlmesh.integrations.github.cicd.config import GithubCICDBotConfig, MergeMethod -from sqlmesh.integrations.github.cicd.controller import ( - BotCommand, - MergeStateStatus, - GithubCheckConclusion, -) -from sqlmesh.integrations.github.cicd.controller import GithubController +from sqlmesh.core.user import User, UserRole from sqlmesh.integrations.github.cicd.command import _update_pr_environment -from sqlmesh.utils.date import to_datetime, now -from tests.integrations.github.cicd.conftest import MockIssueComment +from sqlmesh.integrations.github.cicd.config import (GithubCICDBotConfig, + MergeMethod) +from sqlmesh.integrations.github.cicd.controller import (BotCommand, + GithubCheckConclusion, + GithubController, + MergeStateStatus) +from sqlmesh.utils.date import now, to_datetime from sqlmesh.utils.errors import SQLMeshError +from tests.integrations.github.cicd.conftest import MockIssueComment pytestmark = pytest.mark.github @@ -41,13 +40,21 @@ class _MockLinterRule: controller._console.show_linter_violations( [ RuleViolation( - rule=_MockLinterRule(), violation_msg="Linter warning", violation_range=None + rule=_MockLinterRule(), + violation_msg="Linter warning", + violation_range=None, ) ], _MockModel(), ) controller._console.show_linter_violations( - [RuleViolation(rule=_MockLinterRule(), violation_msg="Linter error", violation_range=None)], + [ + RuleViolation( + rule=_MockLinterRule(), + violation_msg="Linter error", + violation_range=None, + ) + ], _MockModel(), is_error=True, ) @@ -235,9 +242,16 @@ def test_github_controller_approvers( ids=[test[0] for test in is_comment_triggered_params], ) def test_is_comment( - action, comment, is_comment_triggered, github_client, make_event_issue_comment, make_controller + action, + comment, + is_comment_triggered, + github_client, + make_event_issue_comment, + make_controller, ): - controller = make_controller(make_event_issue_comment(action, comment), github_client) + controller = make_controller( + make_event_issue_comment(action, comment), github_client + ) assert controller.is_comment_added == is_comment_triggered @@ -271,7 +285,8 @@ def test_pr_plan_auto_categorization(github_client, make_controller): "tests/fixtures/github/pull_request_synchronized.json", github_client, bot_config=GithubCICDBotConfig( - auto_categorize_changes=custom_categorizer_config, default_pr_start=default_start + auto_categorize_changes=custom_categorizer_config, + default_pr_start=default_start, ), ) assert controller.pr_plan.environment.name == "hello_world_2" @@ -288,7 +303,9 @@ def test_pr_plan_min_intervals(github_client, make_controller): controller = make_controller( "tests/fixtures/github/pull_request_synchronized.json", github_client, - bot_config=GithubCICDBotConfig(default_pr_start="1 day ago", pr_min_intervals=1), + bot_config=GithubCICDBotConfig( + default_pr_start="1 day ago", pr_min_intervals=1 + ), ) assert controller.pr_plan.environment.name == "hello_world_2" assert isinstance(controller.pr_plan, Plan) @@ -302,7 +319,9 @@ def test_pr_plan_preview_window(github_client, make_controller, mocker: MockerFi plan_builder.build.return_value = plan context.config.plan.forward_only = False context.plan_builder.return_value = plan_builder - mocker.patch("sqlmesh.integrations.github.cicd.controller.Context", return_value=context) + mocker.patch( + "sqlmesh.integrations.github.cicd.controller.Context", return_value=context + ) bot_config = GithubCICDBotConfig(default_pr_start="2025-01-01") controller = make_controller( @@ -379,7 +398,8 @@ def test_prod_plan_auto_categorization(github_client, make_controller): "tests/fixtures/github/pull_request_synchronized.json", github_client, bot_config=GithubCICDBotConfig( - auto_categorize_changes=custom_categorizer_config, default_pr_start=default_pr_start + auto_categorize_changes=custom_categorizer_config, + default_pr_start=default_pr_start, ), ) @@ -388,7 +408,9 @@ def test_prod_plan_auto_categorization(github_client, make_controller): assert controller.prod_plan.no_gaps assert not controller._context.apply.called assert controller._context._run_plan_tests.call_args == call(skip_tests=True) - assert controller._prod_plan_builder._categorizer_config == custom_categorizer_config + assert ( + controller._prod_plan_builder._categorizer_config == custom_categorizer_config + ) # default PR start should be ignored for prod plans assert controller.prod_plan.start != default_pr_start @@ -468,7 +490,8 @@ def test_run_tests(github_client, make_controller): "桜" * 65000, None, # ((Max Byte Length) - (Length of ":robot: **SQLMesh Bot Info** :robot:\ntest1\n")) / (Length of "桜") - ":robot: **SQLMesh Bot Info** :robot:\ntest1\n" + ("桜" * int((65535 - 43) / 3)), + ":robot: **SQLMesh Bot Info** :robot:\ntest1\n" + + ("桜" * int((65535 - 43) / 3)), None, ), ] @@ -500,7 +523,9 @@ def test_update_sqlmesh_comment_info( controller = make_controller( "tests/fixtures/github/pull_request_synchronized.json", github_client ) - updated, resp = controller.update_sqlmesh_comment_info(comment, dedup_regex=dedup_regex) + updated, resp = controller.update_sqlmesh_comment_info( + comment, dedup_regex=dedup_regex + ) assert resp.body == resulting_comment if create_comment is None: assert len(created_comments) == 0 @@ -564,7 +589,9 @@ def test_deploy_to_prod_dirty_pr(github_client, make_controller): controller.deploy_to_prod() -def test_try_invalidate_pr_environment(github_client, make_controller, mocker: MockerFixture): +def test_try_invalidate_pr_environment( + github_client, make_controller, mocker: MockerFixture +): invalidate_controller = make_controller( "tests/fixtures/github/pull_request_synchronized.json", github_client ) @@ -582,7 +609,9 @@ def test_try_invalidate_pr_environment(github_client, make_controller, mocker: M no_invalidate_controller._context._state_sync = mocker.MagicMock() no_invalidate_controller.try_invalidate_pr_environment() - assert not no_invalidate_controller._context._state_sync.invalidate_environment.called + assert ( + not no_invalidate_controller._context._state_sync.invalidate_environment.called + ) def test_try_merge_pr(github_client, make_controller, mocker: MockerFixture): @@ -657,13 +686,21 @@ def test_try_merge_pr(github_client, make_controller, mocker: MockerFixture): ids=[test[0] for test in bot_command_parsing_params], ) def test_bot_command_parsing( - action, comment, namespace, command, github_client, make_controller, make_event_issue_comment + action, + comment, + namespace, + command, + github_client, + make_controller, + make_event_issue_comment, ): controller = make_controller( make_event_issue_comment(action, comment), github_client, bot_config=GithubCICDBotConfig( - command_namespace=namespace, enable_deploy_command=True, merge_method=MergeMethod.SQUASH + command_namespace=namespace, + enable_deploy_command=True, + merge_method=MergeMethod.SQUASH, ), ) assert controller.get_command_from_comment() == command @@ -678,7 +715,9 @@ def test_uncategorized( make_mock_issue_comment, tmp_path: pathlib.Path, ): - snapshot_uncategorized = make_snapshot(SqlModel(name="b", query=parse_one("select 1, ds"))) + snapshot_uncategorized = make_snapshot( + SqlModel(name="b", query=parse_one("select 1, ds")) + ) mocker.patch( "sqlmesh.core.plan.Plan.uncategorized", PropertyMock( @@ -741,8 +780,7 @@ def test_get_plan_summary_doesnt_truncate_backfill_list( assert "more ...." not in summary - assert ( - """**Models needing backfill:** + assert """**Models needing backfill:** * `memory.raw.demographics`: [full refresh] * `memory.sushi.active_customers`: [full refresh] * `memory.sushi.count_customers_active`: [full refresh] @@ -759,9 +797,7 @@ def test_get_plan_summary_doesnt_truncate_backfill_list( * `memory.sushi.top_waiters`: [recreate view] * `memory.sushi.waiter_as_customer_by_day`: [2025-06-30 - 2025-07-06] * `memory.sushi.waiter_names`: [full refresh] -* `memory.sushi.waiter_revenue_by_day`: [2025-06-30 - 2025-07-06]""" - in summary - ) +* `memory.sushi.waiter_revenue_by_day`: [2025-06-30 - 2025-07-06]""" in summary def test_get_plan_summary_includes_warnings_and_errors( @@ -780,7 +816,9 @@ def test_get_plan_summary_includes_warnings_and_errors( summary = controller.get_plan_summary(controller.prod_plan) - assert ("> [!WARNING]\n>\n> - Warning 1\n> With multiline\n>\n> - Warning 2\n>\n>") in summary + assert ( + "> [!WARNING]\n>\n> - Warning 1\n> With multiline\n>\n> - Warning 2\n>\n>" + ) in summary assert ( "> Linter warnings for `tests/linter_test.sql`:\n> - mock_linter_rule: Linter warning\n>" ) in summary @@ -822,7 +860,8 @@ def test_get_pr_environment_summary_includes_warnings_and_errors( # completed with an exception triggers a FAILED conclusion and shows errors error_summary = controller.get_pr_environment_summary( - conclusion=GithubCheckConclusion.FAILURE, exception=SQLMeshError("Something broke") + conclusion=GithubCheckConclusion.FAILURE, + exception=SQLMeshError("Something broke"), ) assert "> [!WARNING]\n>\n> - Warning 1\n>\n" in error_summary assert ( @@ -871,7 +910,10 @@ def test_pr_comment_deploy_indicator_includes_command_namespace( comment = created_comments[0].body assert "To **apply** this PR's plan to prod, comment:\n - `/deploy`" not in comment - assert "To **apply** this PR's plan to prod, comment:\n - `#SQLMesh/deploy`" in comment + assert ( + "To **apply** this PR's plan to prod, comment:\n - `#SQLMesh/deploy`" + in comment + ) def test_forward_only_config_falls_back_to_plan_config( @@ -920,7 +962,9 @@ def test_forward_only_config_falls_back_to_plan_config( def test_chunk_up_api_message_preserves_ascii_behavior( github_client, make_event_issue_comment, make_controller ): - controller = make_controller(make_event_issue_comment("created", "test"), github_client) + controller = make_controller( + make_event_issue_comment("created", "test"), github_client + ) max_bytes = controller.MAX_BYTE_LENGTH message = ("a" * max_bytes) + ("b" * max_bytes) + "c" @@ -937,7 +981,9 @@ def test_chunk_up_api_message_preserves_ascii_behavior( def test_chunk_up_api_message_preserves_multibyte_chars( github_client, make_event_issue_comment, make_controller, multibyte_char ): - controller = make_controller(make_event_issue_comment("created", "test"), github_client) + controller = make_controller( + make_event_issue_comment("created", "test"), github_client + ) max_bytes = controller.MAX_BYTE_LENGTH # Place a multibyte character exactly on the byte boundary between two chunks so diff --git a/tests/integrations/github/cicd/test_github_event.py b/tests/integrations/github/cicd/test_github_event.py index 88979b05a8..76b084756c 100644 --- a/tests/integrations/github/cicd/test_github_event.py +++ b/tests/integrations/github/cicd/test_github_event.py @@ -4,26 +4,41 @@ def test_pull_request_review_submit_event(make_event_from_fixture): - event = make_event_from_fixture("tests/fixtures/github/pull_request_review_submit.json") + event = make_event_from_fixture( + "tests/fixtures/github/pull_request_review_submit.json" + ) assert not event.is_pull_request assert event.is_review - assert event.pull_request_url == "https://api.github.com/repos/Codertocat/Hello-World/pulls/2" + assert ( + event.pull_request_url + == "https://api.github.com/repos/Codertocat/Hello-World/pulls/2" + ) def test_pull_request_synchronized_event(make_event_from_fixture): - event = make_event_from_fixture("tests/fixtures/github/pull_request_synchronized.json") + event = make_event_from_fixture( + "tests/fixtures/github/pull_request_synchronized.json" + ) assert event.is_pull_request - assert event.pull_request_url == "https://api.github.com/repos/Codertocat/Hello-World/pulls/2" + assert ( + event.pull_request_url + == "https://api.github.com/repos/Codertocat/Hello-World/pulls/2" + ) def test_github_pull_request_comment(make_event_from_fixture): event = make_event_from_fixture("tests/fixtures/github/pull_request_comment.json") assert event.is_comment - assert event.pull_request_url == "https://api.github.com/repos/Codertocat/Hello-World/pulls/2" + assert ( + event.pull_request_url + == "https://api.github.com/repos/Codertocat/Hello-World/pulls/2" + ) assert event.pull_request_comment_body == "example_comment" -def test_pull_request_synchronized_info(make_event_from_fixture, make_pull_request_info): +def test_pull_request_synchronized_info( + make_event_from_fixture, make_pull_request_info +): pull_request_info = make_pull_request_info( make_event_from_fixture("tests/fixtures/github/pull_request_synchronized.json") ) @@ -33,9 +48,13 @@ def test_pull_request_synchronized_info(make_event_from_fixture, make_pull_reque assert pull_request_info.full_repo_path == "Codertocat/Hello-World" -def test_pull_request_synchronized_enterprise(make_event_from_fixture, make_pull_request_info): +def test_pull_request_synchronized_enterprise( + make_event_from_fixture, make_pull_request_info +): pull_request_info = make_pull_request_info( - make_event_from_fixture("tests/fixtures/github/pull_request_synchronized_enterprise.json") + make_event_from_fixture( + "tests/fixtures/github/pull_request_synchronized_enterprise.json" + ) ) assert pull_request_info.owner == "org" assert pull_request_info.repo == "repo" diff --git a/tests/integrations/github/cicd/test_integration.py b/tests/integrations/github/cicd/test_integration.py index ce357f6d36..4274323520 100644 --- a/tests/integrations/github/cicd/test_integration.py +++ b/tests/integrations/github/cicd/test_integration.py @@ -13,17 +13,17 @@ from pytest_mock.plugin import MockerFixture from sqlglot import exp -from sqlmesh.core.config import CategorizerConfig, Config, ModelDefaultsConfig, LinterConfig +from sqlmesh.core.config import (CategorizerConfig, Config, LinterConfig, + ModelDefaultsConfig) from sqlmesh.core.engine_adapter.shared import DataObject -from sqlmesh.core.user import User, UserRole from sqlmesh.core.model.common import ParsableSql +from sqlmesh.core.user import User, UserRole from sqlmesh.integrations.github.cicd import command -from sqlmesh.integrations.github.cicd.config import GithubCICDBotConfig, MergeMethod -from sqlmesh.integrations.github.cicd.controller import ( - GithubCheckConclusion, - GithubCheckStatus, - GithubController, -) +from sqlmesh.integrations.github.cicd.config import (GithubCICDBotConfig, + MergeMethod) +from sqlmesh.integrations.github.cicd.controller import (GithubCheckConclusion, + GithubCheckStatus, + GithubController) from sqlmesh.utils.errors import CICDBotError, SQLMeshError from tests.integrations.github.cicd.conftest import MockIssueComment @@ -33,11 +33,15 @@ ] -def get_environment_objects(controller: GithubController, environment: str) -> t.List[DataObject]: +def get_environment_objects( + controller: GithubController, environment: str +) -> t.List[DataObject]: return controller._context.engine_adapter.get_data_objects(f"sushi__{environment}") -def get_num_days_loaded(controller: GithubController, environment: str, model: str) -> int: +def get_num_days_loaded( + controller: GithubController, environment: str, model: str +) -> int: return controller._context.engine_adapter.fetchdf( f"SELECT distinct event_date FROM sushi__{environment}.{model}" ).shape[0] @@ -83,7 +87,9 @@ def test_linter( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -227,7 +233,9 @@ def test_merge_pr_has_non_breaking_change( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -246,13 +254,19 @@ def test_merge_pr_has_non_breaking_change( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # Make a non-breaking change model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) github_output_file = tmp_path / "github_output.txt" @@ -261,7 +275,9 @@ def test_merge_pr_has_non_breaking_change( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -274,27 +290,26 @@ def test_merge_pr_has_non_breaking_change( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Directly Modified + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Directly Modified - `memory.sushi.waiter_revenue_by_day` (Non-breaking) **Kind:** INCREMENTAL_BY_TIME_RANGE - **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" - in pr_env_summary - ) - assert ( - """### Indirectly Modified + **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" in pr_env_summary + assert """### Indirectly Modified - `memory.sushi.top_waiters` (Indirect Non-breaking) - **Kind:** VIEW [recreate view]""" - in pr_env_summary - ) + **Kind:** VIEW [recreate view]""" in pr_env_summary assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -304,7 +319,9 @@ def test_merge_pr_has_non_breaking_change( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success expected_prod_plan_directly_modified_summary = """**Directly Modified:** * `memory.sushi.waiter_revenue_by_day` (Non-breaking) @@ -343,7 +360,9 @@ def test_merge_pr_has_non_breaking_change( assert expected_prod_plan_indirectly_modified_summary in prod_plan_preview_summary assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -352,8 +371,13 @@ def test_merge_pr_has_non_breaking_change( assert prod_checks_runs[2]["output"]["title"] == "Deployed to Prod" prod_environment_synced_summary = prod_checks_runs[2]["output"]["summary"] assert "**Generated Prod Plan**" in prod_environment_synced_summary - assert expected_prod_plan_directly_modified_summary in prod_environment_synced_summary - assert expected_prod_plan_indirectly_modified_summary in prod_environment_synced_summary + assert ( + expected_prod_plan_directly_modified_summary in prod_environment_synced_summary + ) + assert ( + expected_prod_plan_indirectly_modified_summary + in prod_environment_synced_summary + ) assert "SQLMesh - Has Required Approval" in controller._check_run_mapping approval_checks_runs = controller._check_run_mapping[ @@ -376,20 +400,21 @@ def test_merge_pr_has_non_breaking_change( ) assert len(get_environment_objects(controller, "hello_world_2")) == 2 - assert get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 - assert "new_col" in get_columns(controller, "hello_world_2", "waiter_revenue_by_day") + assert ( + get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 + ) + assert "new_col" in get_columns( + controller, "hello_world_2", "waiter_revenue_by_day" + ) assert "new_col" in get_columns(controller, None, "waiter_revenue_by_day") assert mock_pull_request.merge.called assert len(created_comments) == 1 comment_body = created_comments[0].body - assert ( - """:robot: **SQLMesh Bot Info** :robot: + assert """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - - `hello_world_2`""" - in comment_body - ) + - `hello_world_2`""" in comment_body assert expected_prod_plan_directly_modified_summary in comment_body assert expected_prod_plan_indirectly_modified_summary in comment_body @@ -438,7 +463,9 @@ def test_merge_pr_has_non_breaking_change_diff_start( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -457,13 +484,19 @@ def test_merge_pr_has_non_breaking_change_diff_start( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # Make a non-breaking change model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) github_output_file = tmp_path / "github_output.txt" @@ -472,7 +505,9 @@ def test_merge_pr_has_non_breaking_change_diff_start( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -485,28 +520,27 @@ def test_merge_pr_has_non_breaking_change_diff_start( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Directly Modified + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Directly Modified - `memory.sushi.waiter_revenue_by_day` (Non-breaking) **Kind:** INCREMENTAL_BY_TIME_RANGE **Dates loaded in PR:** [2022-12-29 - 2022-12-31] - **Dates *not* loaded in PR:** [2022-12-25 - 2022-12-28]""" - in pr_env_summary - ) - assert ( - """### Indirectly Modified + **Dates *not* loaded in PR:** [2022-12-25 - 2022-12-28]""" in pr_env_summary + assert """### Indirectly Modified - `memory.sushi.top_waiters` (Indirect Non-breaking) - **Kind:** VIEW [recreate view]""" - in pr_env_summary - ) + **Kind:** VIEW [recreate view]""" in pr_env_summary assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -516,7 +550,9 @@ def test_merge_pr_has_non_breaking_change_diff_start( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert prod_plan_preview_checks_runs[2]["output"]["title"] == "Prod Plan Preview" expected_prod_plan_directly_modified_summary = """**Directly Modified:** @@ -555,7 +591,9 @@ def test_merge_pr_has_non_breaking_change_diff_start( assert expected_prod_plan_indirectly_modified_summary in prod_plan_preview_summary assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -564,8 +602,13 @@ def test_merge_pr_has_non_breaking_change_diff_start( assert prod_checks_runs[2]["output"]["title"] == "Deployed to Prod" prod_environment_synced_summary = prod_checks_runs[2]["output"]["summary"] assert "**Generated Prod Plan**" in prod_environment_synced_summary - assert expected_prod_plan_directly_modified_summary in prod_environment_synced_summary - assert expected_prod_plan_indirectly_modified_summary in prod_environment_synced_summary + assert ( + expected_prod_plan_directly_modified_summary in prod_environment_synced_summary + ) + assert ( + expected_prod_plan_indirectly_modified_summary + in prod_environment_synced_summary + ) assert "SQLMesh - Has Required Approval" in controller._check_run_mapping approval_checks_runs = controller._check_run_mapping[ @@ -589,20 +632,21 @@ def test_merge_pr_has_non_breaking_change_diff_start( assert len(get_environment_objects(controller, "hello_world_2")) == 2 # 7 days since the prod deploy went through and backfilled the remaining days - assert get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 - assert "new_col" in get_columns(controller, "hello_world_2", "waiter_revenue_by_day") + assert ( + get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 + ) + assert "new_col" in get_columns( + controller, "hello_world_2", "waiter_revenue_by_day" + ) assert "new_col" in get_columns(controller, None, "waiter_revenue_by_day") assert mock_pull_request.merge.called assert len(created_comments) == 1 comment_body = created_comments[0].body - assert ( - """:robot: **SQLMesh Bot Info** :robot: + assert """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - - `hello_world_2`""" - in comment_body - ) + - `hello_world_2`""" in comment_body assert expected_prod_plan_directly_modified_summary in comment_body assert expected_prod_plan_indirectly_modified_summary in comment_body @@ -651,7 +695,9 @@ def test_merge_pr_has_non_breaking_change_no_categorization( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -667,13 +713,19 @@ def test_merge_pr_has_non_breaking_change_no_categorization( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # Make a non-breaking change model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) github_output_file = tmp_path / "github_output.txt" @@ -683,7 +735,9 @@ def test_merge_pr_has_non_breaking_change_no_categorization( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -696,13 +750,18 @@ def test_merge_pr_has_non_breaking_change_no_categorization( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_action_required - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" + assert ( + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( """:warning: Action Required to create or update PR Environment `hello_world_2` :warning: @@ -721,7 +780,9 @@ def test_merge_pr_has_non_breaking_change_no_categorization( assert len(prod_plan_preview_checks_runs) == 2 assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[1]["conclusion"]).is_skipped + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[1]["conclusion"] + ).is_skipped assert ( prod_plan_preview_checks_runs[1]["output"]["title"] == "Skipped generating prod plan preview since PR was not synchronized" @@ -732,7 +793,9 @@ def test_merge_pr_has_non_breaking_change_no_categorization( ) assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed @@ -814,7 +877,9 @@ def test_merge_pr_has_no_changes( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -829,7 +894,11 @@ def test_merge_pr_has_no_changes( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] github_output_file = tmp_path / "github_output.txt" @@ -838,7 +907,9 @@ def test_merge_pr_has_no_changes( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -851,13 +922,18 @@ def test_merge_pr_has_no_changes( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_skipped - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" + assert ( + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) assert ( ":next_track_button: Skipped creating or updating PR Environment `hello_world_2` :next_track_button:\n\nNo changes were detected compared to the prod environment." in pr_checks_runs[2]["output"]["summary"] @@ -871,15 +947,22 @@ def test_merge_pr_has_no_changes( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success expected_prod_plan_summary = ( "**No changes to plan: project files match the `prod` environment**" ) assert prod_plan_preview_checks_runs[2]["output"]["title"] == "Prod Plan Preview" - assert expected_prod_plan_summary in prod_plan_preview_checks_runs[2]["output"]["summary"] + assert ( + expected_prod_plan_summary + in prod_plan_preview_checks_runs[2]["output"]["summary"] + ) assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -916,12 +999,9 @@ def test_merge_pr_has_no_changes( assert len(created_comments) == 1 comment_body = created_comments[0].body - assert ( - f""":robot: **SQLMesh Bot Info** :robot: + assert f""":robot: **SQLMesh Bot Info** :robot:
- :ship: Prod Plan Being Applied""" - in comment_body - ) + :ship: Prod Plan Being Applied""" in comment_body assert expected_prod_plan_summary in comment_body with open(github_output_file, "r", encoding="utf-8") as f: @@ -986,13 +1066,19 @@ def test_no_merge_since_no_deploy_signal( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # Make a non-breaking change model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) github_output_file = tmp_path / "github_output.txt" @@ -1001,7 +1087,9 @@ def test_no_merge_since_no_deploy_signal( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -1014,27 +1102,26 @@ def test_no_merge_since_no_deploy_signal( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Directly Modified + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Directly Modified - `memory.sushi.waiter_revenue_by_day` (Non-breaking) **Kind:** INCREMENTAL_BY_TIME_RANGE - **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" - in pr_env_summary - ) - assert ( - """### Indirectly Modified + **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" in pr_env_summary + assert """### Indirectly Modified - `memory.sushi.top_waiters` (Indirect Non-breaking) - **Kind:** VIEW [recreate view]""" - in pr_env_summary - ) + **Kind:** VIEW [recreate view]""" in pr_env_summary assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -1044,7 +1131,9 @@ def test_no_merge_since_no_deploy_signal( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success expected_prod_plan_directly_modified_summary = """**Directly Modified:** * `memory.sushi.waiter_revenue_by_day` (Non-breaking) @@ -1083,7 +1172,9 @@ def test_no_merge_since_no_deploy_signal( assert expected_prod_plan_indirectly_modified_summary in prod_plan_preview_summary assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed @@ -1112,19 +1203,20 @@ def test_no_merge_since_no_deploy_signal( ) assert len(get_environment_objects(controller, "hello_world_2")) == 2 - assert get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 - assert "new_col" in get_columns(controller, "hello_world_2", "waiter_revenue_by_day") + assert ( + get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 + ) + assert "new_col" in get_columns( + controller, "hello_world_2", "waiter_revenue_by_day" + ) assert "new_col" not in get_columns(controller, None, "waiter_revenue_by_day") assert not mock_pull_request.merge.called assert len(created_comments) == 1 - assert ( - created_comments[0].body - == """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body == """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`""" - ) with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() @@ -1171,7 +1263,9 @@ def test_no_merge_since_no_deploy_signal_no_approvers_defined( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -1189,12 +1283,16 @@ def test_no_merge_since_no_deploy_signal_no_approvers_defined( mock_out_context=False, ) controller._context.plan("prod", no_prompts=True, auto_apply=True) - controller._context.users = [User(username="test", github_username="test_github", roles=[])] + controller._context.users = [ + User(username="test", github_username="test_github", roles=[]) + ] # Make a non-breaking change model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) github_output_file = tmp_path / "github_output.txt" @@ -1203,7 +1301,9 @@ def test_no_merge_since_no_deploy_signal_no_approvers_defined( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -1216,28 +1316,27 @@ def test_no_merge_since_no_deploy_signal_no_approvers_defined( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Directly Modified + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Directly Modified - `memory.sushi.waiter_revenue_by_day` (Non-breaking) **Kind:** INCREMENTAL_BY_TIME_RANGE **Dates loaded in PR:** [2022-12-30 - 2022-12-31] - **Dates *not* loaded in PR:** [2022-12-25 - 2022-12-29]""" - in pr_env_summary - ) - assert ( - """### Indirectly Modified + **Dates *not* loaded in PR:** [2022-12-25 - 2022-12-29]""" in pr_env_summary + assert """### Indirectly Modified - `memory.sushi.top_waiters` (Indirect Non-breaking) - **Kind:** VIEW [recreate view]""" - in pr_env_summary - ) + **Kind:** VIEW [recreate view]""" in pr_env_summary assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -1247,7 +1346,9 @@ def test_no_merge_since_no_deploy_signal_no_approvers_defined( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success expected_prod_plan_directly_modified_summary = """**Directly Modified:** * `memory.sushi.waiter_revenue_by_day` (Non-breaking) @@ -1287,19 +1388,20 @@ def test_no_merge_since_no_deploy_signal_no_approvers_defined( assert "SQLMesh - Has Required Approval" not in controller._check_run_mapping assert len(get_environment_objects(controller, "hello_world_2")) == 2 - assert get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 2 - assert "new_col" in get_columns(controller, "hello_world_2", "waiter_revenue_by_day") + assert ( + get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 2 + ) + assert "new_col" in get_columns( + controller, "hello_world_2", "waiter_revenue_by_day" + ) assert "new_col" not in get_columns(controller, None, "waiter_revenue_by_day") assert not mock_pull_request.merge.called assert len(created_comments) == 1 - assert ( - created_comments[0].body - == """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body == """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`""" - ) with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() @@ -1347,7 +1449,9 @@ def test_deploy_comment_pre_categorized( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -1365,12 +1469,16 @@ def test_deploy_comment_pre_categorized( mock_out_context=False, ) controller._context.plan("prod", no_prompts=True, auto_apply=True) - controller._context.users = [User(username="test", github_username="test_github", roles=[])] + controller._context.users = [ + User(username="test", github_username="test_github", roles=[]) + ] # Make a non-breaking change model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) # Manually categorize the change as non-breaking and don't backfill anything @@ -1388,7 +1496,9 @@ def test_deploy_comment_pre_categorized( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -1401,27 +1511,26 @@ def test_deploy_comment_pre_categorized( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Directly Modified + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Directly Modified - `memory.sushi.waiter_revenue_by_day` (Non-breaking) **Kind:** INCREMENTAL_BY_TIME_RANGE - **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" - in pr_env_summary - ) - assert ( - """### Indirectly Modified + **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" in pr_env_summary + assert """### Indirectly Modified - `memory.sushi.top_waiters` (Indirect Non-breaking) - **Kind:** VIEW [recreate view]""" - in pr_env_summary - ) + **Kind:** VIEW [recreate view]""" in pr_env_summary assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -1431,7 +1540,9 @@ def test_deploy_comment_pre_categorized( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success expected_prod_plan_directly_modified_summary = """**Directly Modified:** * `memory.sushi.waiter_revenue_by_day` (Non-breaking) @@ -1468,7 +1579,9 @@ def test_deploy_comment_pre_categorized( assert expected_prod_plan_indirectly_modified_summary in prod_plan_preview_summary assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -1477,30 +1590,36 @@ def test_deploy_comment_pre_categorized( assert prod_checks_runs[2]["output"]["title"] == "Deployed to Prod" prod_environment_synced_summary = prod_checks_runs[2]["output"]["summary"] assert "**Generated Prod Plan**" in prod_environment_synced_summary - assert expected_prod_plan_directly_modified_summary in prod_environment_synced_summary - assert expected_prod_plan_indirectly_modified_summary in prod_environment_synced_summary + assert ( + expected_prod_plan_directly_modified_summary in prod_environment_synced_summary + ) + assert ( + expected_prod_plan_indirectly_modified_summary + in prod_environment_synced_summary + ) assert "SQLMesh - Has Required Approval" not in controller._check_run_mapping assert len(get_environment_objects(controller, "hello_world_2")) == 2 - assert get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 - assert "new_col" in get_columns(controller, "hello_world_2", "waiter_revenue_by_day") + assert ( + get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 + ) + assert "new_col" in get_columns( + controller, "hello_world_2", "waiter_revenue_by_day" + ) assert "new_col" in get_columns(controller, None, "waiter_revenue_by_day") assert mock_pull_request.merge.called assert len(created_comments) == 1 comment_body = created_comments[0].body - assert ( - """:robot: **SQLMesh Bot Info** :robot: + assert """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2` - :arrow_forward: To **apply** this PR's plan to prod, comment: - `/deploy`
- :ship: Prod Plan Being Applied""" - in comment_body - ) + :ship: Prod Plan Being Applied""" in comment_body assert expected_prod_plan_directly_modified_summary in comment_body assert expected_prod_plan_indirectly_modified_summary in comment_body @@ -1549,7 +1668,9 @@ def test_error_msg_when_applying_plan_with_bug( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -1566,14 +1687,20 @@ def test_error_msg_when_applying_plan_with_bug( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # Make an error by adding a column that doesn't exist model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, query_=ParsableSql( - sql=model.query.select(exp.alias_("non_existing_col", "new_col")).sql(model.dialect) + sql=model.query.select(exp.alias_("non_existing_col", "new_col")).sql( + model.dialect + ) ), ) @@ -1584,7 +1711,9 @@ def test_error_msg_when_applying_plan_with_bug( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -1597,17 +1726,25 @@ def test_error_msg_when_applying_plan_with_bug( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_failure - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" + assert ( + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) summary = pr_checks_runs[2]["output"]["summary"].replace("\n", "") assert '**Skipped models*** `"memory"."sushi"."top_waiters"`' in summary assert '**Failed models*** `"memory"."sushi"."waiter_revenue_by_day"`' in summary - assert 'Binder Error: Referenced column "non_existing_col" not found in FROM clause!' in summary + assert ( + 'Binder Error: Referenced column "non_existing_col" not found in FROM clause!' + in summary + ) assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -1616,7 +1753,9 @@ def test_error_msg_when_applying_plan_with_bug( assert len(prod_plan_preview_checks_runs) == 2 assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[1]["conclusion"]).is_skipped + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[1]["conclusion"] + ).is_skipped assert ( prod_plan_preview_checks_runs[1]["output"]["title"] == "Skipped generating prod plan preview since PR was not synchronized" @@ -1627,7 +1766,9 @@ def test_error_msg_when_applying_plan_with_bug( ) assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 2 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_completed @@ -1708,7 +1849,9 @@ def test_overlapping_changes_models( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -1727,7 +1870,11 @@ def test_overlapping_changes_models( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # These changes have shared children and this ensures we don't repeat the children in the output @@ -1735,7 +1882,9 @@ def test_overlapping_changes_models( model = controller._context.get_model("sushi.customers").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) # Make a breaking change @@ -1749,7 +1898,9 @@ def test_overlapping_changes_models( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -1762,25 +1913,26 @@ def test_overlapping_changes_models( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Directly Modified + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Directly Modified - `memory.sushi.customers` (Non-breaking) **Kind:** FULL [full refresh] - `memory.sushi.waiter_names` (Breaking) - **Kind:** SEED [full refresh]""" - in pr_env_summary - ) - assert ( - """### Indirectly Modified + **Kind:** SEED [full refresh]""" in pr_env_summary + assert """### Indirectly Modified - `memory.sushi.active_customers` (Indirect Non-breaking) **Kind:** CUSTOM [full refresh] @@ -1792,9 +1944,7 @@ def test_overlapping_changes_models( - `memory.sushi.waiter_as_customer_by_day` (Indirect Breaking) **Kind:** INCREMENTAL_BY_TIME_RANGE - **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" - in pr_env_summary - ) + **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" in pr_env_summary assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -1804,7 +1954,9 @@ def test_overlapping_changes_models( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success expected_prod_plan_directly_modified_summary = """**Directly Modified:** * `memory.sushi.customers` (Non-breaking) @@ -1855,7 +2007,9 @@ def test_overlapping_changes_models( assert expected_prod_plan_indirectly_modified_summary in prod_plan_preview_summary assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -1864,8 +2018,13 @@ def test_overlapping_changes_models( assert prod_checks_runs[2]["output"]["title"] == "Deployed to Prod" prod_environment_synced_summary = prod_checks_runs[2]["output"]["summary"] assert "**Generated Prod Plan**" in prod_environment_synced_summary - assert expected_prod_plan_directly_modified_summary in prod_environment_synced_summary - assert expected_prod_plan_indirectly_modified_summary in prod_environment_synced_summary + assert ( + expected_prod_plan_directly_modified_summary in prod_environment_synced_summary + ) + assert ( + expected_prod_plan_indirectly_modified_summary + in prod_environment_synced_summary + ) assert "SQLMesh - Has Required Approval" in controller._check_run_mapping approval_checks_runs = controller._check_run_mapping[ @@ -1894,14 +2053,11 @@ def test_overlapping_changes_models( assert len(created_comments) == 1 comment_body = created_comments[0].body - assert ( - f""":robot: **SQLMesh Bot Info** :robot: + assert f""":robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`
- :ship: Prod Plan Being Applied""" - in comment_body - ) + :ship: Prod Plan Being Applied""" in comment_body assert expected_prod_plan_directly_modified_summary in comment_body assert expected_prod_plan_indirectly_modified_summary in comment_body @@ -1950,7 +2106,9 @@ def test_pr_add_model( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -1970,16 +2128,14 @@ def test_pr_add_model( controller._context.plan("prod", no_prompts=True, auto_apply=True) # Add a model - (controller._context.path / "models" / "cicd_test_model.sql").write_text( - """ + (controller._context.path / "models" / "cicd_test_model.sql").write_text(""" MODEL ( name sushi.cicd_test_model, kind FULL ); select 1; - """ - ) + """) controller._context.load() assert '"memory"."sushi"."cicd_test_model"' in controller._context.models @@ -1989,7 +2145,9 @@ def test_pr_add_model( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -2003,20 +2161,22 @@ def test_pr_add_model( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Added -- `memory.sushi.cicd_test_model` (Breaking) - **Kind:** FULL [full refresh]""" - in pr_env_summary + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Added +- `memory.sushi.cicd_test_model` (Breaking) + **Kind:** FULL [full refresh]""" in pr_env_summary expected_prod_plan_summary = """**Added Models:** - `memory.sushi.cicd_test_model` (Breaking)""" @@ -2029,12 +2189,19 @@ def test_pr_add_model( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert prod_plan_preview_checks_runs[2]["output"]["title"] == "Prod Plan Preview" - assert expected_prod_plan_summary in prod_plan_preview_checks_runs[2]["output"]["summary"] + assert ( + expected_prod_plan_summary + in prod_plan_preview_checks_runs[2]["output"]["summary"] + ) assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -2049,16 +2216,13 @@ def test_pr_add_model( assert len(created_comments) == 1 comment_body = created_comments[0].body - assert ( - """:robot: **SQLMesh Bot Info** :robot: + assert """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2` - :arrow_forward: To **apply** this PR's plan to prod, comment: - `/deploy`
- :ship: Prod Plan Being Applied""" - in comment_body - ) + :ship: Prod Plan Being Applied""" in comment_body assert expected_prod_plan_summary in comment_body assert ( @@ -2104,7 +2268,9 @@ def test_pr_delete_model( mock_pull_request = mock_repo.get_pull() mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -2123,12 +2289,18 @@ def test_pr_delete_model( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # Remove a model model = controller._context.get_model("sushi.top_waiters").copy() del controller._context._models[model.fqn] - controller._context.dag = controller._context.dag.prune(*controller._context._models.keys()) + controller._context.dag = controller._context.dag.prune( + *controller._context._models.keys() + ) github_output_file = tmp_path / "github_output.txt" @@ -2141,7 +2313,9 @@ def test_pr_delete_model( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -2155,19 +2329,21 @@ def test_pr_delete_model( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Removed -- `memory.sushi.top_waiters` (Breaking)""" - in pr_env_summary + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Removed +- `memory.sushi.top_waiters` (Breaking)""" in pr_env_summary expected_prod_plan_summary = """**Removed Models:** - `memory.sushi.top_waiters` (Breaking)""" @@ -2180,12 +2356,19 @@ def test_pr_delete_model( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success assert prod_plan_preview_checks_runs[2]["output"]["title"] == "Prod Plan Preview" - assert expected_prod_plan_summary in prod_plan_preview_checks_runs[2]["output"]["summary"] + assert ( + expected_prod_plan_summary + in prod_plan_preview_checks_runs[2]["output"]["summary"] + ) assert "SQLMesh - Prod Environment Synced" in controller._check_run_mapping - prod_checks_runs = controller._check_run_mapping["SQLMesh - Prod Environment Synced"].all_kwargs + prod_checks_runs = controller._check_run_mapping[ + "SQLMesh - Prod Environment Synced" + ].all_kwargs assert len(prod_checks_runs) == 3 assert GithubCheckStatus(prod_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_checks_runs[1]["status"]).is_in_progress @@ -2222,14 +2405,11 @@ def test_pr_delete_model( assert len(created_comments) == 1 comment_body = created_comments[0].body - assert ( - """:robot: **SQLMesh Bot Info** :robot: + assert """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`
- :ship: Prod Plan Being Applied""" - in comment_body - ) + :ship: Prod Plan Being Applied""" in comment_body assert expected_prod_plan_summary in comment_body with open(github_output_file, "r", encoding="utf-8") as f: @@ -2279,7 +2459,9 @@ def test_has_required_approval_but_not_base_branch( mock_pull_request = mock_repo.get_pull() mock_pull_request.base.ref = "feature/branch" mock_pull_request.get_reviews = mocker.MagicMock( - side_effect=lambda: [make_pull_request_review(username="test_github", state="APPROVED")] + side_effect=lambda: [ + make_pull_request_review(username="test_github", state="APPROVED") + ] ) mock_pull_request.merged = False mock_pull_request.merge = mocker.MagicMock() @@ -2298,13 +2480,19 @@ def test_has_required_approval_but_not_base_branch( ) controller._context.plan("prod", no_prompts=True, auto_apply=True) controller._context.users = [ - User(username="test", github_username="test_github", roles=[UserRole.REQUIRED_APPROVER]) + User( + username="test", + github_username="test_github", + roles=[UserRole.REQUIRED_APPROVER], + ) ] # Make a non-breaking change model = controller._context.get_model("sushi.waiter_revenue_by_day").copy() controller._context.upsert_model( model, - query_=ParsableSql(sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect)), + query_=ParsableSql( + sql=model.query.select(exp.alias_("1", "new_col")).sql(model.dialect) + ), ) github_output_file = tmp_path / "github_output.txt" @@ -2313,7 +2501,9 @@ def test_has_required_approval_but_not_base_branch( command._run_all(controller) assert "SQLMesh - Run Unit Tests" in controller._check_run_mapping - test_checks_runs = controller._check_run_mapping["SQLMesh - Run Unit Tests"].all_kwargs + test_checks_runs = controller._check_run_mapping[ + "SQLMesh - Run Unit Tests" + ].all_kwargs assert len(test_checks_runs) == 3 assert GithubCheckStatus(test_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(test_checks_runs[1]["status"]).is_in_progress @@ -2326,27 +2516,26 @@ def test_has_required_approval_but_not_base_branch( ) assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping - pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs + pr_checks_runs = controller._check_run_mapping[ + "SQLMesh - PR Environment Synced" + ].all_kwargs assert len(pr_checks_runs) == 3 assert GithubCheckStatus(pr_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(pr_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(pr_checks_runs[2]["status"]).is_completed assert GithubCheckConclusion(pr_checks_runs[2]["conclusion"]).is_success - assert pr_checks_runs[2]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" - pr_env_summary = pr_checks_runs[2]["output"]["summary"] assert ( - """### Directly Modified + pr_checks_runs[2]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) + pr_env_summary = pr_checks_runs[2]["output"]["summary"] + assert """### Directly Modified - `memory.sushi.waiter_revenue_by_day` (Non-breaking) **Kind:** INCREMENTAL_BY_TIME_RANGE - **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" - in pr_env_summary - ) - assert ( - """### Indirectly Modified + **Dates loaded in PR:** [2022-12-25 - 2022-12-31]""" in pr_env_summary + assert """### Indirectly Modified - `memory.sushi.top_waiters` (Indirect Non-breaking) - **Kind:** VIEW [recreate view]""" - in pr_env_summary - ) + **Kind:** VIEW [recreate view]""" in pr_env_summary assert "SQLMesh - Prod Plan Preview" in controller._check_run_mapping prod_plan_preview_checks_runs = controller._check_run_mapping[ @@ -2356,7 +2545,9 @@ def test_has_required_approval_but_not_base_branch( assert GithubCheckStatus(prod_plan_preview_checks_runs[0]["status"]).is_queued assert GithubCheckStatus(prod_plan_preview_checks_runs[1]["status"]).is_in_progress assert GithubCheckStatus(prod_plan_preview_checks_runs[2]["status"]).is_completed - assert GithubCheckConclusion(prod_plan_preview_checks_runs[2]["conclusion"]).is_success + assert GithubCheckConclusion( + prod_plan_preview_checks_runs[2]["conclusion"] + ).is_success expected_prod_plan_directly_modified_summary = """**Directly Modified:** * `memory.sushi.waiter_revenue_by_day` (Non-breaking) @@ -2415,19 +2606,20 @@ def test_has_required_approval_but_not_base_branch( ) assert len(get_environment_objects(controller, "hello_world_2")) == 2 - assert get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 - assert "new_col" in get_columns(controller, "hello_world_2", "waiter_revenue_by_day") + assert ( + get_num_days_loaded(controller, "hello_world_2", "waiter_revenue_by_day") == 7 + ) + assert "new_col" in get_columns( + controller, "hello_world_2", "waiter_revenue_by_day" + ) assert "new_col" not in get_columns(controller, None, "waiter_revenue_by_day") assert not mock_pull_request.merge.called assert len(created_comments) == 1 - assert ( - created_comments[0].body - == """:robot: **SQLMesh Bot Info** :robot: + assert created_comments[0].body == """:robot: **SQLMesh Bot Info** :robot: - :eyes: To **review** this PR's changes, use virtual data environment: - `hello_world_2`""" - ) with open(github_output_file, "r", encoding="utf-8") as f: output = f.read() @@ -2480,7 +2672,10 @@ def test_unexpected_error_is_handled( assert "SQLMesh - PR Environment Synced" in controller._check_run_mapping pr_checks_runs = controller._check_run_mapping["SQLMesh - PR Environment Synced"].all_kwargs # type: ignore - assert pr_checks_runs[1]["output"]["title"] == "PR Virtual Data Environment: hello_world_2" + assert ( + pr_checks_runs[1]["output"]["title"] + == "PR Virtual Data Environment: hello_world_2" + ) summary = pr_checks_runs[1]["output"]["summary"] assert ( "**Error:** SQLGlot (local) is using version 'X' which is ahead of 'Y' (remote). Please run a migration" diff --git a/tests/integrations/jupyter/test_magics.py b/tests/integrations/jupyter/test_magics.py index c849dcbfc7..9c4f235c45 100644 --- a/tests/integrations/jupyter/test_magics.py +++ b/tests/integrations/jupyter/test_magics.py @@ -1,11 +1,12 @@ import logging import pathlib import typing as t +from pathlib import Path from unittest.mock import MagicMock import pytest -from bs4 import BeautifulSoup import time_machine +from bs4 import BeautifulSoup from hyperscript import h from IPython.core.error import UsageError from IPython.testing.globalipapp import start_ipython @@ -15,7 +16,6 @@ from sqlmesh import Context, RuntimeEnv from sqlmesh.magics import register_magics -from pathlib import Path logger = logging.getLogger(__name__) @@ -40,10 +40,14 @@ def ip(): @pytest.fixture def notebook(mocker: MockerFixture, ip): - mocker.patch("sqlmesh.RuntimeEnv.get", MagicMock(side_effect=lambda: RuntimeEnv.JUPYTER)) + mocker.patch( + "sqlmesh.RuntimeEnv.get", MagicMock(side_effect=lambda: RuntimeEnv.JUPYTER) + ) mocker.patch( "sqlmesh.core.console.RichConsole", - MagicMock(return_value=RichConsole(force_jupyter=True, color_system="truecolor")), + MagicMock( + return_value=RichConsole(force_jupyter=True, color_system="truecolor") + ), ) register_magics() return ip @@ -70,7 +74,8 @@ def loaded_sushi_context(sushi_context) -> Context: def convert_all_html_output_to_text(): def _convert(output: CapturedIO) -> t.List[str]: return [ - BeautifulSoup(output.data["text/html"]).get_text().strip() for output in output.outputs + BeautifulSoup(output.data["text/html"]).get_text().strip() + for output in output.outputs ] return _convert @@ -87,7 +92,9 @@ def _convert_html_to_tags(html: str) -> t.List[str]: ] def _convert(output: CapturedIO) -> t.List[t.List[str]]: - return [_convert_html_to_tags(output.data["text/html"]) for output in output.outputs] + return [ + _convert_html_to_tags(output.data["text/html"]) for output in output.outputs + ] return _convert @@ -101,10 +108,13 @@ def _convert(output: CapturedIO) -> t.List[str]: return _convert -def test_context(notebook, convert_all_html_output_to_text, get_all_html_output, tmp_path): +def test_context( + notebook, convert_all_html_output_to_text, get_all_html_output, tmp_path +): with capture_output() as output: notebook.run_line_magic( - magic_name="context", line=f"{str(SUSHI_EXAMPLE_PATH)} --log-file-dir {tmp_path}" + magic_name="context", + line=f"{str(SUSHI_EXAMPLE_PATH)} --log-file-dir {tmp_path}", ) assert output.stdout == "" @@ -134,7 +144,9 @@ def test_context(notebook, convert_all_html_output_to_text, get_all_html_output, def test_init(tmp_path, notebook, convert_all_html_output_to_text, get_all_html_output): with pytest.raises(UsageError, match="the following arguments are required: path"): notebook.run_line_magic(magic_name="init", line="") - with pytest.raises(UsageError, match="the following arguments are required: engine"): + with pytest.raises( + UsageError, match="the following arguments are required: engine" + ): notebook.run_line_magic(magic_name="init", line="foo") with capture_output() as output: notebook.run_line_magic(magic_name="init", line=f"{tmp_path} duckdb") @@ -142,7 +154,9 @@ def test_init(tmp_path, notebook, convert_all_html_output_to_text, get_all_html_ assert output.stdout == "" assert output.stderr == "" assert len(output.outputs) == 1 - assert convert_all_html_output_to_text(output) == ["SQLMesh project scaffold created"] + assert convert_all_html_output_to_text(output) == [ + "SQLMesh project scaffold created" + ] assert get_all_html_output(output) == [ str( h( @@ -161,7 +175,10 @@ def test_init(tmp_path, notebook, convert_all_html_output_to_text, get_all_html_ @pytest.mark.slow def test_render( - notebook, sushi_context, convert_all_html_output_to_text, convert_all_html_output_to_tags + notebook, + sushi_context, + convert_all_html_output_to_text, + convert_all_html_output_to_tags, ): with capture_output() as output: notebook.run_line_magic(magic_name="render", line="sushi.top_waiters") @@ -175,10 +192,15 @@ def test_render( @pytest.mark.slow def test_render_no_format( - notebook, sushi_context, convert_all_html_output_to_text, convert_all_html_output_to_tags + notebook, + sushi_context, + convert_all_html_output_to_text, + convert_all_html_output_to_tags, ): with capture_output() as output: - notebook.run_line_magic(magic_name="render", line="sushi.top_waiters --no-format") + notebook.run_line_magic( + magic_name="render", line="sushi.top_waiters --no-format" + ) assert output.stdout == "" assert output.stderr == "" @@ -215,26 +237,27 @@ def test_format(notebook, sushi_context): assert not output.stdout assert not output.stderr assert len(output.outputs) == 0 - assert ( - test_model_path.read_text() - == """MODEL ( + assert test_model_path.read_text() == """MODEL ( name db.test ); SELECT 1 AS foo FROM t""" - ) @pytest.mark.slow -def test_diff(sushi_context, notebook, convert_all_html_output_to_text, get_all_html_output): +def test_diff( + sushi_context, notebook, convert_all_html_output_to_text, get_all_html_output +): with capture_output(): test_model_path = sushi_context.path / "models" / "test_model.sql" test_model_path.write_text("MODEL(name sqlmesh_example.test); SELECT 1 AS foo") sushi_context.load() notebook.run_line_magic(magic_name="plan", line="--no-prompts --auto-apply") - test_model_path.write_text("MODEL(name sqlmesh_example.test); SELECT 1 AS foo, 2 AS bar") + test_model_path.write_text( + "MODEL(name sqlmesh_example.test); SELECT 1 AS foo, 2 AS bar" + ) with capture_output() as output: notebook.run_line_magic(magic_name="diff", line="prod") @@ -289,7 +312,10 @@ def test_diff(sushi_context, notebook, convert_all_html_output_to_text, get_all_ @pytest.mark.slow def test_plan( - notebook, sushi_context, convert_all_html_output_to_text, convert_all_html_output_to_tags + notebook, + sushi_context, + convert_all_html_output_to_text, + convert_all_html_output_to_tags, ): with capture_output() as output: notebook.run_line_magic(magic_name="plan", line="--no-prompts --auto-apply") @@ -336,7 +362,9 @@ def test_run_dag( assert any("[2K" in text for text in html_text_actual) assert any("Executing model batches" in text for text in html_text_actual) assert any("✔ Model batches executed" in text for text in html_text_actual) - assert any("Run finished for environment 'prod'" in text for text in html_text_actual) + assert any( + "Run finished for environment 'prod'" in text for text in html_text_actual + ) # Check the final messages final_outputs = [text for text in html_text_actual if text.strip()] @@ -349,7 +377,9 @@ def test_run_dag( pattern = r'font-weight: bold">0.\d{2}s ' import re - actual_html_output[i] = re.sub(pattern, 'font-weight: bold">0.00s ', chunk) + actual_html_output[i] = re.sub( + pattern, 'font-weight: bold">0.00s ', chunk + ) expected_html_output = [ str( h( @@ -372,7 +402,9 @@ def test_run_dag( {"style": RICH_PRE_STYLE}, h( "span", - {"style": "color: #000080; text-decoration-color: #000080; font-weight: bold"}, + { + "style": "color: #000080; text-decoration-color: #000080; font-weight: bold" + }, "Executing model batches", autoescape=False, ), @@ -672,9 +704,7 @@ def test_create_test(notebook, sushi_context): test_file = sushi_context.path / "tests" / "test_top_waiters.yaml" assert test_file.exists() - assert ( - test_file.read_text() - == """test_top_waiters: + assert test_file.read_text() == """test_top_waiters: model: '"memory"."sushi"."top_waiters"' inputs: '"memory"."sushi"."waiter_revenue_by_day"': @@ -682,7 +712,6 @@ def test_create_test(notebook, sushi_context): outputs: query: [] """ - ) def test_test(notebook, sushi_context): @@ -725,7 +754,9 @@ def test_audit(notebook, loaded_sushi_context, convert_all_html_output_to_text): def test_fetchdf(notebook, sushi_context): with capture_output() as output: - notebook.run_cell_magic(magic_name="fetchdf", line="my_result", cell="SELECT 1 AS foo") + notebook.run_cell_magic( + magic_name="fetchdf", line="my_result", cell="SELECT 1 AS foo" + ) assert not output.stdout assert not output.stderr @@ -734,7 +765,9 @@ def test_fetchdf(notebook, sushi_context): assert notebook.user_ns["my_result"].to_dict() == {"foo": {0: 1}} -def test_info(notebook, sushi_context, convert_all_html_output_to_text, get_all_html_output): +def test_info( + notebook, sushi_context, convert_all_html_output_to_text, get_all_html_output +): with capture_output() as output: notebook.run_line_magic(magic_name="info", line="--verbose") @@ -797,7 +830,9 @@ def test_migrate( @pytest.mark.slow -def test_create_external_models(notebook, loaded_sushi_context, convert_all_html_output_to_text): +def test_create_external_models( + notebook, loaded_sushi_context, convert_all_html_output_to_text +): external_model_file = loaded_sushi_context.path / "external_models.yaml" external_model_file.unlink() assert not external_model_file.exists() @@ -810,24 +845,25 @@ def test_create_external_models(notebook, loaded_sushi_context, convert_all_html assert len(output.outputs) == 0 assert external_model_file.exists() - assert ( - external_model_file.read_text() - == """- name: '"memory"."raw"."demographics"' + assert external_model_file.read_text() == """- name: '"memory"."raw"."demographics"' columns: customer_id: INT zip: TEXT gateway: duckdb """ - ) @pytest.mark.slow @time_machine.travel(FREEZE_TIME) def test_table_diff(notebook, loaded_sushi_context, convert_all_html_output_to_text): with capture_output(): - loaded_sushi_context.plan("dev", no_prompts=True, auto_apply=True, include_unmodified=True) + loaded_sushi_context.plan( + "dev", no_prompts=True, auto_apply=True, include_unmodified=True + ) with capture_output() as output: - notebook.run_line_magic(magic_name="table_diff", line="dev:prod --model sushi.top_waiters") + notebook.run_line_magic( + magic_name="table_diff", line="dev:prod --model sushi.top_waiters" + ) assert not output.stdout assert not output.stderr @@ -871,7 +907,9 @@ def test_lint(notebook, sushi_context): assert "Linter warnings for" in output.outputs[0].data["text/plain"] with capture_output() as output: - notebook.run_line_magic(magic_name="lint", line="--models sushi.items sushi.raw_marketing") + notebook.run_line_magic( + magic_name="lint", line="--models sushi.items sushi.raw_marketing" + ) assert len(output.outputs) == 2 assert "Linter warnings for" in output.outputs[0].data["text/plain"] diff --git a/tests/lsp/test_code_actions.py b/tests/lsp/test_code_actions.py index 509f49f9b1..d45444a411 100644 --- a/tests/lsp/test_code_actions.py +++ b/tests/lsp/test_code_actions.py @@ -1,6 +1,8 @@ -import typing as t import os +import typing as t + from lsprotocol import types + from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext from sqlmesh.lsp.uri import URI @@ -18,7 +20,11 @@ def test_code_actions_with_linting(copy_to_temp_path: t.Callable): with config_path.open("r") as f: lines = f.readlines() lines = [ - line.replace("enabled=False,", "enabled=True,") if "enabled=False," in line else line + ( + line.replace("enabled=False,", "enabled=True,") + if "enabled=False," in line + else line + ) for line in lines ] with config_path.open("w") as f: @@ -51,7 +57,9 @@ def test_code_actions_with_linting(copy_to_temp_path: t.Callable): lsp_context = LSPContext(context) # Get diagnostics (linting violations) - violations = lsp_context.lint_model(URI.from_path(sushi_path / "models" / "latest_order.sql")) + violations = lsp_context.lint_model( + URI.from_path(sushi_path / "models" / "latest_order.sql") + ) uri = URI.from_path(sushi_path / "models" / "latest_order.sql") @@ -102,7 +110,8 @@ def test_code_actions_with_linting(copy_to_temp_path: t.Callable): assert first_action.edit is not None assert first_action.edit.changes is not None assert ( - URI.from_path(sushi_path / "models" / "latest_order.sql").value in first_action.edit.changes + URI.from_path(sushi_path / "models" / "latest_order.sql").value + in first_action.edit.changes ) # The fix should replace SELECT * with specific columns @@ -167,7 +176,8 @@ def test_code_actions_create_file(copy_to_temp_path: t.Callable) -> None: params = types.CodeActionParams( text_document=types.TextDocumentIdentifier(uri=uri.value), range=types.Range( - start=types.Position(line=0, character=0), end=types.Position(line=1, character=0) + start=types.Position(line=0, character=0), + end=types.Position(line=1, character=0), ), context=types.CodeActionContext(diagnostics=diagnostics), ) @@ -177,6 +187,10 @@ def test_code_actions_create_file(copy_to_temp_path: t.Callable) -> None: action = next(a for a in actions if isinstance(a, types.CodeAction)) assert action.edit is not None assert action.edit.document_changes is not None - create_file = [c for c in action.edit.document_changes if isinstance(c, types.CreateFile)] + create_file = [ + c for c in action.edit.document_changes if isinstance(c, types.CreateFile) + ] assert create_file, "Expected a CreateFile operation" - assert create_file[0].uri == URI.from_path(sushi_path / "external_models.yaml").value + assert ( + create_file[0].uri == URI.from_path(sushi_path / "external_models.yaml").value + ) diff --git a/tests/lsp/test_completions.py b/tests/lsp/test_completions.py index e0772c1a96..dfbdd4ff64 100644 --- a/tests/lsp/test_completions.py +++ b/tests/lsp/test_completions.py @@ -1,14 +1,12 @@ from sqlglot import Tokenizer + from sqlmesh.core.context import Context -from sqlmesh.lsp.completions import ( - get_keywords_from_tokenizer, - get_sql_completions, - extract_keywords_from_content, -) +from sqlmesh.lsp.completions import (extract_keywords_from_content, + get_keywords_from_tokenizer, + get_sql_completions) from sqlmesh.lsp.context import LSPContext from sqlmesh.lsp.uri import URI - TOKENIZER_KEYWORDS = set(Tokenizer.KEYWORDS.keys()) @@ -26,7 +24,9 @@ def test_get_macros(): context = Context(paths=["examples/sushi"]) lsp_context = LSPContext(context) - file_path = next(key for key in lsp_context.map.keys() if key.name == "active_customers.sql") + file_path = next( + key for key in lsp_context.map.keys() if key.name == "active_customers.sql" + ) with open(file_path, "r", encoding="utf-8") as f: file_content = f.read() @@ -69,7 +69,9 @@ def test_get_sql_completions_with_context_and_file_uri(): context = Context(paths=["examples/sushi"]) lsp_context = LSPContext(context) - file_uri = next(key for key in lsp_context.map.keys() if key.name == "active_customers.sql") + file_uri = next( + key for key in lsp_context.map.keys() if key.name == "active_customers.sql" + ) completions = LSPContext.get_completions(lsp_context, URI.from_path(file_uri)) assert len(completions.keywords) > len(TOKENIZER_KEYWORDS) assert "sushi.active_customers" not in completions.models @@ -116,8 +118,12 @@ def test_get_sql_completions_with_file_content(): WHERE my_custom_column > 100 """ - file_uri = next(key for key in lsp_context.map.keys() if key.name == "active_customers.sql") - completions = LSPContext.get_completions(lsp_context, URI.from_path(file_uri), content) + file_uri = next( + key for key in lsp_context.map.keys() if key.name == "active_customers.sql" + ) + completions = LSPContext.get_completions( + lsp_context, URI.from_path(file_uri), content + ) # Check that SQL keywords are included assert any(k in ["SELECT", "FROM", "WHERE", "JOIN"] for k in completions.keywords) @@ -135,16 +141,20 @@ def test_get_sql_completions_with_file_content(): # Check that file keywords come after SQL keywords # SQL keywords should appear first in the list sql_keyword_indices = [ - i for i, k in enumerate(keywords_list) if k in ["SELECT", "FROM", "WHERE", "JOIN"] + i + for i, k in enumerate(keywords_list) + if k in ["SELECT", "FROM", "WHERE", "JOIN"] ] file_keyword_indices = [ - i for i, k in enumerate(keywords_list) if k in ["my_custom_column", "my_custom_table"] + i + for i, k in enumerate(keywords_list) + if k in ["my_custom_column", "my_custom_table"] ] if sql_keyword_indices and file_keyword_indices: - assert max(sql_keyword_indices) < min(file_keyword_indices), ( - "SQL keywords should come before file keywords" - ) + assert max(sql_keyword_indices) < min( + file_keyword_indices + ), "SQL keywords should come before file keywords" def test_get_sql_completions_with_partial_cte_query(): @@ -161,8 +171,12 @@ def test_get_sql_completions_with_partial_cte_query(): SELECT * FROM """ - file_uri = next(key for key in lsp_context.map.keys() if key.name == "active_customers.sql") - completions = LSPContext.get_completions(lsp_context, URI.from_path(file_uri), content) + file_uri = next( + key for key in lsp_context.map.keys() if key.name == "active_customers.sql" + ) + completions = LSPContext.get_completions( + lsp_context, URI.from_path(file_uri), content + ) # Check that CTE names are included in the keywords keywords_list = completions.keywords diff --git a/tests/lsp/test_diagnostics.py b/tests/lsp/test_diagnostics.py index 96167d47e5..ae50d8b3df 100644 --- a/tests/lsp/test_diagnostics.py +++ b/tests/lsp/test_diagnostics.py @@ -15,9 +15,11 @@ def test_diagnostic_on_sushi(tmp_path, copy_to_temp_path) -> None: with active_customers_path.open("r") as f: lines = f.readlines() lines = [ - line.replace("SELECT customer_id, zip", "SELECT *") - if "SELECT customer_id, zip" in line - else line + ( + line.replace("SELECT customer_id, zip", "SELECT *") + if "SELECT customer_id, zip" in line + else line + ) for line in lines ] with active_customers_path.open("w") as f: @@ -28,7 +30,11 @@ def test_diagnostic_on_sushi(tmp_path, copy_to_temp_path) -> None: with config_path.open("r") as f: lines = f.readlines() lines = [ - line.replace("enabled=False,", "enabled=True,") if "enabled=False," in line else line + ( + line.replace("enabled=False,", "enabled=True,") + if "enabled=False," in line + else line + ) for line in lines ] with config_path.open("w") as f: @@ -45,7 +51,9 @@ def test_diagnostic_on_sushi(tmp_path, copy_to_temp_path) -> None: assert len(lsp_diagnostics) > 0 # Get the no select star diagnostic - select_star_diagnostic = [diag for diag in lsp_diagnostics if diag.rule.name == "noselectstar"] + select_star_diagnostic = [ + diag for diag in lsp_diagnostics if diag.rule.name == "noselectstar" + ] assert len(select_star_diagnostic) == 1 diagnostic = select_star_diagnostic[0] diff --git a/tests/lsp/test_document_highlight.py b/tests/lsp/test_document_highlight.py index e6ce0ae7ec..68afe5ea1c 100644 --- a/tests/lsp/test_document_highlight.py +++ b/tests/lsp/test_document_highlight.py @@ -1,4 +1,4 @@ -from lsprotocol.types import Position, DocumentHighlightKind +from lsprotocol.types import DocumentHighlightKind, Position from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext, ModelTarget @@ -28,7 +28,9 @@ def test_get_document_highlights_cte(): assert len(ranges) >= 2 # Should have definition + usage # Test highlighting CTE definition - position on "current_marketing" definition - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) highlights = get_document_highlights(lsp_context, test_uri, position) assert highlights is not None @@ -40,7 +42,9 @@ def test_get_document_highlights_cte(): assert DocumentHighlightKind.Read in highlight_kinds # CTE usage # Test highlighting CTE usage - position on "current_marketing" usage - position = Position(line=ranges[1].start.line, character=ranges[1].start.character + 4) + position = Position( + line=ranges[1].start.line, character=ranges[1].start.character + 4 + ) highlights = get_document_highlights(lsp_context, test_uri, position) assert highlights is not None @@ -94,7 +98,9 @@ def test_get_document_highlights_multiple_ctes(): highlights = get_document_highlights(lsp_context, test_uri, position) assert highlights is not None - assert len(highlights) == len(outer_ranges) # Should match all occurrences of outer CTE + assert len(highlights) == len( + outer_ranges + ) # Should match all occurrences of outer CTE # Test the inner CTE - "current_marketing" (not outer) inner_ranges = find_ranges_from_regex(read_file, r"current_marketing(?!_outer)") diff --git a/tests/lsp/test_hints.py b/tests/lsp/test_hints.py index 99851a1361..5e669f1d64 100644 --- a/tests/lsp/test_hints.py +++ b/tests/lsp/test_hints.py @@ -1,12 +1,11 @@ """Tests for type hinting SQLMesh models""" import pytest - from sqlglot import exp, parse_one from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext, ModelTarget -from sqlmesh.lsp.hints import get_hints, _get_type_hints_for_model_from_query +from sqlmesh.lsp.hints import _get_type_hints_for_model_from_query, get_hints from sqlmesh.lsp.uri import URI @@ -24,12 +23,14 @@ def test_hints() -> None: customer_revenue_lifetime_path = next( path for path, info in lsp_context.map.items() - if isinstance(info, ModelTarget) and "sushi.customer_revenue_lifetime" in info.names + if isinstance(info, ModelTarget) + and "sushi.customer_revenue_lifetime" in info.names ) customer_revenue_by_day_path = next( path for path, info in lsp_context.map.items() - if isinstance(info, ModelTarget) and "sushi.customer_revenue_by_day" in info.names + if isinstance(info, ModelTarget) + and "sushi.customer_revenue_by_day" in info.names ) active_customers_uri = URI.from_path(active_customers_path) @@ -151,7 +152,9 @@ def test_alias_cast_hints() -> None: @pytest.mark.fast def test_simple_cte_hints() -> None: """Don't add type hints if the expression is already a cast""" - query = parse_one("WITH t AS (SELECT a FROM b) SELECT a AS c FROM t", dialect="postgres") + query = parse_one( + "WITH t AS (SELECT a FROM b) SELECT a AS c FROM t", dialect="postgres" + ) result = _get_type_hints_for_model_from_query( query=query, diff --git a/tests/lsp/test_reference.py b/tests/lsp/test_reference.py index 6aae4b869e..6e877763e1 100644 --- a/tests/lsp/test_reference.py +++ b/tests/lsp/test_reference.py @@ -1,7 +1,8 @@ from sqlmesh.core.context import Context from sqlmesh.core.linter.rule import Position -from sqlmesh.lsp.context import LSPContext, ModelTarget, AuditTarget -from sqlmesh.lsp.reference import ModelReference, get_model_definitions_for_a_path, by_position +from sqlmesh.lsp.context import AuditTarget, LSPContext, ModelTarget +from sqlmesh.lsp.reference import (ModelReference, by_position, + get_model_definitions_for_a_path) from sqlmesh.lsp.uri import URI @@ -63,9 +64,7 @@ def test_reference_with_alias() -> None: assert str(references[0].path).endswith("orders.py") assert get_string_from_range(read_file, references[0].range) == "sushi.orders" - assert ( - references[0].markdown_description - == """Table of sushi orders. + assert references[0].markdown_description == """Table of sushi orders. | Column | Type | Description | |--------|------|-------------| @@ -75,7 +74,6 @@ def test_reference_with_alias() -> None: | start_ts | INT | | | end_ts | INT | | | event_date | DATE | |""" - ) assert str(references[1].path).endswith("order_items.py") assert get_string_from_range(read_file, references[1].range) == "sushi.order_items" assert str(references[2].path).endswith("items.py") @@ -99,7 +97,9 @@ def test_standalone_audit_reference() -> None: if isinstance(info, ModelTarget) and "sushi.items" in info.names ) - references = get_model_definitions_for_a_path(lsp_context, URI.from_path(audit_path)) + references = get_model_definitions_for_a_path( + lsp_context, URI.from_path(audit_path) + ) assert len(references) == 1 assert references[0].path == items_path @@ -124,7 +124,9 @@ def get_string_from_range(file_lines, range_obj) -> str: return line_content[start_character:end_character] # Reference spans multiple lines - result = file_lines[start_line][start_character:] # First line from start_character to end + result = file_lines[start_line][ + start_character: + ] # First line from start_character to end for line_num in range(start_line + 1, end_line): # Middle lines (if any) result += file_lines[line_num] result += file_lines[end_line][:end_character] # Last line up to end_character @@ -157,7 +159,9 @@ def test_filter_references_by_position() -> None: for i, reference in enumerate(all_references): # Position inside the reference - should return exactly one reference middle_line = (reference.range.start.line + reference.range.end.line) // 2 - middle_char = (reference.range.start.character + reference.range.end.character) // 2 + middle_char = ( + reference.range.start.character + reference.range.end.character + ) // 2 position_inside = Position(line=middle_line, character=middle_char) filtered = list(filter(by_position(position_inside), all_references)) assert len(filtered) == 1 @@ -176,15 +180,16 @@ def test_filter_references_by_position() -> None: position_outside = Position(line=outside_line, character=outside_char) filtered_outside = list(filter(by_position(position_outside), all_references)) - assert reference not in filtered_outside, ( - f"Reference {i} should not match position outside its range" - ) + assert ( + reference not in filtered_outside + ), f"Reference {i} should not match position outside its range" # Test case: cursor at beginning of file - no references should match position_start = Position(line=0, character=0) filtered_start = list(filter(by_position(position_start), all_references)) assert len(filtered_start) == 0 or all( - ref.range.start.line == 0 and ref.range.start.character <= 0 for ref in filtered_start + ref.range.start.line == 0 and ref.range.start.character <= 0 + for ref in filtered_start ) # Test case: cursor at end of file - no references should match (unless there's a reference at the end) diff --git a/tests/lsp/test_reference_cte.py b/tests/lsp/test_reference_cte.py index 9bc74bc990..8ba50eefd4 100644 --- a/tests/lsp/test_reference_cte.py +++ b/tests/lsp/test_reference_cte.py @@ -1,10 +1,12 @@ import re +import typing as t + +from lsprotocol.types import Position, Range + from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext, ModelTarget from sqlmesh.lsp.reference import CTEReference, get_references from sqlmesh.lsp.uri import URI -from lsprotocol.types import Range, Position -import typing as t def test_cte_parsing(): @@ -24,8 +26,12 @@ def test_cte_parsing(): # Find position of the cte reference ranges = find_ranges_from_regex(read_file, r"current_marketing(?!_outer)") assert len(ranges) == 2 - position = Position(line=ranges[1].start.line, character=ranges[1].start.character + 4) - references = get_references(lsp_context, URI.from_path(sushi_customers_path), position) + position = Position( + line=ranges[1].start.line, character=ranges[1].start.character + 4 + ) + references = get_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) assert len(references) == 1 assert references[0].path == sushi_customers_path assert isinstance(references[0], CTEReference) @@ -39,8 +45,12 @@ def test_cte_parsing(): # Find the position of the current_marketing_outer reference ranges = find_ranges_from_regex(read_file, r"current_marketing_outer") assert len(ranges) == 2 - position = Position(line=ranges[1].start.line, character=ranges[1].start.character + 4) - references = get_references(lsp_context, URI.from_path(sushi_customers_path), position) + position = Position( + line=ranges[1].start.line, character=ranges[1].start.character + 4 + ) + references = get_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) assert len(references) == 1 assert references[0].path == sushi_customers_path assert isinstance(references[0], CTEReference) diff --git a/tests/lsp/test_reference_cte_find_all.py b/tests/lsp/test_reference_cte_find_all.py index dabe1589e2..42f7cc626e 100644 --- a/tests/lsp/test_reference_cte_find_all.py +++ b/tests/lsp/test_reference_cte_find_all.py @@ -1,4 +1,5 @@ from lsprotocol.types import Position + from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext, ModelTarget from sqlmesh.lsp.reference import get_cte_references @@ -24,8 +25,12 @@ def test_cte_find_all_references(): assert len(ranges) == 2 # regex finds 2 occurrences (definition and FROM clause) # Click on the CTE definition - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) - references = get_cte_references(lsp_context, URI.from_path(sushi_customers_path), position) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) + references = get_cte_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) # Should find the definition, FROM clause, and column prefix usages assert len(references) == 4 # definition + FROM + 2 column prefix uses assert all(ref.path == sushi_customers_path for ref in references) @@ -36,13 +41,15 @@ def test_cte_find_all_references(): ref_range.start.line == expected_range.start.line and ref_range.start.character == expected_range.start.character for ref_range in reference_ranges - ), ( - f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" - ) + ), f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" # Click on the CTE usage - position = Position(line=ranges[1].start.line, character=ranges[1].start.character + 4) - references = get_cte_references(lsp_context, URI.from_path(sushi_customers_path), position) + position = Position( + line=ranges[1].start.line, character=ranges[1].start.character + 4 + ) + references = get_cte_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) # Should find the same references assert len(references) == 4 # definition + FROM + 2 column prefix uses @@ -54,9 +61,7 @@ def test_cte_find_all_references(): ref_range.start.line == expected_range.start.line and ref_range.start.character == expected_range.start.character for ref_range in reference_ranges - ), ( - f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" - ) + ), f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" def test_cte_find_all_references_outer(): @@ -77,8 +82,12 @@ def test_cte_find_all_references_outer(): assert len(ranges) == 2 # Click on the CTE definition - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) - references = get_cte_references(lsp_context, URI.from_path(sushi_customers_path), position) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) + references = get_cte_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) # Should find both the definition and the usage assert len(references) == 2 @@ -91,13 +100,15 @@ def test_cte_find_all_references_outer(): ref_range.start.line == expected_range.start.line and ref_range.start.character == expected_range.start.character for ref_range in reference_ranges - ), ( - f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" - ) + ), f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" # Click on the CTE usage - position = Position(line=ranges[1].start.line, character=ranges[1].start.character + 4) - references = get_cte_references(lsp_context, URI.from_path(sushi_customers_path), position) + position = Position( + line=ranges[1].start.line, character=ranges[1].start.character + 4 + ) + references = get_cte_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) # Should find the same references assert len(references) == 2 @@ -109,9 +120,7 @@ def test_cte_find_all_references_outer(): ref_range.start.line == expected_range.start.line and ref_range.start.character == expected_range.start.character for ref_range in reference_ranges - ), ( - f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" - ) + ), f"Expected to find reference at line {expected_range.start.line}, char {expected_range.start.character}" def test_cte_no_references_on_non_cte(): @@ -132,8 +141,12 @@ def test_cte_no_references_on_non_cte(): ranges = find_ranges_from_regex(read_file, r"sushi\.orders") assert len(ranges) >= 1 - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) - references = get_cte_references(lsp_context, URI.from_path(sushi_customers_path), position) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) + references = get_cte_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) # Should find no references since this is not a CTE assert len(references) == 0 diff --git a/tests/lsp/test_reference_external_model.py b/tests/lsp/test_reference_external_model.py index 25de22f10f..a69e756a12 100644 --- a/tests/lsp/test_reference_external_model.py +++ b/tests/lsp/test_reference_external_model.py @@ -1,4 +1,5 @@ import os +import typing as t from pathlib import Path from sqlmesh import Config @@ -10,7 +11,6 @@ from sqlmesh.lsp.uri import URI from sqlmesh.utils.lineage import ExternalModelReference from tests.utils.test_filesystem import create_temp_file -import typing as t def test_reference() -> None: @@ -54,7 +54,9 @@ def test_unregistered_external_model(tmp_path: Path): lsp_context = LSPContext(ctx) uri = URI.from_path(model_path) - references = get_references(lsp_context, uri, Position(line=0, character=len(contents) - 3)) + references = get_references( + lsp_context, uri, Position(line=0, character=len(contents) - 3) + ) assert len(references) == 1 reference = references[0] diff --git a/tests/lsp/test_reference_macro.py b/tests/lsp/test_reference_macro.py index 3ee7c48b3b..4dffe17a4a 100644 --- a/tests/lsp/test_reference_macro.py +++ b/tests/lsp/test_reference_macro.py @@ -1,6 +1,7 @@ from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext, ModelTarget -from sqlmesh.lsp.reference import MacroReference, get_macro_definitions_for_a_path +from sqlmesh.lsp.reference import (MacroReference, + get_macro_definitions_for_a_path) from sqlmesh.lsp.uri import URI diff --git a/tests/lsp/test_reference_macro_find_all.py b/tests/lsp/test_reference_macro_find_all.py index 328924599a..c033ac9405 100644 --- a/tests/lsp/test_reference_macro_find_all.py +++ b/tests/lsp/test_reference_macro_find_all.py @@ -1,16 +1,13 @@ from lsprotocol.types import Position + from sqlmesh.core.context import Context +from sqlmesh.core.linter.helpers import Position as SQLMeshPosition +from sqlmesh.core.linter.helpers import Range as SQLMeshRange +from sqlmesh.core.linter.helpers import read_range_from_file from sqlmesh.lsp.context import LSPContext, ModelTarget -from sqlmesh.lsp.reference import ( - get_macro_find_all_references, - get_macro_definitions_for_a_path, -) +from sqlmesh.lsp.reference import (get_macro_definitions_for_a_path, + get_macro_find_all_references) from sqlmesh.lsp.uri import URI -from sqlmesh.core.linter.helpers import ( - read_range_from_file, - Range as SQLMeshRange, - Position as SQLMeshPosition, -) def test_find_all_references_for_macro_add_one(): @@ -29,16 +26,22 @@ def test_find_all_references_for_macro_add_one(): macro_references = get_macro_definitions_for_a_path(lsp_context, top_waiters_uri) # Find the @ADD_ONE reference - add_one_ref = next((ref for ref in macro_references if ref.range.start.line == 12), None) + add_one_ref = next( + (ref for ref in macro_references if ref.range.start.line == 12), None + ) assert add_one_ref is not None, "Should find @ADD_ONE reference in top_waiters" # Click on the @ADD_ONE macro at line 13, character 5 (the @ symbol) position = Position(line=12, character=5) - all_references = get_macro_find_all_references(lsp_context, top_waiters_uri, position) + all_references = get_macro_find_all_references( + lsp_context, top_waiters_uri, position + ) # Should find at least 2 references: the definition and the usage in top_waiters - assert len(all_references) >= 2, f"Expected at least 2 references, found {len(all_references)}" + assert ( + len(all_references) >= 2 + ), f"Expected at least 2 references, found {len(all_references)}" # Verify the macro definition is included definition_refs = [ref for ref in all_references if "utils.py" in str(ref.path)] @@ -56,7 +59,9 @@ def test_find_all_references_for_macro_add_one(): for expected_file, expectations in expected_files.items(): file_refs = [ref for ref in all_references if expected_file in str(ref.path)] - assert len(file_refs) >= 1, f"Should find at least one reference in {expected_file}" + assert ( + len(file_refs) >= 1 + ), f"Should find at least one reference in {expected_file}" file_ref = file_refs[0] file_path = file_ref.path @@ -72,9 +77,9 @@ def test_find_all_references_for_macro_add_one(): # Read the content at the reference location content = read_range_from_file(file_path, sqlmesh_range) - assert content.startswith(expectations["expected_content"]), ( - f"Expected content to start with '{expectations['expected_content']}', got: {content}" - ) + assert content.startswith( + expectations["expected_content"] + ), f"Expected content to start with '{expectations['expected_content']}', got: {content}" def test_find_all_references_for_macro_multiply(): @@ -93,21 +98,29 @@ def test_find_all_references_for_macro_multiply(): macro_references = get_macro_definitions_for_a_path(lsp_context, top_waiters_uri) # Find the @MULTIPLY reference - multiply_ref = next((ref for ref in macro_references if ref.range.start.line == 13), None) + multiply_ref = next( + (ref for ref in macro_references if ref.range.start.line == 13), None + ) assert multiply_ref is not None, "Should find @MULTIPLY reference in top_waiters" # Click on the @MULTIPLY macro at line 14, character 5 (the @ symbol) position = Position(line=13, character=5) - all_references = get_macro_find_all_references(lsp_context, top_waiters_uri, position) + all_references = get_macro_find_all_references( + lsp_context, top_waiters_uri, position + ) # Should find at least 2 references: the definition and the usage - assert len(all_references) >= 2, f"Expected at least 2 references, found {len(all_references)}" + assert ( + len(all_references) >= 2 + ), f"Expected at least 2 references, found {len(all_references)}" # Verify both definition and usage are included - assert any("utils.py" in str(ref.path) for ref in all_references), ( - "Should include macro definition" - ) - assert any("top_waiters" in str(ref.path) for ref in all_references), "Should include usage" + assert any( + "utils.py" in str(ref.path) for ref in all_references + ), "Should include macro definition" + assert any( + "top_waiters" in str(ref.path) for ref in all_references + ), "Should include usage" def test_find_all_references_for_sql_literal_macro(): @@ -126,15 +139,23 @@ def test_find_all_references_for_sql_literal_macro(): macro_references = get_macro_definitions_for_a_path(lsp_context, top_waiters_uri) # Find the @SQL_LITERAL reference - sql_literal_ref = next((ref for ref in macro_references if ref.range.start.line == 14), None) - assert sql_literal_ref is not None, "Should find @SQL_LITERAL reference in top_waiters" + sql_literal_ref = next( + (ref for ref in macro_references if ref.range.start.line == 14), None + ) + assert ( + sql_literal_ref is not None + ), "Should find @SQL_LITERAL reference in top_waiters" # Click on the @SQL_LITERAL macro position = Position(line=14, character=5) - all_references = get_macro_find_all_references(lsp_context, top_waiters_uri, position) + all_references = get_macro_find_all_references( + lsp_context, top_waiters_uri, position + ) # For user-defined macros in utils.py, should find references - assert len(all_references) >= 2, f"Expected at least 2 references, found {len(all_references)}" + assert ( + len(all_references) >= 2 + ), f"Expected at least 2 references, found {len(all_references)}" def test_find_references_from_outside_macro_position(): @@ -152,15 +173,21 @@ def test_find_references_from_outside_macro_position(): # Click on a position that is not on a macro position = Position(line=0, character=0) # First line, which is a comment - all_references = get_macro_find_all_references(lsp_context, top_waiters_uri, position) + all_references = get_macro_find_all_references( + lsp_context, top_waiters_uri, position + ) # Should return empty list when not on a macro - assert len(all_references) == 0, "Should not find macro references when not on a macro" + assert ( + len(all_references) == 0 + ), "Should not find macro references when not on a macro" def test_multi_repo_macro_references(): """Test finding macro references across multiple repositories.""" - context = Context(paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory") + context = Context( + paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory" + ) lsp_context = LSPContext(context) # Find model 'd' which uses macros from repo_2 @@ -177,19 +204,22 @@ def test_multi_repo_macro_references(): # Click on the second macro reference which appears under the same name in repo_1 ('dup') first_ref = macro_references[1] position = Position( - line=first_ref.range.start.line, character=first_ref.range.start.character + 1 + line=first_ref.range.start.line, + character=first_ref.range.start.character + 1, ) all_references = get_macro_find_all_references(lsp_context, d_uri, position) # Should find the definition and usage - assert len(all_references) == 2, f"Expected 2 references, found {len(all_references)}" + assert ( + len(all_references) == 2 + ), f"Expected 2 references, found {len(all_references)}" # Verify references from repo_2 - assert any("repo_2" in str(ref.path) for ref in all_references), ( - "Should find macro in repo_2" - ) + assert any( + "repo_2" in str(ref.path) for ref in all_references + ), "Should find macro in repo_2" # But not references in repo_1 since despite identical name they're different macros - assert not any("repo_1" in str(ref.path) for ref in all_references), ( - "Shouldn't find macro in repo_1" - ) + assert not any( + "repo_1" in str(ref.path) for ref in all_references + ), "Shouldn't find macro in repo_1" diff --git a/tests/lsp/test_reference_macro_multi.py b/tests/lsp/test_reference_macro_multi.py index 3902c0b275..7299ba7b3c 100644 --- a/tests/lsp/test_reference_macro_multi.py +++ b/tests/lsp/test_reference_macro_multi.py @@ -1,11 +1,14 @@ from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext, ModelTarget -from sqlmesh.lsp.reference import MacroReference, get_macro_definitions_for_a_path +from sqlmesh.lsp.reference import (MacroReference, + get_macro_definitions_for_a_path) from sqlmesh.lsp.uri import URI def test_macro_references_multirepo() -> None: - context = Context(paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory") + context = Context( + paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory" + ) lsp_context = LSPContext(context) d_path = next( @@ -20,5 +23,7 @@ def test_macro_references_multirepo() -> None: assert len(macro_references) == 2 for ref in macro_references: assert isinstance(ref, MacroReference) - assert str(URI.from_path(ref.path).value).endswith("multi/repo_2/macros/__init__.py") + assert str(URI.from_path(ref.path).value).endswith( + "multi/repo_2/macros/__init__.py" + ) assert ref.target_range is not None diff --git a/tests/lsp/test_reference_model_column_prefix.py b/tests/lsp/test_reference_model_column_prefix.py index 082ee9c8e6..02b66b6223 100644 --- a/tests/lsp/test_reference_model_column_prefix.py +++ b/tests/lsp/test_reference_model_column_prefix.py @@ -36,17 +36,20 @@ def test_model_reference_with_column_prefix(): assert from_clause_range is not None, "Should find FROM clause with sushi.orders" position = Position( - line=from_clause_range.start.line, character=from_clause_range.start.character + 6 + line=from_clause_range.start.line, + character=from_clause_range.start.character + 6, ) - model_refs = get_all_references(lsp_context, URI.from_path(sushi_customers_path), position) + model_refs = get_all_references( + lsp_context, URI.from_path(sushi_customers_path), position + ) assert len(model_refs) >= 6 # Verify that we have the FROM clause reference - assert any(ref.range.start.line == from_clause_range.start.line for ref in model_refs), ( - "Should find FROM clause reference" - ) + assert any( + ref.range.start.line == from_clause_range.start.line for ref in model_refs + ), "Should find FROM clause reference" def test_column_prefix_references_are_found(): @@ -66,21 +69,25 @@ def test_column_prefix_references_are_found(): ranges = find_ranges_from_regex(read_file, r"sushi\.orders") # Should find exactly 1 in FROM clause with column prefix - assert len(ranges) == 1, f"Expected 1 occurrence of 'sushi.orders', found {len(ranges)}" + assert ( + len(ranges) == 1 + ), f"Expected 1 occurrence of 'sushi.orders', found {len(ranges)}" # Verify we have the expected lines line_contents = [read_file[r.start.line].strip() for r in ranges] # Should find FROM clause - assert any("FROM sushi.orders" in content for content in line_contents), ( - "Should find FROM clause with sushi.orders" - ) + assert any( + "FROM sushi.orders" in content for content in line_contents + ), "Should find FROM clause with sushi.orders" def test_quoted_uppercase_table_and_column_references(tmp_path: Path): # Initialize example project in temporary directory with case sensitive normalization init_example_project( - tmp_path, engine_type="duckdb", dialect="duckdb,normalization_strategy=case_sensitive" + tmp_path, + engine_type="duckdb", + dialect="duckdb,normalization_strategy=case_sensitive", ) # Create a model with quoted uppercase schema and table names @@ -142,7 +149,9 @@ def test_quoted_uppercase_table_and_column_references(tmp_path: Path): ranges = find_ranges_from_regex(read_file, r'"SUSHI"\.orders') # Should find 3 occurrences: FROM clause and 2 in WHERE clause with column prefix - assert len(ranges) == 3, f"Expected 3 occurrences of '\"SUSHI\".orders', found {len(ranges)}" + assert ( + len(ranges) == 3 + ), f"Expected 3 occurrences of '\"SUSHI\".orders', found {len(ranges)}" # Click on the table reference in FROM clause from_clause_range = None @@ -155,16 +164,19 @@ def test_quoted_uppercase_table_and_column_references(tmp_path: Path): assert from_clause_range is not None, 'Should find FROM clause with "SUSHI".orders' position = Position( - line=from_clause_range.start.line, character=from_clause_range.start.character + 5 + line=from_clause_range.start.line, + character=from_clause_range.start.character + 5, ) - model_refs = get_all_references(lsp_context, URI.from_path(quoted_test_model_path), position) + model_refs = get_all_references( + lsp_context, URI.from_path(quoted_test_model_path), position + ) # Should find only references to "SUSHI".orders (3 total: FROM clause and 2 column prefixes in WHERE) # The lowercase sushi.orders should NOT be included if case sensitivity is working - assert len(model_refs) == 4, ( - f'Expected exactly 3 references for "SUSHI".orders, found {len(model_refs)}' - ) + assert ( + len(model_refs) == 4 + ), f'Expected exactly 3 references for "SUSHI".orders, found {len(model_refs)}' # Verify that we have all 3 references ref_lines = [ref.range.start.line for ref in model_refs] @@ -175,15 +187,17 @@ def test_quoted_uppercase_table_and_column_references(tmp_path: Path): assert from_line in ref_lines, "Should find FROM clause reference" for where_line in where_lines: - assert where_line in ref_lines, f"Should find WHERE clause reference on line {where_line}" + assert ( + where_line in ref_lines + ), f"Should find WHERE clause reference on line {where_line}" # Now test that lowercase sushi.orders references are separate lowercase_ranges = find_ranges_from_regex(read_file, r"sushi\.orders") # Should find 2 occurrences: FROM clause and 1 in WHERE clause - assert len(lowercase_ranges) == 2, ( - f"Expected 2 occurrences of 'sushi.orders', found {len(lowercase_ranges)}" - ) + assert ( + len(lowercase_ranges) == 2 + ), f"Expected 2 occurrences of 'sushi.orders', found {len(lowercase_ranges)}" # Click on the lowercase table reference lowercase_from_range = None @@ -196,7 +210,8 @@ def test_quoted_uppercase_table_and_column_references(tmp_path: Path): assert lowercase_from_range is not None, "Should find FROM clause with sushi.orders" lowercase_position = Position( - line=lowercase_from_range.start.line, character=lowercase_from_range.start.character + 5 + line=lowercase_from_range.start.line, + character=lowercase_from_range.start.character + 5, ) lowercase_refs = get_all_references( @@ -204,6 +219,6 @@ def test_quoted_uppercase_table_and_column_references(tmp_path: Path): ) # Should find only references to lowercase sushi.orders, NOT the uppercase ones - assert len(lowercase_refs) == 3, ( - f"Expected exactly 2 references for sushi.orders, found {len(lowercase_refs)}" - ) + assert ( + len(lowercase_refs) == 3 + ), f"Expected exactly 2 references for sushi.orders, found {len(lowercase_refs)}" diff --git a/tests/lsp/test_reference_model_find_all.py b/tests/lsp/test_reference_model_find_all.py index cd9c0a3a1c..2e5be01f39 100644 --- a/tests/lsp/test_reference_model_find_all.py +++ b/tests/lsp/test_reference_model_find_all.py @@ -1,10 +1,9 @@ from lsprotocol.types import Position + from sqlmesh.core.context import Context -from sqlmesh.lsp.context import LSPContext, ModelTarget, AuditTarget -from sqlmesh.lsp.reference import ( - get_model_find_all_references, - get_model_definitions_for_a_path, -) +from sqlmesh.lsp.context import AuditTarget, LSPContext, ModelTarget +from sqlmesh.lsp.reference import (get_model_definitions_for_a_path, + get_model_find_all_references) from sqlmesh.lsp.uri import URI from tests.lsp.test_reference_cte import find_ranges_from_regex @@ -28,11 +27,15 @@ def test_find_references_for_model_usages(): assert len(ranges) >= 1, "Should find at least one reference to sushi.orders" # Click on the model reference - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 6) - references = get_model_find_all_references(lsp_context, URI.from_path(customers_path), position) - assert len(references) >= 6, ( - f"Expected at least 6 references to sushi.orders (including column prefix), found {len(references)}" + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 6 ) + references = get_model_find_all_references( + lsp_context, URI.from_path(customers_path), position + ) + assert ( + len(references) >= 6 + ), f"Expected at least 6 references to sushi.orders (including column prefix), found {len(references)}" # Verify expected files are present reference_files = {str(ref.path) for ref in references} @@ -45,9 +48,9 @@ def test_find_references_for_model_usages(): "waiter_revenue_by_day", ] for pattern in expected_patterns: - assert any(pattern in uri for uri in reference_files), ( - f"Missing reference in file containing '{pattern}'" - ) + assert any( + pattern in uri for uri in reference_files + ), f"Missing reference in file containing '{pattern}'" # Verify exact ranges for each reference pattern # Note: customers file has multiple references due to column prefix support @@ -79,9 +82,9 @@ def test_find_references_for_model_usages(): assert pattern in refs_by_pattern, f"Missing references for pattern '{pattern}'" actual_refs = refs_by_pattern[pattern] - assert len(actual_refs) == len(expected_range_list), ( - f"Expected {len(expected_range_list)} references for {pattern}, found {len(actual_refs)}" - ) + assert len(actual_refs) == len( + expected_range_list + ), f"Expected {len(expected_range_list)} references for {pattern}, found {len(actual_refs)}" # Sort both actual and expected by line number for consistent comparison actual_refs_sorted = sorted( @@ -89,23 +92,28 @@ def test_find_references_for_model_usages(): ) expected_sorted = sorted(expected_range_list, key=lambda r: (r[0], r[1])) - for i, (ref, expected_range) in enumerate(zip(actual_refs_sorted, expected_sorted)): - expected_start_line, expected_start_char, expected_end_line, expected_end_char = ( - expected_range - ) - - assert ref.range.start.line == expected_start_line, ( - f"Expected {pattern} reference #{i + 1} start line {expected_start_line}, found {ref.range.start.line}" - ) - assert ref.range.start.character == expected_start_char, ( - f"Expected {pattern} reference #{i + 1} start character {expected_start_char}, found {ref.range.start.character}" - ) - assert ref.range.end.line == expected_end_line, ( - f"Expected {pattern} reference #{i + 1} end line {expected_end_line}, found {ref.range.end.line}" - ) - assert ref.range.end.character == expected_end_char, ( - f"Expected {pattern} reference #{i + 1} end character {expected_end_char}, found {ref.range.end.character}" - ) + for i, (ref, expected_range) in enumerate( + zip(actual_refs_sorted, expected_sorted) + ): + ( + expected_start_line, + expected_start_char, + expected_end_line, + expected_end_char, + ) = expected_range + + assert ( + ref.range.start.line == expected_start_line + ), f"Expected {pattern} reference #{i + 1} start line {expected_start_line}, found {ref.range.start.line}" + assert ( + ref.range.start.character == expected_start_char + ), f"Expected {pattern} reference #{i + 1} start character {expected_start_char}, found {ref.range.start.character}" + assert ( + ref.range.end.line == expected_end_line + ), f"Expected {pattern} reference #{i + 1} end line {expected_end_line}, found {ref.range.end.line}" + assert ( + ref.range.end.character == expected_end_char + ), f"Expected {pattern} reference #{i + 1} end character {expected_end_char}, found {ref.range.end.character}" def test_find_references_for_marketing_model(): @@ -123,25 +131,30 @@ def test_find_references_for_marketing_model(): # Find sushi.marketing reference marketing_ranges = find_ranges_from_regex(read_file, r"sushi\.marketing") - assert len(marketing_ranges) >= 1, "Should find at least one reference to sushi.marketing" + assert ( + len(marketing_ranges) >= 1 + ), "Should find at least one reference to sushi.marketing" position = Position( - line=marketing_ranges[0].start.line, character=marketing_ranges[0].start.character + 8 + line=marketing_ranges[0].start.line, + character=marketing_ranges[0].start.character + 8, + ) + references = get_model_find_all_references( + lsp_context, URI.from_path(customers_path), position ) - references = get_model_find_all_references(lsp_context, URI.from_path(customers_path), position) # sushi.marketing should have exactly 2 references: model itself + customers usage - assert len(references) == 2, ( - f"Expected exactly 2 references to sushi.marketing, found {len(references)}" - ) + assert ( + len(references) == 2 + ), f"Expected exactly 2 references to sushi.marketing, found {len(references)}" # Verify files are present reference_files = {str(ref.path) for ref in references} expected_patterns = ["marketing", "customers"] for pattern in expected_patterns: - assert any(pattern in uri for uri in reference_files), ( - f"Missing reference in file containing '{pattern}'" - ) + assert any( + pattern in uri for uri in reference_files + ), f"Missing reference in file containing '{pattern}'" def test_find_references_for_python_model(): @@ -152,7 +165,8 @@ def test_find_references_for_python_model(): revenue_path = next( path for path, info in lsp_context.map.items() - if isinstance(info, ModelTarget) and "sushi.customer_revenue_by_day" in info.names + if isinstance(info, ModelTarget) + and "sushi.customer_revenue_by_day" in info.names ) with open(revenue_path, "r", encoding="utf-8") as file: @@ -165,7 +179,9 @@ def test_find_references_for_python_model(): position = Position( line=items_ranges[0].start.line, character=items_ranges[0].start.character + 6 ) - references = get_model_find_all_references(lsp_context, URI.from_path(revenue_path), position) + references = get_model_find_all_references( + lsp_context, URI.from_path(revenue_path), position + ) assert len(references) == 5 # Verify expected files @@ -180,9 +196,9 @@ def test_find_references_for_python_model(): "assert_item_price_above_zero", ] for pattern in expected_patterns: - assert any(pattern in uri for uri in reference_files), ( - f"Missing reference in file containing '{pattern}'" - ) + assert any( + pattern in uri for uri in reference_files + ), f"Missing reference in file containing '{pattern}'" def test_waiter_revenue_by_day_multiple_references(): @@ -203,9 +219,9 @@ def test_waiter_revenue_by_day_multiple_references(): waiter_revenue_ranges = find_ranges_from_regex( top_waiters_file, r"sushi\.waiter_revenue_by_day" ) - assert len(waiter_revenue_ranges) >= 2, ( - "Should find at least 2 references to sushi.waiter_revenue_by_day in top_waiters" - ) + assert ( + len(waiter_revenue_ranges) >= 2 + ), "Should find at least 2 references to sushi.waiter_revenue_by_day in top_waiters" # Click on the first reference position = Position( @@ -217,20 +233,20 @@ def test_waiter_revenue_by_day_multiple_references(): ) # Should find model definition + 3 references in top_waiters = 4 total - assert len(references) == 4, ( - f"Expected exactly 4 references to sushi.waiter_revenue_by_day, found {len(references)}" - ) + assert ( + len(references) == 4 + ), f"Expected exactly 4 references to sushi.waiter_revenue_by_day, found {len(references)}" # Count references in top_waiters file top_waiters_refs = [ref for ref in references if "top_waiters" in str(ref.path)] - assert len(top_waiters_refs) == 3, ( - f"Expected exactly 3 references in top_waiters, found {len(top_waiters_refs)}" - ) + assert ( + len(top_waiters_refs) == 3 + ), f"Expected exactly 3 references in top_waiters, found {len(top_waiters_refs)}" # Verify model definition is included - assert any("waiter_revenue_by_day" in str(ref.path) for ref in references), ( - "Should include model definition" - ) + assert any( + "waiter_revenue_by_day" in str(ref.path) for ref in references + ), "Should include model definition" def test_precise_character_positions(): @@ -247,28 +263,44 @@ def test_precise_character_positions(): # Click on 's' in "sushi" - should work position = Position(line=30, character=7) - references = get_model_find_all_references(lsp_context, URI.from_path(customers_path), position) + references = get_model_find_all_references( + lsp_context, URI.from_path(customers_path), position + ) assert len(references) > 0, "Should find references when clicking on 's' in 'sushi'" # Click on '.' between sushi and orders - should work position = Position(line=30, character=12) - references = get_model_find_all_references(lsp_context, URI.from_path(customers_path), position) + references = get_model_find_all_references( + lsp_context, URI.from_path(customers_path), position + ) assert len(references) > 0, "Should find references when clicking on '.' separator" # Click on 'o' in "orders" - should work position = Position(line=30, character=13) - references = get_model_find_all_references(lsp_context, URI.from_path(customers_path), position) - assert len(references) > 0, "Should find references when clicking on 'o' in 'orders'" + references = get_model_find_all_references( + lsp_context, URI.from_path(customers_path), position + ) + assert ( + len(references) > 0 + ), "Should find references when clicking on 'o' in 'orders'" # Click just before "sushi" - should not work position = Position(line=30, character=6) - references = get_model_find_all_references(lsp_context, URI.from_path(customers_path), position) - assert len(references) == 0, "Should not find references when clicking just before 'sushi'" + references = get_model_find_all_references( + lsp_context, URI.from_path(customers_path), position + ) + assert ( + len(references) == 0 + ), "Should not find references when clicking just before 'sushi'" # Click just after "orders" - should not work position = Position(line=30, character=21) - references = get_model_find_all_references(lsp_context, URI.from_path(customers_path), position) - assert len(references) == 0, "Should not find references when clicking just after 'orders'" + references = get_model_find_all_references( + lsp_context, URI.from_path(customers_path), position + ) + assert ( + len(references) == 0 + ), "Should not find references when clicking just after 'orders'" def test_audit_model_references(): @@ -277,7 +309,9 @@ def test_audit_model_references(): lsp_context = LSPContext(context) # Find audit files - audit_paths = [path for path, info in lsp_context.map.items() if isinstance(info, AuditTarget)] + audit_paths = [ + path for path, info in lsp_context.map.items() if isinstance(info, AuditTarget) + ] if audit_paths: audit_path = audit_paths[0] @@ -288,13 +322,16 @@ def test_audit_model_references(): # Click on the first reference which is: sushi.items first_ref = refs[0] position = Position( - line=first_ref.range.start.line, character=first_ref.range.start.character + 1 + line=first_ref.range.start.line, + character=first_ref.range.start.character + 1, ) references = get_model_find_all_references( lsp_context, URI.from_path(audit_path), position ) - assert len(references) == 5, "Should find references from audit files as well" + assert ( + len(references) == 5 + ), "Should find references from audit files as well" reference_files = {str(ref.path) for ref in references} @@ -307,6 +344,6 @@ def test_audit_model_references(): "assert_item_price_above_zero", ] for pattern in expected_patterns: - assert any(pattern in uri for uri in reference_files), ( - f"Missing reference in file containing '{pattern}'" - ) + assert any( + pattern in uri for uri in reference_files + ), f"Missing reference in file containing '{pattern}'" diff --git a/tests/lsp/test_rename_cte.py b/tests/lsp/test_rename_cte.py index 4ca1002c2e..f99b834b7b 100644 --- a/tests/lsp/test_rename_cte.py +++ b/tests/lsp/test_rename_cte.py @@ -1,4 +1,5 @@ from lsprotocol.types import Position + from sqlmesh.core.context import Context from sqlmesh.lsp.context import LSPContext, ModelTarget from sqlmesh.lsp.rename import prepare_rename, rename_symbol @@ -24,7 +25,9 @@ def test_prepare_rename_cte(): assert len(ranges) == 2 # Click on the CTE definition - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) result = prepare_rename(lsp_context, URI.from_path(sushi_customers_path), position) assert result is not None @@ -32,7 +35,9 @@ def test_prepare_rename_cte(): assert result.range == ranges[0] # Should return the definition range # Test clicking on CTE usage - position = Position(line=ranges[1].start.line, character=ranges[1].start.character + 4) + position = Position( + line=ranges[1].start.line, character=ranges[1].start.character + 4 + ) result = prepare_rename(lsp_context, URI.from_path(sushi_customers_path), position) assert result is not None @@ -58,7 +63,9 @@ def test_prepare_rename_cte_outer(): assert len(ranges) == 2 # Click on the CTE definition - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) result = prepare_rename(lsp_context, URI.from_path(sushi_customers_path), position) assert result is not None @@ -83,7 +90,9 @@ def test_prepare_rename_non_cte(): ranges = find_ranges_from_regex(read_file, r"sushi\.orders") assert len(ranges) >= 1 - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) result = prepare_rename(lsp_context, URI.from_path(sushi_customers_path), position) assert result is None @@ -107,7 +116,9 @@ def test_rename_cte(): assert len(ranges) == 2 # Click on the CTE definition - position = Position(line=ranges[0].start.line, character=ranges[0].start.character + 4) + position = Position( + line=ranges[0].start.line, character=ranges[0].start.character + 4 + ) workspace_edit = rename_symbol( lsp_context, URI.from_path(sushi_customers_path), position, "new_marketing" ) @@ -130,9 +141,7 @@ def test_rename_cte(): edit_range.start.line == expected_range.start.line and edit_range.start.character == expected_range.start.character for edit_range in edit_ranges - ), ( - f"Expected to find edit at line {expected_range.start.line}, char {expected_range.start.character}" - ) + ), f"Expected to find edit at line {expected_range.start.line}, char {expected_range.start.character}" # Verify that all edits have the new name assert all(edit.new_text == "new_marketing" for edit in edits) @@ -157,7 +166,9 @@ def test_rename_cte(): # Verify the edited content edited_content = "".join(lines) assert "new_marketing" in edited_content - assert "current_marketing" not in edited_content.replace("current_marketing_outer", "") + assert "current_marketing" not in edited_content.replace( + "current_marketing_outer", "" + ) assert edited_content.count("new_marketing") == 4 assert ( " SELECT new_marketing.* FROM new_marketing WHERE new_marketing.customer_id != 100\n" @@ -183,9 +194,14 @@ def test_rename_cte_outer(): assert len(ranges) == 2 # Click on the CTE usage - position = Position(line=ranges[1].start.line, character=ranges[1].start.character + 4) + position = Position( + line=ranges[1].start.line, character=ranges[1].start.character + 4 + ) workspace_edit = rename_symbol( - lsp_context, URI.from_path(sushi_customers_path), position, "new_marketing_outer" + lsp_context, + URI.from_path(sushi_customers_path), + position, + "new_marketing_outer", ) assert workspace_edit is not None @@ -204,9 +220,7 @@ def test_rename_cte_outer(): edit_range.start.line == expected_range.start.line and edit_range.start.character == expected_range.start.character for edit_range in edit_ranges - ), ( - f"Expected to find edit at line {expected_range.start.line}, char {expected_range.start.character}" - ) + ), f"Expected to find edit at line {expected_range.start.line}, char {expected_range.start.character}" # Verify that all edits have the new name assert all(edit.new_text == "new_marketing_outer" for edit in edits) diff --git a/tests/setup.py b/tests/setup.py index ab48a3128f..4867cbe981 100644 --- a/tests/setup.py +++ b/tests/setup.py @@ -1,5 +1,6 @@ -import setuptools from pathlib import Path + +import setuptools import toml # type: ignore # This relies on `make package-tests` copying the sqlmesh pyproject.toml into tests/ so we can reference it diff --git a/tests/test_forking.py b/tests/test_forking.py index d11379a158..0ddb7bcf66 100644 --- a/tests/test_forking.py +++ b/tests/test_forking.py @@ -1,10 +1,10 @@ +import concurrent.futures import os + import pytest from sqlmesh import Context from sqlmesh.core.model import schema -import concurrent.futures - pytestmark = pytest.mark.isolated @@ -13,7 +13,9 @@ def test_parallel_load(assert_exp_eq, mocker): mocker.patch("sqlmesh.core.constants.MAX_FORK_WORKERS", 2) spy_update_schemas = mocker.spy(schema, "_update_model_schemas") - process_pool_executor = mocker.spy(concurrent.futures.ProcessPoolExecutor, "__init__") + process_pool_executor = mocker.spy( + concurrent.futures.ProcessPoolExecutor, "__init__" + ) as_completed = mocker.spy(concurrent.futures, "as_completed") context = Context(paths="examples/sushi") @@ -72,8 +74,12 @@ def test_parallel_load(assert_exp_eq, mocker): def test_parallel_load_multi_repo(assert_exp_eq, mocker): mocker.patch("sqlmesh.core.constants.MAX_FORK_WORKERS", 2) - process_pool_executor = mocker.spy(concurrent.futures.ProcessPoolExecutor, "__init__") - context = Context(paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory") + process_pool_executor = mocker.spy( + concurrent.futures.ProcessPoolExecutor, "__init__" + ) + context = Context( + paths=["examples/multi/repo_1", "examples/multi/repo_2"], gateway="memory" + ) if hasattr(os, "fork"): executor_args = process_pool_executor.call_args diff --git a/tests/utils/pandas.py b/tests/utils/pandas.py index b9451f4545..9ce2180553 100644 --- a/tests/utils/pandas.py +++ b/tests/utils/pandas.py @@ -19,5 +19,7 @@ def compare_dataframes( actual: pd.DataFrame, expected: pd.DataFrame, msg: str = "DataFrame", **kwargs ) -> None: actual = actual.sort_values(by=actual.columns.to_list()).reset_index(drop=True) - expected = expected.sort_values(by=expected.columns.to_list()).reset_index(drop=True) + expected = expected.sort_values(by=expected.columns.to_list()).reset_index( + drop=True + ) pd.testing.assert_frame_equal(actual, expected, obj=msg, **kwargs) diff --git a/tests/utils/test_aws.py b/tests/utils/test_aws.py index 905cc00dfe..f97d838543 100644 --- a/tests/utils/test_aws.py +++ b/tests/utils/test_aws.py @@ -1,6 +1,7 @@ import pytest -from sqlmesh.utils.errors import SQLMeshError, ConfigError -from sqlmesh.utils.aws import validate_s3_uri, parse_s3_uri + +from sqlmesh.utils.aws import parse_s3_uri, validate_s3_uri +from sqlmesh.utils.errors import ConfigError, SQLMeshError def test_validate_s3_uri(): diff --git a/tests/utils/test_cache.py b/tests/utils/test_cache.py index e6e041e30a..60e322a820 100644 --- a/tests/utils/test_cache.py +++ b/tests/utils/test_cache.py @@ -62,7 +62,9 @@ def test_optimized_query_cache(tmp_path: Path, mocker: MockerFixture): assert model._query_renderer._optimized_cache is not None -def test_optimized_query_cache_missing_rendered_query(tmp_path: Path, mocker: MockerFixture): +def test_optimized_query_cache_missing_rendered_query( + tmp_path: Path, mocker: MockerFixture +): model = SqlModel( name="test_model", query=parse_one("SELECT a FROM tbl"), @@ -85,15 +87,13 @@ def test_optimized_query_cache_missing_rendered_query(tmp_path: Path, mocker: Mo def test_optimized_query_cache_macro_def_change(tmp_path: Path, mocker: MockerFixture): - expressions = d.parse( - """ + expressions = d.parse(""" MODEL (name db.table); @DEF(filter_, a = 1); SELECT a FROM (SELECT 1 AS a) WHERE @filter_; - """ - ) + """) model = t.cast(SqlModel, load_sql_based_model(expressions)) cache = OptimizedQueryCache(tmp_path) @@ -110,15 +110,13 @@ def test_optimized_query_cache_macro_def_change(tmp_path: Path, mocker: MockerFi ) # Change the filter_ definition - new_expressions = d.parse( - """ + new_expressions = d.parse(""" MODEL (name db.table); @DEF(filter_, a = 2); SELECT a FROM (SELECT 1 AS a) WHERE @filter_; - """ - ) + """) new_model = t.cast(SqlModel, load_sql_based_model(new_expressions)) assert not cache.with_optimized_query(new_model) @@ -133,7 +131,9 @@ def test_optimized_query_cache_macro_def_change(tmp_path: Path, mocker: MockerFi ) -def test_file_cache_init_handles_stale_file(tmp_path: Path, mocker: MockerFixture) -> None: +def test_file_cache_init_handles_stale_file( + tmp_path: Path, mocker: MockerFixture +) -> None: cache: FileCache[_TestEntry] = FileCache(tmp_path) stale_file = tmp_path / f"{cache._cache_version}__fake_deleted_model_9999999999" diff --git a/tests/utils/test_concurrency.py b/tests/utils/test_concurrency.py index 5e1e4326f7..69321c17c1 100644 --- a/tests/utils/test_concurrency.py +++ b/tests/utils/test_concurrency.py @@ -2,11 +2,9 @@ from pytest_mock.plugin import MockerFixture from sqlmesh.core.snapshot import SnapshotId -from sqlmesh.utils.concurrency import ( - NodeExecutionFailedError, - concurrent_apply_to_snapshots, - concurrent_apply_to_values, -) +from sqlmesh.utils.concurrency import (NodeExecutionFailedError, + concurrent_apply_to_snapshots, + concurrent_apply_to_values) @pytest.mark.parametrize("tasks_num", [1, 2]) @@ -67,7 +65,9 @@ def raise_(): @pytest.mark.parametrize("tasks_num", [1, 2]) -def test_concurrent_apply_to_snapshots_return_failed_skipped(mocker: MockerFixture, tasks_num: int): +def test_concurrent_apply_to_snapshots_return_failed_skipped( + mocker: MockerFixture, tasks_num: int +): snapshot_a = mocker.Mock() snapshot_a.snapshot_id = SnapshotId(name="model_a", identifier="snapshot_a") snapshot_a.parents = [] @@ -139,7 +139,11 @@ def raise_(snapshot): assert len(errors) == 1 assert errors[0].node == failed_snapshot.snapshot_id - assert set(skipped) == {snapshot_a.snapshot_id, snapshot_b.snapshot_id, snapshot_c.snapshot_id} + assert set(skipped) == { + snapshot_a.snapshot_id, + snapshot_b.snapshot_id, + snapshot_c.snapshot_id, + } @pytest.mark.parametrize("tasks_num", [1, 3]) diff --git a/tests/utils/test_connection_pool.py b/tests/utils/test_connection_pool.py index c5926a3824..e5e8387a21 100644 --- a/tests/utils/test_connection_pool.py +++ b/tests/utils/test_connection_pool.py @@ -3,11 +3,9 @@ from pytest_mock.plugin import MockerFixture -from sqlmesh.utils.connection_pool import ( - SingletonConnectionPool, - ThreadLocalConnectionPool, - ThreadLocalSharedConnectionPool, -) +from sqlmesh.utils.connection_pool import (SingletonConnectionPool, + ThreadLocalConnectionPool, + ThreadLocalSharedConnectionPool) def test_singleton_connection_pool_get(mocker: MockerFixture): diff --git a/tests/utils/test_conversions.py b/tests/utils/test_conversions.py index 1e1b62f77e..4cb9747474 100644 --- a/tests/utils/test_conversions.py +++ b/tests/utils/test_conversions.py @@ -4,7 +4,8 @@ import pytest -from sqlmesh.utils.conversions import ensure_bool, make_serializable, try_str_to_bool +from sqlmesh.utils.conversions import (ensure_bool, make_serializable, + try_str_to_bool) class TestTryStrToBool: @@ -93,7 +94,9 @@ def test_boolean_strings(self, input_val: str, expected: bool) -> None: ("0", True), # String "0" is truthy (non-empty) ], ) - def test_other_strings_use_bool_conversion(self, input_val: str, expected: bool) -> None: + def test_other_strings_use_bool_conversion( + self, input_val: str, expected: bool + ) -> None: """Non-boolean strings fall back to bool() conversion.""" assert ensure_bool(input_val) is expected diff --git a/tests/utils/test_date.py b/tests/utils/test_date.py index c926507ccd..bfe7f294ab 100644 --- a/tests/utils/test_date.py +++ b/tests/utils/test_date.py @@ -1,27 +1,16 @@ import typing as t from datetime import date, datetime +import pandas as pd # noqa: TID253 import pytest import time_machine from sqlglot import exp -import pandas as pd # noqa: TID253 -from sqlmesh.utils.date import ( - UTC, - TimeLike, - date_dict, - format_tz_datetime, - is_categorical_relative_expression, - is_relative, - make_inclusive, - make_ts_exclusive, - to_datetime, - to_time_column, - to_timestamp, - to_ts, - to_tstz, - to_utc_timestamp, -) +from sqlmesh.utils.date import (UTC, TimeLike, date_dict, format_tz_datetime, + is_categorical_relative_expression, + is_relative, make_inclusive, make_ts_exclusive, + to_datetime, to_time_column, to_timestamp, + to_ts, to_tstz, to_utc_timestamp) def test_to_datetime() -> None: @@ -85,7 +74,12 @@ def test_to_timestamp() -> None: "start_in, end_in, start_out, end_out", [ ("2020-01-01", "2020-01-01", "2020-01-01", "2020-01-01 23:59:59.999999+00:00"), - ("2020-01-01", date(2020, 1, 1), "2020-01-01", "2020-01-01 23:59:59.999999+00:00"), + ( + "2020-01-01", + date(2020, 1, 1), + "2020-01-01", + "2020-01-01 23:59:59.999999+00:00", + ), ( date(2020, 1, 1), date(2020, 1, 1), @@ -116,7 +110,13 @@ def test_make_inclusive(start_in, end_in, start_out, end_out) -> None: @pytest.mark.parametrize( "start_in, end_in, start_out, end_out, dialect", [ - ("2020-01-01", "2020-01-01", "2020-01-01", "2020-01-01 23:59:59.999999999+00:00", "tsql"), + ( + "2020-01-01", + "2020-01-01", + "2020-01-01", + "2020-01-01 23:59:59.999999999+00:00", + "tsql", + ), ( "2020-01-01", date(2020, 1, 1), @@ -197,8 +197,13 @@ def test_to_ts(): def test_to_tstz(): - assert to_tstz(datetime(2020, 1, 1).replace(tzinfo=UTC)) == "2020-01-01 00:00:00+00:00" - assert to_tstz(datetime(2020, 1, 1).replace(tzinfo=None)) == "2020-01-01 00:00:00+00:00" + assert ( + to_tstz(datetime(2020, 1, 1).replace(tzinfo=UTC)) == "2020-01-01 00:00:00+00:00" + ) + assert ( + to_tstz(datetime(2020, 1, 1).replace(tzinfo=None)) + == "2020-01-01 00:00:00+00:00" + ) @pytest.mark.parametrize( @@ -270,13 +275,17 @@ def test_to_time_column( result: str, ): assert ( - to_time_column(time_column, time_column_type, dialect, time_column_format).sql(dialect) + to_time_column(time_column, time_column_type, dialect, time_column_format).sql( + dialect + ) == result ) def test_date_dict(): - resp = date_dict("2020-01-02 01:00:00", "2020-01-01 00:00:00", "2020-01-02 00:00:00") + resp = date_dict( + "2020-01-02 01:00:00", "2020-01-01 00:00:00", "2020-01-02 00:00:00" + ) assert resp == { "latest_dt": datetime(2020, 1, 2, 1, 0, 0, tzinfo=UTC), "execution_dt": datetime(2020, 1, 2, 1, 0, 0, tzinfo=UTC), @@ -352,7 +361,10 @@ def test_tsql_date_dict(start, end, expected_start_dt, expected_end_dt): def test_format_tz_datetime(): test_datetime = to_datetime("2020-01-01 00:00:00") assert format_tz_datetime(test_datetime) == "2020-01-01 12:00AM UTC" - assert format_tz_datetime(test_datetime, format_string=None) == "2020-01-01 00:00:00+00:00" + assert ( + format_tz_datetime(test_datetime, format_string=None) + == "2020-01-01 00:00:00+00:00" + ) def test_is_relative(): diff --git a/tests/utils/test_filesystem.py b/tests/utils/test_filesystem.py index 9b3d6f6895..8cc90fe862 100644 --- a/tests/utils/test_filesystem.py +++ b/tests/utils/test_filesystem.py @@ -1,7 +1,9 @@ import pathlib -def create_temp_file(tmp_path: pathlib.Path, filepath: pathlib.Path, contents: str) -> pathlib.Path: +def create_temp_file( + tmp_path: pathlib.Path, filepath: pathlib.Path, contents: str +) -> pathlib.Path: target_filepath = tmp_path / filepath target_filepath.parent.mkdir(parents=True, exist_ok=True) target_filepath.write_text(contents) diff --git a/tests/utils/test_git_client.py b/tests/utils/test_git_client.py index 13eecf294b..d73318a5d3 100644 --- a/tests/utils/test_git_client.py +++ b/tests/utils/test_git_client.py @@ -1,6 +1,8 @@ import subprocess from pathlib import Path + import pytest + from sqlmesh.utils.git import GitClient @@ -8,7 +10,9 @@ def git_repo(tmp_path: Path) -> Path: repo_path = tmp_path / "test_repo" repo_path.mkdir() - subprocess.run(["git", "init", "-b", "main"], cwd=repo_path, check=True, capture_output=True) + subprocess.run( + ["git", "init", "-b", "main"], cwd=repo_path, check=True, capture_output=True + ) return repo_path @@ -17,7 +21,9 @@ def test_git_uncommitted_changes(git_repo: Path): test_file = git_repo / "model.sql" test_file.write_text("SELECT 1 AS a") - subprocess.run(["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True) + subprocess.run( + ["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True + ) subprocess.run( [ "git", @@ -42,7 +48,9 @@ def test_git_uncommitted_changes(git_repo: Path): assert uncommitted[0].name == "model.sql" # stage the change and test that it is still detected - subprocess.run(["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True) + subprocess.run( + ["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True + ) uncommitted = git_client.list_uncommitted_changed_files() assert len(uncommitted) == 1 assert uncommitted[0].name == "model.sql" @@ -74,7 +82,9 @@ def test_git_both_staged_and_unstaged_changes(git_repo: Path): # stage file1 file1.write_text("SELECT 10") - subprocess.run(["git", "add", "model1.sql"], cwd=git_repo, check=True, capture_output=True) + subprocess.run( + ["git", "add", "model1.sql"], cwd=git_repo, check=True, capture_output=True + ) # modify file2 but don't stage it! file2.write_text("SELECT 20") @@ -90,7 +100,9 @@ def test_git_untracked_files(git_repo: Path): git_client = GitClient(git_repo) initial_file = git_repo / "initial.sql" initial_file.write_text("SELECT 0") - subprocess.run(["git", "add", "initial.sql"], cwd=git_repo, check=True, capture_output=True) + subprocess.run( + ["git", "add", "initial.sql"], cwd=git_repo, check=True, capture_output=True + ) subprocess.run( [ "git", @@ -124,7 +136,9 @@ def test_git_committed_changes(git_repo: Path): test_file = git_repo / "model.sql" test_file.write_text("SELECT 1") - subprocess.run(["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True) + subprocess.run( + ["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True + ) subprocess.run( [ "git", @@ -149,7 +163,9 @@ def test_git_committed_changes(git_repo: Path): ) test_file.write_text("SELECT 2") - subprocess.run(["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True) + subprocess.run( + ["git", "add", "model.sql"], cwd=git_repo, check=True, capture_output=True + ) subprocess.run( [ "git", diff --git a/tests/utils/test_helpers.py b/tests/utils/test_helpers.py index 20a544512e..f30538e149 100644 --- a/tests/utils/test_helpers.py +++ b/tests/utils/test_helpers.py @@ -1,22 +1,37 @@ -import pytest from functools import wraps + +import pytest from sqlglot import expressions from sqlglot.optimizer.annotate_types import annotate_types -from sqlmesh.core.console import set_console, get_console, TerminalConsole - +from sqlmesh.core.console import TerminalConsole, get_console, set_console from sqlmesh.utils import columns_to_types_all_known @pytest.mark.parametrize( "columns_to_types, expected", [ - ({"a": expressions.DataType.build("INT"), "b": expressions.DataType.build("INT")}, True), ( - {"a": expressions.DataType.build("UNKNOWN"), "b": expressions.DataType.build("INT")}, + { + "a": expressions.DataType.build("INT"), + "b": expressions.DataType.build("INT"), + }, + True, + ), + ( + { + "a": expressions.DataType.build("UNKNOWN"), + "b": expressions.DataType.build("INT"), + }, + False, + ), + ( + { + "a": expressions.DataType.build("NULL"), + "b": expressions.DataType.build("INT"), + }, False, ), - ({"a": expressions.DataType.build("NULL"), "b": expressions.DataType.build("INT")}, False), ( { "a": expressions.DataType.build("INT"), @@ -68,7 +83,11 @@ False, ), ( - {"a": annotate_types(expressions.DataType.build("VARCHAR(MAX)", dialect="redshift"))}, + { + "a": annotate_types( + expressions.DataType.build("VARCHAR(MAX)", dialect="redshift") + ) + }, True, ), ], diff --git a/tests/utils/test_jinja.py b/tests/utils/test_jinja.py index 01eeb47412..47cefda876 100644 --- a/tests/utils/test_jinja.py +++ b/tests/utils/test_jinja.py @@ -3,15 +3,9 @@ from base64 import b64encode from sqlmesh.utils import AttributeDict, yaml -from sqlmesh.utils.jinja import ( - ENVIRONMENT, - JinjaMacroRegistry, - MacroExtractor, - MacroReference, - MacroReturnVal, - call_name, - nodes, -) +from sqlmesh.utils.jinja import (ENVIRONMENT, JinjaMacroRegistry, + MacroExtractor, MacroReference, + MacroReturnVal, call_name, nodes) def test_macro_registry_render(): @@ -52,7 +46,10 @@ def test_macro_registry_render(): ] assert ( - extractor.extract("""{% set foo = bar | replace("'", "\\"") %}""", dialect="bigquery") == {} + extractor.extract( + """{% set foo = bar | replace("'", "\\"") %}""", dialect="bigquery" + ) + == {} ) @@ -70,7 +67,9 @@ def test_macro_registry_render_nested_self_package_references(): registry.add_macros(extractor.extract(package_a), package="package_a") - rendered = registry.build_environment().from_string("{{ package_a.macro_a_c() }}").render() + rendered = ( + registry.build_environment().from_string("{{ package_a.macro_a_c() }}").render() + ) assert rendered == "macro_a_a" @@ -86,7 +85,9 @@ def test_macro_registry_render_private_macros(): registry.add_macros(extractor.extract(package_a), package="package_a") - rendered = registry.build_environment().from_string("{{ package_a.macro_a_b() }}").render() + rendered = ( + registry.build_environment().from_string("{{ package_a.macro_a_b() }}").render() + ) assert rendered == "macro_a_a" @@ -173,7 +174,10 @@ def test_macro_registry_trim(): ) assert set(trimmed_registry_for_package_b.packages) == {"package_a", "package_b"} assert set(trimmed_registry_for_package_b.packages["package_a"]) == {"macro_a_a"} - assert set(trimmed_registry_for_package_b.packages["package_b"]) == {"macro_b_a", "macro_b_b"} + assert set(trimmed_registry_for_package_b.packages["package_b"]) == { + "macro_b_a", + "macro_b_b", + } assert not trimmed_registry_for_package_b.root_macros @@ -197,7 +201,9 @@ def macro_return(val): def test_global_objs(): - original_registry = JinjaMacroRegistry(global_objs={"target": AttributeDict({"test": "value"})}) + original_registry = JinjaMacroRegistry( + global_objs={"target": AttributeDict({"test": "value"})} + ) deserialized_registry = JinjaMacroRegistry.parse_raw(original_registry.json()) assert deserialized_registry.global_objs["target"].test == "value" @@ -280,7 +286,9 @@ def test_macro_registry_top_level_packages(): def test_find_call_names(): - jinja_str = "{{ local_macro() }}{{ package.package_macro() }}{{ 'stringval'.function() }}" + jinja_str = ( + "{{ local_macro() }}{{ package.package_macro() }}{{ 'stringval'.function() }}" + ) [call_name(node) for node in ENVIRONMENT.parse(jinja_str).find_all(nodes.Call)] == [ ("local_macro",), ("package", "package_macro"), @@ -302,7 +310,9 @@ def test_dbt_adapter_macro_scope(): registry.add_macros(macros, package="package_a") - rendered = registry.build_environment().from_string("{{ spark__macro_a() }}").render() + rendered = ( + registry.build_environment().from_string("{{ spark__macro_a() }}").render() + ) assert rendered.strip() == "macro_a" @@ -314,7 +324,11 @@ def test_macro_registry_to_expressions_sorted(): "schema": "main", "nested": {"foo": "bar", "baz": "bing"}, }, - "orders": {"schema": "main", "database": "jaffle_shop", "nested_list": ["b", "a", "c"]}, + "orders": { + "schema": "main", + "database": "jaffle_shop", + "nested_list": ["b", "a", "c"], + }, } ) @@ -339,7 +353,9 @@ def test_builtin_base64_filters(): env = JinjaMacroRegistry().build_environment() assert env.from_string("{{ value | b64decode }}").render(value=encoded) == "secret" assert env.from_string("{{ 'secret' | b64encode }}").render() == encoded - assert env.from_string("{{ 'secret' | b64encode | b64decode }}").render() == "secret" + assert ( + env.from_string("{{ 'secret' | b64encode | b64decode }}").render() == "secret" + ) # The same filters are available when rendering Jinja in config YAML files. config = yaml.load(f'env_vars:\n TOKEN: "{{{{ "{encoded}" | b64decode }}}}"') @@ -349,7 +365,9 @@ def test_builtin_base64_filters(): def test_builtin_b64decode_with_env_var(monkeypatch): # Real-world use case: a base64-encoded secret stored in an environment variable # is decoded inline in config YAML via env_var(...) piped through b64decode. - monkeypatch.setenv("SNOWFLAKE_PW_B64", b64encode(b"super-secret-pw").decode("utf-8")) + monkeypatch.setenv( + "SNOWFLAKE_PW_B64", b64encode(b"super-secret-pw").decode("utf-8") + ) config = yaml.load("password: \"{{ env_var('SNOWFLAKE_PW_B64') | b64decode }}\"") assert config == {"password": "super-secret-pw"} diff --git a/tests/utils/test_metaprogramming.py b/tests/utils/test_metaprogramming.py index 1f1431b963..a1740f4f6b 100644 --- a/tests/utils/test_metaprogramming.py +++ b/tests/utils/test_metaprogramming.py @@ -1,11 +1,10 @@ import ast +import re import typing as t from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path -from tenacity import retry, stop_after_attempt -import re import pandas as pd # noqa: TID253 import pytest import sqlglot @@ -14,24 +13,18 @@ from sqlglot import exp as expressions from sqlglot.expressions import SQLGLOT_META, to_table from sqlglot.optimizer.pushdown_projections import SELECT_ALL +from tenacity import retry, stop_after_attempt import tests.utils.test_date as test_date -from sqlmesh.core.dialect import normalize_model_name from sqlmesh.core import constants as c +from sqlmesh.core.dialect import normalize_model_name from sqlmesh.core.macros import RuntimeStage from sqlmesh.utils.errors import SQLMeshError -from sqlmesh.utils.metaprogramming import ( - Executable, - ExecutableKind, - _dict_sort, - _resolve_import_module, - build_env, - func_globals, - normalize_source, - prepare_env, - print_exception, - serialize_env, -) +from sqlmesh.utils.metaprogramming import (Executable, ExecutableKind, + _dict_sort, _resolve_import_module, + build_env, func_globals, + normalize_source, prepare_env, + print_exception, serialize_env) def test_print_exception(mocker: MockerFixture): @@ -236,9 +229,7 @@ def closure(z: int): return closure(y) + other_func(Y)""" ) - assert ( - normalize_source(other_func) - == """def other_func(a: int): + assert normalize_source(other_func) == """def other_func(a: int): import sqlglot sqlglot.parse_one('1') pd.DataFrame([{'x': 1}]) @@ -246,7 +237,6 @@ def closure(z: int): my_lambda() obj = MyClass(a) return X + a + W + obj.compute_with_reference()""" - ) def test_serialize_env_error() -> None: @@ -290,7 +280,8 @@ def closure(z: int): "Z": Executable(payload="3", kind=ExecutableKind.VALUE), "W": Executable(payload="0", kind=ExecutableKind.VALUE), "_GeneratorContextManager": Executable( - payload="from contextlib import _GeneratorContextManager", kind=ExecutableKind.IMPORT + payload="from contextlib import _GeneratorContextManager", + kind=ExecutableKind.IMPORT, ), "contextmanager": Executable( payload="from contextlib import contextmanager", kind=ExecutableKind.IMPORT @@ -354,9 +345,12 @@ def get_value(self): ), "pd": Executable(payload="import pandas as pd", kind=ExecutableKind.IMPORT), "sqlglot": Executable(kind=ExecutableKind.IMPORT, payload="import sqlglot"), - "exp": Executable(kind=ExecutableKind.IMPORT, payload="import sqlglot.expressions as exp"), + "exp": Executable( + kind=ExecutableKind.IMPORT, payload="import sqlglot.expressions as exp" + ), "expressions": Executable( - kind=ExecutableKind.IMPORT, payload="import sqlglot.expressions as expressions" + kind=ExecutableKind.IMPORT, + payload="import sqlglot.expressions as expressions", ), "func": Executable( payload="""@contextmanager @@ -395,11 +389,16 @@ def sample_context_manager(): name="sample_context_manager", path="test_metaprogramming.py", ), - "wraps": Executable(payload="from functools import wraps", kind=ExecutableKind.IMPORT), + "wraps": Executable( + payload="from functools import wraps", kind=ExecutableKind.IMPORT + ), "functools": Executable(payload="import functools", kind=ExecutableKind.IMPORT), - "retry": Executable(payload="from tenacity import retry", kind=ExecutableKind.IMPORT), + "retry": Executable( + payload="from tenacity import retry", kind=ExecutableKind.IMPORT + ), "stop_after_attempt": Executable( - payload="from tenacity.stop import stop_after_attempt", kind=ExecutableKind.IMPORT + payload="from tenacity.stop import stop_after_attempt", + kind=ExecutableKind.IMPORT, ), "wrapped_f": Executable( payload='''@retry(stop=stop_after_attempt(3)) @@ -463,7 +462,9 @@ def function_with_custom_decorator(): serialized_env = serialize_env(env, path=path) # type: ignore assert prepare_env(serialized_env) - expected_env = {k: Executable(**v.dict(), is_metadata=True) for k, v in expected_env.items()} + expected_env = { + k: Executable(**v.dict(), is_metadata=True) for k, v in expected_env.items() + } # Every object is treated as "metadata only", transitively assert all(is_metadata for (_, is_metadata) in env.values()) @@ -497,7 +498,8 @@ def test_serialize_env_with_enum_import_appearing_in_two_functions() -> None: expected_env = { "RuntimeStage": Executable( - payload="from sqlmesh.core.macros import RuntimeStage", kind=ExecutableKind.IMPORT + payload="from sqlmesh.core.macros import RuntimeStage", + kind=ExecutableKind.IMPORT, ), "macro1": Executable( payload="""def macro1(): @@ -562,17 +564,29 @@ def test_dict_sort_mixed_key_types(): def test_dict_sort_nested_structures(): """Test dict_sort with deeply nested dictionaries.""" - nested1 = {"outer": {"z": 26, "a": 1}, "list": [3, {"y": 2, "x": 1}], "simple": "value"} + nested1 = { + "outer": {"z": 26, "a": 1}, + "list": [3, {"y": 2, "x": 1}], + "simple": "value", + } - nested2 = {"simple": "value", "list": [3, {"x": 1, "y": 2}], "outer": {"a": 1, "z": 26}} + nested2 = { + "simple": "value", + "list": [3, {"x": 1, "y": 2}], + "outer": {"a": 1, "z": 26}, + } repr1 = _dict_sort(nested1) repr2 = _dict_sort(nested2) assert repr1 != repr2 # Verify structure is maintained with sorted keys - expected1 = "{'list': [3, {'y': 2, 'x': 1}], 'outer': {'z': 26, 'a': 1}, 'simple': 'value'}" - expected2 = "{'list': [3, {'x': 1, 'y': 2}], 'outer': {'a': 1, 'z': 26}, 'simple': 'value'}" + expected1 = ( + "{'list': [3, {'y': 2, 'x': 1}], 'outer': {'z': 26, 'a': 1}, 'simple': 'value'}" + ) + expected2 = ( + "{'list': [3, {'x': 1, 'y': 2}], 'outer': {'a': 1, 'z': 26}, 'simple': 'value'}" + ) assert repr1 == expected1 assert repr2 == expected2 diff --git a/tests/utils/test_pydantic.py b/tests/utils/test_pydantic.py index 9234589218..e655edcde4 100644 --- a/tests/utils/test_pydantic.py +++ b/tests/utils/test_pydantic.py @@ -1,15 +1,13 @@ import typing as t -import pytest from functools import cached_property import pydantic +import pytest from sqlmesh.utils.date import TimeLike, to_date, to_datetime -from sqlmesh.utils.pydantic import ( - PydanticModel, - get_concrete_types_from_typehint, - validation_error_message, -) +from sqlmesh.utils.pydantic import (PydanticModel, + get_concrete_types_from_typehint, + validation_error_message) def test_datetime_date_serialization() -> None: diff --git a/tests/utils/test_windows.py b/tests/utils/test_windows.py index 196589d9c2..bf6a866dcb 100644 --- a/tests/utils/test_windows.py +++ b/tests/utils/test_windows.py @@ -1,6 +1,9 @@ -import pytest from pathlib import Path -from sqlmesh.utils.windows import IS_WINDOWS, WINDOWS_LONGPATH_PREFIX, fix_windows_path + +import pytest + +from sqlmesh.utils.windows import (IS_WINDOWS, WINDOWS_LONGPATH_PREFIX, + fix_windows_path) @pytest.mark.skipif( @@ -31,9 +34,13 @@ def test_fix_windows_path(): # paths with relative sections need to have relative sections resolved before they can be used # since the \\?\ prefix doesnt work for paths with relative sections - assert fix_windows_path(Path("c:\\foo\\..\\bar")) == Path(WINDOWS_LONGPATH_PREFIX + "c:\\bar") + assert fix_windows_path(Path("c:\\foo\\..\\bar")) == Path( + WINDOWS_LONGPATH_PREFIX + "c:\\bar" + ) # also check that relative sections are still resolved if they are added to a previously prefixed path base = fix_windows_path(Path("c:\\foo")) assert base == Path(WINDOWS_LONGPATH_PREFIX + "c:\\foo") - assert fix_windows_path(base / ".." / "bar") == Path(WINDOWS_LONGPATH_PREFIX + "c:\\bar") + assert fix_windows_path(base / ".." / "bar") == Path( + WINDOWS_LONGPATH_PREFIX + "c:\\bar" + ) diff --git a/tests/utils/test_yaml.py b/tests/utils/test_yaml.py index 5a2e04e5be..937309b509 100644 --- a/tests/utils/test_yaml.py +++ b/tests/utils/test_yaml.py @@ -1,7 +1,7 @@ import os +from decimal import Decimal import pytest -from decimal import Decimal import sqlmesh.utils.yaml as yaml from sqlmesh.utils.errors import SQLMeshError @@ -73,7 +73,9 @@ def test_load_keep_last_duplicate_key() -> None: } # Test keeping last key - assert yaml.load(input_str, allow_duplicate_keys=True, keep_last_duplicate_key=True) == { + assert yaml.load( + input_str, allow_duplicate_keys=True, keep_last_duplicate_key=True + ) == { "name": "third_name", "foo": "bar", "mapping": {"key": "third_value"}, diff --git a/tests/web/conftest.py b/tests/web/conftest.py index 6b6fcaad29..a286208f00 100644 --- a/tests/web/conftest.py +++ b/tests/web/conftest.py @@ -3,9 +3,8 @@ import pytest from fastapi import FastAPI -from sqlmesh.core.context import Context from sqlmesh.core.console import set_console - +from sqlmesh.core.context import Context from web.server.console import api_console from web.server.settings import Settings, get_loaded_context, get_settings @@ -23,11 +22,9 @@ def get_settings_override() -> Settings: return Settings(project_path=tmp_path) config = tmp_path / "config.py" - config.write_text( - """from sqlmesh.core.config import Config, ModelDefaultsConfig + config.write_text("""from sqlmesh.core.config import Config, ModelDefaultsConfig config = Config(model_defaults=ModelDefaultsConfig(dialect='')) - """ - ) + """) web_app.dependency_overrides[get_settings] = get_settings_override yield tmp_path diff --git a/tests/web/test_lineage.py b/tests/web/test_lineage.py index 0cffd3ecc3..cd248b4df7 100644 --- a/tests/web/test_lineage.py +++ b/tests/web/test_lineage.py @@ -94,7 +94,9 @@ def test_get_lineage(client: TestClient, web_sushi_context: Context) -> None: } -def test_get_lineage_managed_columns(client: TestClient, web_sushi_context: Context) -> None: +def test_get_lineage_managed_columns( + client: TestClient, web_sushi_context: Context +) -> None: # Get lineage of managed column response = client.get("/api/lineage/sushi.marketing/valid_from") assert response.status_code == 200 @@ -126,7 +128,9 @@ def test_get_lineage_single_model(client: TestClient, project_context: Context) assert response_json['"bar"']["col"]["models"] == {} -def test_get_lineage_external_model(client: TestClient, project_context: Context) -> None: +def test_get_lineage_external_model( + client: TestClient, project_context: Context +) -> None: project_tmp_path = project_context.path models_dir = project_tmp_path / "models" models_dir.mkdir() @@ -151,23 +155,17 @@ def test_get_lineage_cte(client: TestClient, project_context: Context) -> None: models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); + foo_sql_file.write_text("""MODEL (name foo); WITH my_cte AS ( SELECT col FROM bar ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT col FROM baz;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT col FROM baz;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_table;""" - ) + baz_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_table;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -187,28 +185,24 @@ def test_get_lineage_cte(client: TestClient, project_context: Context) -> None: assert response_json['"baz"']["col"]["models"] == {'"external_table"': ["col"]} -def test_get_lineage_cte_downstream(client: TestClient, project_context: Context) -> None: +def test_get_lineage_cte_downstream( + client: TestClient, project_context: Context +) -> None: project_tmp_path = project_context.path models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); - SELECT col FROM bar;""" - ) + foo_sql_file.write_text("""MODEL (name foo); + SELECT col FROM bar;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); + bar_sql_file.write_text("""MODEL (name bar); WITH my_cte AS ( SELECT col FROM baz ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_table;""" - ) + baz_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_table;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -238,39 +232,38 @@ def test_get_lineage_join(client: TestClient, project_context: Context) -> None: SELECT id, bar.quantity * baz.price AS col FROM bar JOIN baz ON bar.id = baz.id;""" ) bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT id, quantity FROM external_bar;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT id, quantity FROM external_bar;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); - SELECT id, price FROM external_baz;""" - ) + baz_sql_file.write_text("""MODEL (name baz); + SELECT id, price FROM external_baz;""") project_context.load() response = client.get("/api/lineage/foo/col") assert response.status_code == 200, response.json() response_json = response.json() - assert response_json['"foo"']["col"]["models"] == {'"bar"': ["quantity"], '"baz"': ["price"]} - assert response_json['"bar"']["quantity"]["models"] == {'"external_bar"': ["quantity"]} + assert response_json['"foo"']["col"]["models"] == { + '"bar"': ["quantity"], + '"baz"': ["price"], + } + assert response_json['"bar"']["quantity"]["models"] == { + '"external_bar"': ["quantity"] + } assert response_json['"baz"']["price"]["models"] == {'"external_baz"': ["price"]} -def test_get_lineage_multiple_columns(client: TestClient, project_context: Context) -> None: +def test_get_lineage_multiple_columns( + client: TestClient, project_context: Context +) -> None: project_tmp_path = project_context.path models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); - SELECT id, bar.value * bar.multiplier AS col FROM bar;""" - ) + foo_sql_file.write_text("""MODEL (name foo); + SELECT id, bar.value * bar.multiplier AS col FROM bar;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT id, value, multiplier FROM external_bar;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT id, value, multiplier FROM external_bar;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -279,7 +272,9 @@ def test_get_lineage_multiple_columns(client: TestClient, project_context: Conte assert "value" in response_json['"foo"']["col"]["models"]['"bar"'] assert "multiplier" in response_json['"foo"']["col"]["models"]['"bar"'] assert response_json['"bar"']["value"]["models"] == {'"external_bar"': ["value"]} - assert response_json['"bar"']["multiplier"]["models"] == {'"external_bar"': ["multiplier"]} + assert response_json['"bar"']["multiplier"]["models"] == { + '"external_bar"': ["multiplier"] + } def test_get_lineage_union(client: TestClient, project_context: Context) -> None: @@ -287,69 +282,66 @@ def test_get_lineage_union(client: TestClient, project_context: Context) -> None models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); + foo_sql_file.write_text("""MODEL (name foo); SELECT col FROM bar UNION - SELECT col FROM baz;""" - ) + SELECT col FROM baz;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT col FROM external_bar;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT col FROM external_bar;""") bar_sql_file = models_dir / "baz.sql" - bar_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_baz;""" - ) + bar_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_baz;""") project_context.load() response = client.get("/api/lineage/foo/col") assert response.status_code == 200, response.json() response_json = response.json() - assert response_json['"foo"']["col"]["models"] == {'"bar"': ["col"], '"baz"': ["col"]} + assert response_json['"foo"']["col"]["models"] == { + '"bar"': ["col"], + '"baz"': ["col"], + } # Models only response = client.get("/api/lineage/foo/col?models_only=1") assert response.status_code == 200, response.json() response_json = response.json() - assert response_json['"foo"']["col"]["models"] == {'"bar"': ["col"], '"baz"': ["col"]} + assert response_json['"foo"']["col"]["models"] == { + '"bar"': ["col"], + '"baz"': ["col"], + } -def test_get_lineage_union_downstream(client: TestClient, project_context: Context) -> None: +def test_get_lineage_union_downstream( + client: TestClient, project_context: Context +) -> None: project_tmp_path = project_context.path models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); - SELECT col FROM bar;""" - ) + foo_sql_file.write_text("""MODEL (name foo); + SELECT col FROM bar;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); + bar_sql_file.write_text("""MODEL (name bar); SELECT col FROM baz UNION - SELECT col FROM qwe;""" - ) + SELECT col FROM qwe;""") bar_sql_file = models_dir / "baz.sql" - bar_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_baz;""" - ) + bar_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_baz;""") bar_sql_file = models_dir / "qwe.sql" - bar_sql_file.write_text( - """MODEL (name qwe); - SELECT col FROM external_qwe;""" - ) + bar_sql_file.write_text("""MODEL (name qwe); + SELECT col FROM external_qwe;""") project_context.load() response = client.get("/api/lineage/foo/col") assert response.status_code == 200, response.json() response_json = response.json() assert response_json['"foo"']["col"]["models"] == {'"bar"': ["col"]} - assert response_json['"bar"']["col"]["models"] == {'"baz"': ["col"], '"qwe"': ["col"]} + assert response_json['"bar"']["col"]["models"] == { + '"baz"': ["col"], + '"qwe"': ["col"], + } assert response_json['"baz"']["col"]["models"] == {'"external_baz"': ["col"]} assert response_json['"qwe"']["col"]["models"] == {'"external_qwe"': ["col"]} @@ -358,7 +350,10 @@ def test_get_lineage_union_downstream(client: TestClient, project_context: Conte assert response.status_code == 200, response.json() response_json = response.json() assert response_json['"foo"']["col"]["models"] == {'"bar"': ["col"]} - assert response_json['"bar"']["col"]["models"] == {'"baz"': ["col"], '"qwe"': ["col"]} + assert response_json['"bar"']["col"]["models"] == { + '"baz"': ["col"], + '"qwe"': ["col"], + } assert response_json['"baz"']["col"]["models"] == {'"external_baz"': ["col"]} assert response_json['"qwe"']["col"]["models"] == {'"external_qwe"': ["col"]} @@ -368,32 +363,29 @@ def test_get_lineage_cte_union(client: TestClient, project_context: Context) -> models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); + foo_sql_file.write_text("""MODEL (name foo); WITH my_cte AS ( SELECT col FROM bar UNION SELECT col FROM baz ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT col FROM external_bar;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT col FROM external_bar;""") bar_sql_file = models_dir / "baz.sql" - bar_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_baz;""" - ) + bar_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_baz;""") project_context.load() response = client.get("/api/lineage/foo/col") assert response.status_code == 200, response.json() response_json = response.json() assert response_json['"foo"']["col"]["models"] == {'"foo": my_cte': ["col"]} - assert response_json['"foo": my_cte']["col"]["models"] == {'"bar"': ["col"], '"baz"': ["col"]} + assert response_json['"foo": my_cte']["col"]["models"] == { + '"bar"': ["col"], + '"baz"': ["col"], + } assert response_json['"bar"']["col"]["models"] == {'"external_bar"': ["col"]} assert response_json['"baz"']["col"]["models"] == {'"external_baz"': ["col"]} @@ -401,40 +393,37 @@ def test_get_lineage_cte_union(client: TestClient, project_context: Context) -> response = client.get("/api/lineage/foo/col?models_only=1") assert response.status_code == 200, response.json() response_json = response.json() - assert response_json['"foo"']["col"]["models"] == {'"bar"': ["col"], '"baz"': ["col"]} + assert response_json['"foo"']["col"]["models"] == { + '"bar"': ["col"], + '"baz"': ["col"], + } assert response_json['"bar"']["col"]["models"] == {'"external_bar"': ["col"]} assert response_json['"baz"']["col"]["models"] == {'"external_baz"': ["col"]} -def test_get_lineage_cte_union_downstream(client: TestClient, project_context: Context) -> None: +def test_get_lineage_cte_union_downstream( + client: TestClient, project_context: Context +) -> None: project_tmp_path = project_context.path models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); - SELECT col FROM bar;""" - ) + foo_sql_file.write_text("""MODEL (name foo); + SELECT col FROM bar;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); + bar_sql_file.write_text("""MODEL (name bar); WITH my_cte AS ( SELECT col FROM baz UNION SELECT col FROM qwe ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_baz;""" - ) + baz_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_baz;""") baz_sql_file = models_dir / "qwe.sql" - baz_sql_file.write_text( - """MODEL (name qwe); - SELECT col FROM external_qwe;""" - ) + baz_sql_file.write_text("""MODEL (name qwe); + SELECT col FROM external_qwe;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -442,7 +431,10 @@ def test_get_lineage_cte_union_downstream(client: TestClient, project_context: C response_json = response.json() assert response_json['"foo"']["col"]["models"] == {'"bar"': ["col"]} assert response_json['"bar"']["col"]["models"] == {'"bar": my_cte': ["col"]} - assert response_json['"bar": my_cte']["col"]["models"] == {'"baz"': ["col"], '"qwe"': ["col"]} + assert response_json['"bar": my_cte']["col"]["models"] == { + '"baz"': ["col"], + '"qwe"': ["col"], + } assert response_json['"baz"']["col"]["models"] == {'"external_baz"': ["col"]} assert response_json['"qwe"']["col"]["models"] == {'"external_qwe"': ["col"]} @@ -451,7 +443,10 @@ def test_get_lineage_cte_union_downstream(client: TestClient, project_context: C assert response.status_code == 200, response.json() response_json = response.json() assert response_json['"foo"']["col"]["models"] == {'"bar"': ["col"]} - assert response_json['"bar"']["col"]["models"] == {'"baz"': ["col"], '"qwe"': ["col"]} + assert response_json['"bar"']["col"]["models"] == { + '"baz"': ["col"], + '"qwe"': ["col"], + } assert response_json['"baz"']["col"]["models"] == {'"external_baz"': ["col"]} assert response_json['"qwe"']["col"]["models"] == {'"external_qwe"': ["col"]} @@ -463,25 +458,19 @@ def test_get_lineage_cte_downstream_union_downstream( models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); - SELECT col FROM bar;""" - ) + foo_sql_file.write_text("""MODEL (name foo); + SELECT col FROM bar;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); + bar_sql_file.write_text("""MODEL (name bar); WITH my_cte AS ( SELECT * FROM baz ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); + baz_sql_file.write_text("""MODEL (name baz); SELECT col FROM external_table1 UNION - SELECT col FROM external_table2;""" - ) + SELECT col FROM external_table2;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -514,8 +503,7 @@ def test_get_lineage_nested_cte_union_downstream( models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); + foo_sql_file.write_text("""MODEL (name foo); WITH my_cte2 AS ( SELECT * FROM bar ), my_cte1 AS ( @@ -523,16 +511,13 @@ def test_get_lineage_nested_cte_union_downstream( UNION SELECT col FROM external_table1 ) - SELECT col FROM my_cte1;""" - ) + SELECT col FROM my_cte1;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); + bar_sql_file.write_text("""MODEL (name bar); SELECT col FROM external_table2 UNION SELECT col FROM external_table3 - ;""" - ) + ;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -568,20 +553,14 @@ def test_get_lineage_subquery(client: TestClient, project_context: Context) -> N models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); - SELECT col FROM (SELECT col FROM bar) my_dt;""" - ) + foo_sql_file.write_text("""MODEL (name foo); + SELECT col FROM (SELECT col FROM bar) my_dt;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT col FROM baz;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT col FROM baz;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_table;""" - ) + baz_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_table;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -601,31 +580,27 @@ def test_get_lineage_subquery(client: TestClient, project_context: Context) -> N assert response_json['"baz"']["col"]["models"] == {'"external_table"': ["col"]} -def test_get_lineage_cte_name_collision(client: TestClient, project_context: Context) -> None: +def test_get_lineage_cte_name_collision( + client: TestClient, project_context: Context +) -> None: project_tmp_path = project_context.path models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); + foo_sql_file.write_text("""MODEL (name foo); WITH my_cte AS ( SELECT col FROM bar ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); + bar_sql_file.write_text("""MODEL (name bar); WITH my_cte AS ( SELECT col FROM baz ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_table;""" - ) + baz_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_table;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -653,20 +628,14 @@ def test_get_lineage_derived_table_alias_collision( models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); - SELECT col FROM (SELECT col FROM bar) my_dt;""" - ) + foo_sql_file.write_text("""MODEL (name foo); + SELECT col FROM (SELECT col FROM bar) my_dt;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT col FROM (SELECT col FROM baz) my_dt;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT col FROM (SELECT col FROM baz) my_dt;""") baz_sql_file = models_dir / "baz.sql" - baz_sql_file.write_text( - """MODEL (name baz); - SELECT col FROM external_table;""" - ) + baz_sql_file.write_text("""MODEL (name baz); + SELECT col FROM external_table;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -692,8 +661,7 @@ def test_get_lineage_constants(client: TestClient, project_context: Context) -> models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); + foo_sql_file.write_text("""MODEL (name foo); WITH my_cte AS ( SELECT col FROM bar UNION @@ -701,13 +669,10 @@ def test_get_lineage_constants(client: TestClient, project_context: Context) -> UNION SELECT 1 as col FROM external_table ) - SELECT col FROM my_cte;""" - ) + SELECT col FROM my_cte;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT col FROM external_table;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT col FROM external_table;""") project_context.load() response = client.get("/api/lineage/foo/col") @@ -725,13 +690,14 @@ def test_get_lineage_constants(client: TestClient, project_context: Context) -> assert response_json['"bar"']["col"]["models"] == {'"external_table"': ["col"]} -def test_get_lineage_quoted_columns(client: TestClient, project_context: Context) -> None: +def test_get_lineage_quoted_columns( + client: TestClient, project_context: Context +) -> None: project_tmp_path = project_context.path models_dir = project_tmp_path / "models" models_dir.mkdir() foo_sql_file = models_dir / "foo.sql" - foo_sql_file.write_text( - """MODEL (name foo); + foo_sql_file.write_text("""MODEL (name foo); WITH my_cte AS ( SELECT col as "@col" FROM bar UNION @@ -739,13 +705,10 @@ def test_get_lineage_quoted_columns(client: TestClient, project_context: Context UNION SELECT 1 as "@col" FROM external_table ) - SELECT "@col" FROM my_cte;""" - ) + SELECT "@col" FROM my_cte;""") bar_sql_file = models_dir / "bar.sql" - bar_sql_file.write_text( - """MODEL (name bar); - SELECT col FROM external_table;""" - ) + bar_sql_file.write_text("""MODEL (name bar); + SELECT col FROM external_table;""") project_context.load() response = client.get("/api/lineage/foo/@col") diff --git a/tests/web/test_main.py b/tests/web/test_main.py index b20947c49d..f671baae34 100644 --- a/tests/web/test_main.py +++ b/tests/web/test_main.py @@ -82,11 +82,9 @@ def test_get_file_not_found(client: TestClient) -> None: def test_get_file_invalid_path(client: TestClient, project_tmp_path: Path) -> None: config = project_tmp_path / "config.py" - config.write_text( - """from sqlmesh.core.config import Config, ModelDefaultsConfig + config.write_text("""from sqlmesh.core.config import Config, ModelDefaultsConfig config = Config(ignore_patterns=["*.txt"], model_defaults=ModelDefaultsConfig(dialect='')) - """ - ) + """) foo_txt = project_tmp_path / "foo.txt" foo_txt.touch() @@ -149,7 +147,9 @@ def test_rename_file(client: TestClient, project_tmp_path: Path) -> None: assert not txt_file.exists() -def test_rename_file_and_keep_content(client: TestClient, project_tmp_path: Path) -> None: +def test_rename_file_and_keep_content( + client: TestClient, project_tmp_path: Path +) -> None: txt_file = project_tmp_path / "foo.txt" txt_file.write_text("bar") @@ -190,7 +190,9 @@ def test_rename_file_already_exists(client: TestClient, project_tmp_path: Path) assert not foo_file.exists() -def test_rename_file_to_existing_directory(client: TestClient, project_tmp_path: Path) -> None: +def test_rename_file_to_existing_directory( + client: TestClient, project_tmp_path: Path +) -> None: foo_file = project_tmp_path / "foo.txt" foo_file.touch() existing_dir = project_tmp_path / "existing_dir" @@ -232,7 +234,9 @@ def test_create_directory(client: TestClient, project_tmp_path: Path) -> None: } -def test_create_directory_already_exists(client: TestClient, project_tmp_path: Path) -> None: +def test_create_directory_already_exists( + client: TestClient, project_tmp_path: Path +) -> None: new_dir = project_tmp_path / "new_dir" new_dir.mkdir() @@ -257,7 +261,9 @@ def test_rename_directory(client: TestClient, project_tmp_path: Path) -> None: } -def test_rename_directory_already_exists_empty(client: TestClient, project_tmp_path: Path) -> None: +def test_rename_directory_already_exists_empty( + client: TestClient, project_tmp_path: Path +) -> None: new_dir = project_tmp_path / "new_dir" new_dir.mkdir() existing_dir = project_tmp_path / "renamed_dir" @@ -291,7 +297,9 @@ def test_rename_directory_already_exists_not_empty( assert new_dir.exists() -def test_rename_directory_to_existing_file(client: TestClient, project_tmp_path: Path) -> None: +def test_rename_directory_to_existing_file( + client: TestClient, project_tmp_path: Path +) -> None: new_dir = project_tmp_path / "new_dir" new_dir.mkdir() existing_file = project_tmp_path / "foo.txt" @@ -317,7 +325,9 @@ def test_delete_directory_not_found(client: TestClient, project_tmp_path: Path) assert response.status_code == 404 -def test_delete_directory_not_a_directory(client: TestClient, project_tmp_path: Path) -> None: +def test_delete_directory_not_a_directory( + client: TestClient, project_tmp_path: Path +) -> None: txt_file = project_tmp_path / "foo.txt" txt_file.touch() @@ -381,7 +391,9 @@ def test_plan_test_failures( async def test_cancel(client: TestClient) -> None: client.app.state.circuit_breaker = threading.Event() # type: ignore transport = ASGITransport(client.app) # type: ignore - async with AsyncClient(transport=transport, base_url="http://testserver") as _client: + async with AsyncClient( + transport=transport, base_url="http://testserver" + ) as _client: await _client.post("/api/plan", json={"environment": "dev"}) response = await _client.post("/api/plan/cancel") assert response.status_code == 204 @@ -430,7 +442,9 @@ def test_modules(client: TestClient) -> None: def test_fetchdf(client: TestClient, web_sushi_context: Context) -> None: - response = client.post("/api/commands/fetchdf", json={"sql": "SELECT * from sushi.top_waiters"}) + response = client.post( + "/api/commands/fetchdf", json={"sql": "SELECT * from sushi.top_waiters"} + ) assert response.status_code == 200 with pa.ipc.open_stream(response.content) as reader: df = reader.read_pandas() @@ -498,7 +512,9 @@ def test_get_environments(client: TestClient, project_context: Context) -> None: plan_id="", suffix_target="schema", ) - assert response_json["pinned_environments"] == list(project_context.config.pinned_environments) + assert response_json["pinned_environments"] == list( + project_context.config.pinned_environments + ) assert ( response_json["default_target_environment"] == project_context.config.default_target_environment @@ -554,7 +570,9 @@ def test_test(client: TestClient, web_sushi_context: Context) -> None: assert response_json["failures"] == [] # Single test - response = client.get("/api/commands/test", params={"test": "tests/test_order_items.yaml"}) + response = client.get( + "/api/commands/test", params={"test": "tests/test_order_items.yaml"} + ) assert response.status_code == 200 response_json = response.json() assert response_json["tests_run"] == 1 @@ -570,8 +588,7 @@ def test_test_failure(client: TestClient, project_context: Context) -> None: tests_dir = project_context.path / "tests" tests_dir.mkdir() test_file = tests_dir / "test_foo.yaml" - test_file.write_text( - """test_foo: + test_file.write_text("""test_foo: model: foo outputs: query: @@ -579,8 +596,7 @@ def test_test_failure(client: TestClient, project_context: Context) -> None: vars: start: 2022-01-01 end: 2022-01-01 - latest: 2022-01-01""" - ) + latest: 2022-01-01""") project_context.load() response = client.get("/api/commands/test") diff --git a/vscode/extension/tests/tcloud/mock_tcloud/__init__.py b/vscode/extension/tests/tcloud/mock_tcloud/__init__.py index 98ad152e0e..79e27d406e 100644 --- a/vscode/extension/tests/tcloud/mock_tcloud/__init__.py +++ b/vscode/extension/tests/tcloud/mock_tcloud/__init__.py @@ -1 +1 @@ -# Mock tcloud package \ No newline at end of file +# Mock tcloud package diff --git a/vscode/extension/tests/tcloud/mock_tcloud/cli.py b/vscode/extension/tests/tcloud/mock_tcloud/cli.py index 55de42ca81..3c613d349b 100755 --- a/vscode/extension/tests/tcloud/mock_tcloud/cli.py +++ b/vscode/extension/tests/tcloud/mock_tcloud/cli.py @@ -12,6 +12,7 @@ import click + def get_auth_state_file(): """Get the path to the auth state file in the current working directory""" return Path.cwd() / ".tcloud_auth_state.json" @@ -66,11 +67,11 @@ def cli(ctx: click.Context, project: str, version: bool) -> None: version_state = load_version_state() print(version_state["version"]) ctx.exit(0) - + if ctx.invoked_subcommand is None: click.echo(ctx.get_help()) ctx.exit(0) - + ctx.ensure_object(dict) ctx.obj["project"] = project @@ -143,7 +144,7 @@ def sqlmesh_lsp(ctx: click.Context, args) -> None: """Run SQLMesh LSP server""" # For testing purposes, we'll simulate the LSP server starting print("Starting SQLMesh LSP server...", flush=True) - + # Get the path to sqlmesh in the same environment as this script bin_dir = os.path.dirname(sys.executable) sqlmesh_path = os.path.join(bin_dir, "sqlmesh") diff --git a/web/client/src/workers/sqlglot/sqlglot.py b/web/client/src/workers/sqlglot/sqlglot.py index 435a79823f..a8edd77bf2 100644 --- a/web/client/src/workers/sqlglot/sqlglot.py +++ b/web/client/src/workers/sqlglot/sqlglot.py @@ -19,7 +19,9 @@ def parse_to_json(sql: str, read: DialectType = None) -> str: return json.dumps( [ exp.dump() if exp else {} - for exp in sqlglot.parse(sql, read=read, error_level=sqlglot.ErrorLevel.IGNORE) + for exp in sqlglot.parse( + sql, read=read, error_level=sqlglot.ErrorLevel.IGNORE + ) ] ) @@ -39,14 +41,19 @@ def get_dialect(name: str = "") -> str: def format(sql: str = "", read: DialectType = None) -> str: return "\n".join( - sqlglot.transpile(sql, read=read, error_level=sqlglot.errors.ErrorLevel.IGNORE, pretty=True) + sqlglot.transpile( + sql, read=read, error_level=sqlglot.errors.ErrorLevel.IGNORE, pretty=True + ) ) def validate(sql: str = "", read: DialectType = None) -> str: try: sqlglot.transpile( - sql, read=read, pretty=False, unsupported_level=sqlglot.errors.ErrorLevel.IMMEDIATE + sql, + read=read, + pretty=False, + unsupported_level=sqlglot.errors.ErrorLevel.IMMEDIATE, ) except sqlglot.errors.ParseError: return json.dumps(False) diff --git a/web/server/api/endpoints/__init__.py b/web/server/api/endpoints/__init__.py index db5d18b8c9..48b073d2b6 100644 --- a/web/server/api/endpoints/__init__.py +++ b/web/server/api/endpoints/__init__.py @@ -1,18 +1,8 @@ from fastapi import APIRouter -from web.server.api.endpoints import ( - commands, - directories, - environments, - events, - files, - lineage, - meta, - models, - modules, - plan, - table_diff, -) +from web.server.api.endpoints import (commands, directories, environments, + events, files, lineage, meta, models, + modules, plan, table_diff) api_router = APIRouter() api_router.include_router(commands.router, prefix="/commands") diff --git a/web/server/api/endpoints/commands.py b/web/server/api/endpoints/commands.py index 5db3c85d66..eb2e7459c3 100644 --- a/web/server/api/endpoints/commands.py +++ b/web/server/api/endpoints/commands.py @@ -2,9 +2,11 @@ import asyncio import typing as t +from pathlib import Path from fastapi import APIRouter, Body, Depends, Request, Response from starlette.status import HTTP_204_NO_CONTENT + from sqlmesh.core.console import Verbosity from sqlmesh.core.context import Context from sqlmesh.core.snapshot.definition import SnapshotChangeCategory @@ -15,12 +17,8 @@ from web.server.console import api_console from web.server.exceptions import ApiException from web.server.settings import get_loaded_context -from web.server.utils import ( - ArrowStreamingResponse, - df_to_pyarrow_bytes, - run_in_executor, -) -from pathlib import Path +from web.server.utils import (ArrowStreamingResponse, df_to_pyarrow_bytes, + run_in_executor) router = APIRouter() @@ -154,7 +152,9 @@ async def test( message="Unable to run tests", origin="API -> commands -> test", ) - context.console.log_test_results(result, context.test_connection_config._engine_adapter.DIALECT) + context.console.log_test_results( + result, context.test_connection_config._engine_adapter.DIALECT + ) def _test_path(test: ModelTest) -> t.Optional[str]: if path := test.path_relative_to(context.path): @@ -168,7 +168,9 @@ def _test_path(test: ModelTest) -> t.Optional[str]: path=_test_path(test), tb=tb, ) - for test, tb in ((t.cast(ModelTest, test), tb) for test, tb in result.errors) + for test, tb in ( + (t.cast(ModelTest, test), tb) for test, tb in result.errors + ) ], failures=[ models.TestErrorOrFailure( @@ -176,7 +178,9 @@ def _test_path(test: ModelTest) -> t.Optional[str]: path=_test_path(test), tb=tb, ) - for test, tb in ((t.cast(ModelTest, test), tb) for test, tb in result.failures) + for test, tb in ( + (t.cast(ModelTest, test), tb) for test, tb in result.failures + ) ], skipped=[ models.TestSkipped( @@ -236,9 +240,13 @@ def _run_plan_apply( ) -> None: """Run plan apply""" plan_options = plan_options or models.PlanOptions() - tracker_apply = models.PlanApplyStageTracker(environment=environment, plan_options=plan_options) + tracker_apply = models.PlanApplyStageTracker( + environment=environment, plan_options=plan_options + ) api_console.start_plan_tracker(tracker_apply) - plan_builder = get_plan_builder(context, plan_options, environment, plan_dates, categories) + plan_builder = get_plan_builder( + context, plan_options, environment, plan_dates, categories + ) plan = plan_builder.build() tracker_apply.start = plan.start tracker_apply.end = plan.end diff --git a/web/server/api/endpoints/environments.py b/web/server/api/endpoints/environments.py index 1598e02cc4..7e600c4760 100644 --- a/web/server/api/endpoints/environments.py +++ b/web/server/api/endpoints/environments.py @@ -17,7 +17,9 @@ async def get_environments( ) -> Environments: """Get the environments""" try: - environments = {env.name: env for env in context.state_reader.get_environments()} + environments = { + env.name: env for env in context.state_reader.get_environments() + } except Exception: raise ApiException( message="Unable to get environments", diff --git a/web/server/api/endpoints/files.py b/web/server/api/endpoints/files.py index db58fce55e..4fb5927b1d 100644 --- a/web/server/api/endpoints/files.py +++ b/web/server/api/endpoints/files.py @@ -13,12 +13,8 @@ from web.server import models from web.server.console import api_console from web.server.exceptions import ApiException -from web.server.settings import ( - Settings, - get_context, - get_path_to_model_mapping, - get_settings, -) +from web.server.settings import (Settings, get_context, + get_path_to_model_mapping, get_settings) from web.server.utils import replace_file, validate_path router = APIRouter() @@ -57,10 +53,14 @@ async def write_file( path_or_new_path = path if new_path: path_or_new_path = validate_path(new_path, settings) - replace_file(settings.project_path / path, settings.project_path / path_or_new_path) + replace_file( + settings.project_path / path, settings.project_path / path_or_new_path + ) else: full_path = settings.project_path / path - config, _ = context.config_for_path(Path(path_or_new_path)) if context else (None, None) + config, _ = ( + context.config_for_path(Path(path_or_new_path)) if context else (None, None) + ) if ( config and config.ui.format_on_save @@ -151,8 +151,12 @@ def walk_path( ) ) elif entry.is_file(follow_symlinks=False): - files.append(models.File(name=entry.name, path=str(relative_path.as_posix()))) - return sorted(directories, key=lambda x: x.name), sorted(files, key=lambda x: x.name) + files.append( + models.File(name=entry.name, path=str(relative_path.as_posix())) + ) + return sorted(directories, key=lambda x: x.name), sorted( + files, key=lambda x: x.name + ) directories, files = walk_path(path) relative_path = str(Path(path).relative_to(settings.project_path)) diff --git a/web/server/api/endpoints/lineage.py b/web/server/api/endpoints/lineage.py index 5bef4601fd..008f786438 100644 --- a/web/server/api/endpoints/lineage.py +++ b/web/server/api/endpoints/lineage.py @@ -29,7 +29,9 @@ def get_source_name( ) -> str: table = node.expression.find(exp.Table) if table: - return normalize_model_name(table, default_catalog=default_catalog, dialect=dialect) + return normalize_model_name( + table, default_catalog=default_catalog, dialect=dialect + ) if node.reference_node_name: # CTE name or derived table alias return f"{model_name}: {node.reference_node_name}" @@ -91,7 +93,10 @@ def create_lineage_adjacency_list( if table: column_name = get_column_name(d) dependencies[table].add(column_name) - if isinstance(d.expression, exp.Table) and (table, column_name) not in visited: + if ( + isinstance(d.expression, exp.Table) + and (table, column_name) not in visited + ): nodes.append((table, column_name)) visited.add((table, column_name)) @@ -141,7 +146,9 @@ def column_lineage( try: model_name = context.get_model(model_name).fqn if models_only: - return create_models_only_lineage_adjacency_list(model_name, column_name, context) + return create_models_only_lineage_adjacency_list( + model_name, column_name, context + ) return create_lineage_adjacency_list(model_name, column_name, context) except Exception: raise ApiException( diff --git a/web/server/api/endpoints/meta.py b/web/server/api/endpoints/meta.py index f159f1992a..680a85385e 100644 --- a/web/server/api/endpoints/meta.py +++ b/web/server/api/endpoints/meta.py @@ -23,7 +23,9 @@ def get_api_meta( has_running_task = False if models.Modules.PLANS in settings.modules: - has_running_task = hasattr(request.app.state, "task") and not request.app.state.task.done() + has_running_task = ( + hasattr(request.app.state, "task") and not request.app.state.task.done() + ) api_console.log_event_plan_overview() api_console.log_event_plan_apply() diff --git a/web/server/api/endpoints/models.py b/web/server/api/endpoints/models.py index 21a7b93eb0..40919084de 100644 --- a/web/server/api/endpoints/models.py +++ b/web/server/api/endpoints/models.py @@ -58,16 +58,23 @@ def serialize_all_models( ) -def serialize_model(context: Context, model: Model, render_query: bool = False) -> models.Model: +def serialize_model( + context: Context, model: Model, render_query: bool = False +) -> models.Model: type = _get_model_type(model) default_catalog = model.default_catalog dialect = model.dialect or "SQLGlot" time_column = ( - f"{model.time_column.column} | {model.time_column.format}" if model.time_column else None + f"{model.time_column.column} | {model.time_column.format}" + if model.time_column + else None ) tags = ", ".join(model.tags) if model.tags else None partitioned_by = ( - ", ".join(expr.sql(pretty=True, dialect=model.dialect) for expr in model.partitioned_by) + ", ".join( + expr.sql(pretty=True, dialect=model.dialect) + for expr in model.partitioned_by + ) if model.partitioned_by else None ) @@ -85,9 +92,13 @@ def serialize_model(context: Context, model: Model, render_query: bool = False) description = model.column_descriptions.get(name) if not description and render_query: # The column name is already normalized in `columns_to_types`, so we need to quote it - description = column_description(context, model.name, name, quote_column=True) + description = column_description( + context, model.name, name, quote_column=True + ) - columns.append(models.Column(name=name, type=str(data_type), description=description)) + columns.append( + models.Column(name=name, type=str(data_type), description=description) + ) details = models.ModelDetails( owner=model.owner, @@ -102,7 +113,9 @@ def serialize_model(context: Context, model: Model, render_query: bool = False) time_column=time_column, tags=tags, references=[ - models.Reference(name=ref.name, expression=ref.expression.sql(), unique=ref.unique) + models.Reference( + name=ref.name, expression=ref.expression.sql(), unique=ref.unique + ) for ref in model.all_references ], partitioned_by=partitioned_by, @@ -117,7 +130,9 @@ def serialize_model(context: Context, model: Model, render_query: bool = False) sql = None if render_query: query = model.render_query() or ( - model.query if hasattr(model, "query") else exp.select('"FAILED TO RENDER QUERY"') + model.query + if hasattr(model, "query") + else exp.select('"FAILED TO RENDER QUERY"') ) sql = query.sql(pretty=True, dialect=model.dialect) @@ -125,7 +140,9 @@ def serialize_model(context: Context, model: Model, render_query: bool = False) return models.Model( name=model.name, fqn=model.fqn, - path=str(path.absolute().relative_to(context.path).as_posix()) if path else None, + path=( + str(path.absolute().relative_to(context.path).as_posix()) if path else None + ), full_path=str(path.absolute().as_posix()) if path else None, dialect=dialect, columns=columns, diff --git a/web/server/api/endpoints/plan.py b/web/server/api/endpoints/plan.py index ba4feb34d3..0b4dbca061 100644 --- a/web/server/api/endpoints/plan.py +++ b/web/server/api/endpoints/plan.py @@ -35,7 +35,12 @@ async def initiate_plan( plan_options = plan_options or models.PlanOptions() request.app.state.task = asyncio.create_task( run_in_executor( - get_plan_builder, context, plan_options, environment, plan_dates, categories + get_plan_builder, + context, + plan_options, + environment, + plan_dates, + categories, ) ) else: @@ -83,7 +88,9 @@ def get_plan_builder( categories: t.Optional[t.Dict[str, SnapshotChangeCategory]] = None, ) -> PlanBuilder: try: - return _get_plan_builder(context, plan_options, environment, plan_dates, categories) + return _get_plan_builder( + context, plan_options, environment, plan_dates, categories + ) except ApiException as e: raise e except Exception as e: @@ -133,7 +140,9 @@ def _get_plan_changes(context: Context, plan: Plan) -> models.PlanChanges: def _get_plan_backfills(context: Context, plan: Plan) -> t.Dict[str, t.Any]: """Get plan backfills""" merged_intervals = context.scheduler().merged_missing_intervals() - batches = context.scheduler().batch_intervals(merged_intervals, None, EnvironmentNamingInfo()) + batches = context.scheduler().batch_intervals( + merged_intervals, None, EnvironmentNamingInfo() + ) tasks = {snapshot.name: len(intervals) for snapshot, intervals in batches.items()} snapshots = plan.context_diff.snapshots default_catalog = context.default_catalog @@ -166,7 +175,9 @@ def _get_plan_builder( plan_dates: t.Optional[models.PlanDates] = None, categories: t.Optional[t.Dict[str, SnapshotChangeCategory]] = None, ) -> PlanBuilder: - tracker = models.PlanOverviewStageTracker(environment=environment, plan_options=plan_options) + tracker = models.PlanOverviewStageTracker( + environment=environment, plan_options=plan_options + ) api_console.start_plan_tracker(tracker) tracker_stage_validate = models.PlanStageValidation() tracker.add_stage(stage=models.PlanStage.validation, data=tracker_stage_validate) diff --git a/web/server/api/endpoints/table_diff.py b/web/server/api/endpoints/table_diff.py index b0167ed032..ba82477e0e 100644 --- a/web/server/api/endpoints/table_diff.py +++ b/web/server/api/endpoints/table_diff.py @@ -6,7 +6,8 @@ from sqlglot import exp from sqlmesh.core.context import Context -from web.server.models import ProcessedSampleData, RowDiff, SchemaDiff, TableDiff +from web.server.models import (ProcessedSampleData, RowDiff, SchemaDiff, + TableDiff) from web.server.settings import get_loaded_context router = APIRouter() @@ -14,8 +15,8 @@ def _cells_match(x: t.Any, y: t.Any) -> bool: # lazily import pandas and numpy as we do in core - import pandas as pd import numpy as np + import pandas as pd def _normalize(val: t.Any) -> t.Any: if pd.isnull(val): @@ -33,12 +34,16 @@ def _process_sample_data( if row_diff.joined_sample.shape[0] == 0: return ProcessedSampleData( column_differences=[], - source_only=row_diff.s_sample.replace({pd.NA: None}).to_dict("records") - if row_diff.s_sample.shape[0] > 0 - else [], - target_only=row_diff.t_sample.replace({pd.NA: None}).to_dict("records") - if row_diff.t_sample.shape[0] > 0 - else [], + source_only=( + row_diff.s_sample.replace({pd.NA: None}).to_dict("records") + if row_diff.s_sample.shape[0] > 0 + else [] + ), + target_only=( + row_diff.t_sample.replace({pd.NA: None}).to_dict("records") + if row_diff.t_sample.shape[0] > 0 + else [] + ), ) keys: list[str] = [] @@ -100,12 +105,16 @@ def _process_sample_data( return ProcessedSampleData( column_differences=column_differences, - source_only=row_diff.s_sample.replace({pd.NA: None}).to_dict("records") - if row_diff.s_sample.shape[0] > 0 - else [], - target_only=row_diff.t_sample.replace({pd.NA: None}).to_dict("records") - if row_diff.t_sample.shape[0] > 0 - else [], + source_only=( + row_diff.s_sample.replace({pd.NA: None}).to_dict("records") + if row_diff.s_sample.shape[0] > 0 + else [] + ), + target_only=( + row_diff.t_sample.replace({pd.NA: None}).to_dict("records") + if row_diff.t_sample.shape[0] > 0 + else [] + ), ) diff --git a/web/server/console.py b/web/server/console.py index 871aaefbb1..00f7945334 100644 --- a/web/server/console.py +++ b/web/server/console.py @@ -3,13 +3,16 @@ import asyncio import json import typing as t + from fastapi.encoders import jsonable_encoder from sse_starlette.sse import ServerSentEvent -from sqlmesh.core.snapshot.definition import Interval, Intervals + from sqlmesh.core.console import TerminalConsole from sqlmesh.core.environment import EnvironmentNamingInfo from sqlmesh.core.plan.definition import EvaluatablePlan -from sqlmesh.core.snapshot import Snapshot, SnapshotInfoLike, SnapshotTableInfo, SnapshotId +from sqlmesh.core.snapshot import (Snapshot, SnapshotId, SnapshotInfoLike, + SnapshotTableInfo) +from sqlmesh.core.snapshot.definition import Interval, Intervals from sqlmesh.core.snapshot.execution_tracker import QueryExecutionStats from sqlmesh.core.test import ModelTest from sqlmesh.core.test.result import ModelTextTestResult @@ -69,7 +72,9 @@ def stop_creation_progress(self, success: bool = True) -> None: if self.is_cancelling_plan(): self.finish_plan_cancellation() else: - self.stop_plan_tracker(tracker=self.plan_apply_stage_tracker, success=success) + self.stop_plan_tracker( + tracker=self.plan_apply_stage_tracker, success=success + ) def start_restate_progress(self) -> None: if self.plan_apply_stage_tracker: @@ -87,7 +92,9 @@ def stop_restate_progress(self, success: bool) -> None: if self.is_cancelling_plan(): self.finish_plan_cancellation() else: - self.stop_plan_tracker(tracker=self.plan_apply_stage_tracker, success=success) + self.stop_plan_tracker( + tracker=self.plan_apply_stage_tracker, success=success + ) def start_evaluation_progress( self, @@ -101,7 +108,8 @@ def start_evaluation_progress( if self.plan_apply_stage_tracker: batch_sizes = { - snapshot: len(intervals) for snapshot, intervals in batched_intervals.items() + snapshot: len(intervals) + for snapshot, intervals in batched_intervals.items() } tasks = { snapshot.name: models.BackfillTask( @@ -109,7 +117,9 @@ def start_evaluation_progress( total=total_tasks, start=now_timestamp(), name=snapshot.name, - view_name=snapshot.display_name(environment_naming_info, default_catalog), + view_name=snapshot.display_name( + environment_naming_info, default_catalog + ), ) for snapshot, total_tasks in batch_sizes.items() } @@ -168,7 +178,9 @@ def stop_evaluation_progress(self, success: bool = True) -> None: if self.is_cancelling_plan(): self.finish_plan_cancellation() else: - self.stop_plan_tracker(tracker=self.plan_apply_stage_tracker, success=success) + self.stop_plan_tracker( + tracker=self.plan_apply_stage_tracker, success=success + ) def start_promotion_progress( self, @@ -188,7 +200,9 @@ def start_promotion_progress( self.log_event_plan_apply() - def update_promotion_progress(self, snapshot: SnapshotInfoLike, promoted: bool) -> None: + def update_promotion_progress( + self, snapshot: SnapshotInfoLike, promoted: bool + ) -> None: if self.plan_apply_stage_tracker and self.plan_apply_stage_tracker.promote: self.plan_apply_stage_tracker.promote.update( {"num_tasks": self.plan_apply_stage_tracker.promote.num_tasks + 1} @@ -204,7 +218,9 @@ def stop_promotion_progress(self, success: bool = True) -> None: if self.is_cancelling_plan(): self.finish_plan_cancellation() else: - self.stop_plan_tracker(tracker=self.plan_apply_stage_tracker, success=success) + self.stop_plan_tracker( + tracker=self.plan_apply_stage_tracker, success=success + ) def start_plan_tracker( self, @@ -233,7 +249,10 @@ def stop_plan_tracker( ], success: bool = True, ) -> None: - if isinstance(tracker, models.PlanApplyStageTracker) and self.plan_apply_stage_tracker: + if ( + isinstance(tracker, models.PlanApplyStageTracker) + and self.plan_apply_stage_tracker + ): self.stop_plan_tracker_stages(self.plan_apply_stage_tracker, False) self.plan_apply_stage_tracker.stop(success) self.log_event_plan_apply() @@ -245,7 +264,10 @@ def stop_plan_tracker( self.stop_plan_tracker_stages(self.plan_overview_stage_tracker, False) self.plan_overview_stage_tracker.stop(success) self.log_event_plan_overview() - elif isinstance(tracker, models.PlanCancelStageTracker) and self.plan_cancel_stage_tracker: + elif ( + isinstance(tracker, models.PlanCancelStageTracker) + and self.plan_cancel_stage_tracker + ): self.stop_plan_tracker_stages(self.plan_cancel_stage_tracker, False) self.plan_cancel_stage_tracker.stop(success) self.log_event_plan_cancel() @@ -261,7 +283,9 @@ def log_event( ) ) - def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> None: + def log_test_results( + self, result: ModelTextTestResult, target_dialect: str + ) -> None: if result.wasSuccessful(): self.log_event( event=models.EventName.TESTS, @@ -299,21 +323,31 @@ def log_test_results(self, result: ModelTextTestResult, target_dialect: str) -> def log_event_plan_apply(self) -> None: self.log_event( event=models.EventName.PLAN_APPLY, - data=self.plan_apply_stage_tracker.dict() if self.plan_apply_stage_tracker else {}, + data=( + self.plan_apply_stage_tracker.dict() + if self.plan_apply_stage_tracker + else {} + ), ) def log_event_plan_overview(self) -> None: self.log_event( event=models.EventName.PLAN_OVERVIEW, data=( - self.plan_overview_stage_tracker.dict() if self.plan_overview_stage_tracker else {} + self.plan_overview_stage_tracker.dict() + if self.plan_overview_stage_tracker + else {} ), ) def log_event_plan_cancel(self) -> None: self.log_event( event=models.EventName.PLAN_CANCEL, - data=self.plan_cancel_stage_tracker.dict() if self.plan_cancel_stage_tracker else {}, + data=( + self.plan_cancel_stage_tracker.dict() + if self.plan_cancel_stage_tracker + else {} + ), ) def log_exception(self, exception: ApiException) -> None: @@ -329,7 +363,10 @@ def log_exception(self, exception: ApiException) -> None: self.stop_plan_tracker(self.plan_apply_stage_tracker, False) def is_cancelling_plan(self) -> bool: - return bool(self.plan_cancel_stage_tracker and not self.plan_cancel_stage_tracker.meta.done) + return bool( + self.plan_cancel_stage_tracker + and not self.plan_cancel_stage_tracker.meta.done + ) def stop_plan_tracker_stages( self, @@ -346,12 +383,18 @@ def stop_plan_tracker_stages( return stages = ( - [attr for attr in tracker.__fields__ if not attr.startswith("__")] if tracker else [] + [attr for attr in tracker.__fields__ if not attr.startswith("__")] + if tracker + else [] ) for key in stages: stage = getattr(tracker, key) - if isinstance(stage, models.Trackable) and stage.meta and not stage.meta.done: + if ( + isinstance(stage, models.Trackable) + and stage.meta + and not stage.meta.done + ): stage.stop(success) tracker.stop(success) diff --git a/web/server/models.py b/web/server/models.py index 935822cf46..f3a06f1c0b 100644 --- a/web/server/models.py +++ b/web/server/models.py @@ -13,13 +13,11 @@ from sqlmesh.core.environment import Environment, EnvironmentNamingInfo from sqlmesh.core.node import IntervalUnit, NodeType from sqlmesh.core.plan.definition import Plan -from sqlmesh.core.snapshot.definition import ( - Snapshot, - SnapshotChangeCategory, - SnapshotId, -) +from sqlmesh.core.snapshot.definition import (Snapshot, SnapshotChangeCategory, + SnapshotId) from sqlmesh.utils.date import TimeLike, now_timestamp -from sqlmesh.utils.pydantic import PydanticModel, ValidationInfo, field_validator, validation_data +from sqlmesh.utils.pydantic import (PydanticModel, ValidationInfo, + field_validator, validation_data) SUPPORTED_EXTENSIONS = {".py", ".sql", ".yaml", ".yml", ".csv"} @@ -203,13 +201,17 @@ def get_view_name( default_catalog: t.Optional[str], ) -> str: return ( - snapshots[snapshot_id].display_name(environment_naming_info, default_catalog) + snapshots[snapshot_id].display_name( + environment_naming_info, default_catalog + ) if snapshot_id in snapshots else snapshot_id.name ) @staticmethod - def get_node_type(snapshots: t.Dict[SnapshotId, Snapshot], snapshot_id: SnapshotId) -> NodeType: + def get_node_type( + snapshots: t.Dict[SnapshotId, Snapshot], snapshot_id: SnapshotId + ) -> NodeType: return snapshots[snapshot_id].node_type @@ -248,7 +250,7 @@ def _get_parents( for parent in current.parents: if parent.name not in plan.context_diff.modified_snapshots: continue - (snapshot, _) = plan.context_diff.modified_snapshots[parent.name] + snapshot, _ = plan.context_diff.modified_snapshots[parent.name] parents = ( parents | {parent.name} @@ -265,7 +267,9 @@ def _get_parents( direct.append( ChangeDirect( name=name, - view_name=current.display_name(environment_naming_info, default_catalog), + view_name=current.display_name( + environment_naming_info, default_catalog + ), node_type=current.node_type, diff=plan.context_diff.text_diff(name), change_category=current.change_category, @@ -276,7 +280,9 @@ def _get_parents( indirect.append( ChangeIndirect( name=name, - view_name=current.display_name(environment_naming_info, default_catalog), + view_name=current.display_name( + environment_naming_info, default_catalog + ), node_type=current.node_type, parents=_get_parents(current), ) @@ -285,7 +291,9 @@ def _get_parents( metadata.append( ChangeDisplay( name=name, - view_name=current.display_name(environment_naming_info, default_catalog), + view_name=current.display_name( + environment_naming_info, default_catalog + ), node_type=current.node_type, ) ) @@ -401,7 +409,9 @@ def validate_schema( ) -> t.Dict[str, str]: if isinstance(v, dict): # Handle modified field which has tuples of (source_type, target_type) - if info.field_name == "modified" and any(isinstance(val, tuple) for val in v.values()): + if info.field_name == "modified" and any( + isinstance(val, tuple) for val in v.values() + ): return { k: f"{str(val[0])} → {str(val[1])}" for k, val in v.items() diff --git a/web/server/openapi.py b/web/server/openapi.py index 302688dd53..377a56c673 100644 --- a/web/server/openapi.py +++ b/web/server/openapi.py @@ -15,7 +15,9 @@ def generate_openapi_spec(app: FastAPI, path: str) -> None: parser = argparse.ArgumentParser(description="Generate OpenAPI specification") parser.add_argument( - "--output", default="web/client/openapi.json", help="Path to output OpenAPI spec file" + "--output", + default="web/client/openapi.json", + help="Path to output OpenAPI spec file", ) args = parser.parse_args() diff --git a/web/server/settings.py b/web/server/settings.py index d893f52afa..e19c5da848 100644 --- a/web/server/settings.py +++ b/web/server/settings.py @@ -99,7 +99,9 @@ def get_loaded_context( ) -> t.Generator[Context, None, None]: try: with get_loaded_context_lock: - yield _get_loaded_context(settings.project_path, settings.config, settings.gateway) + yield _get_loaded_context( + settings.project_path, settings.config, settings.gateway + ) except Exception: raise ApiException( message="Unable to create a loaded context", diff --git a/web/server/utils.py b/web/server/utils.py index 868425e76e..80daaaeea5 100644 --- a/web/server/utils.py +++ b/web/server/utils.py @@ -12,10 +12,10 @@ from starlette.status import HTTP_404_NOT_FOUND, HTTP_422_UNPROCESSABLE_ENTITY from sqlmesh.core import constants as c +from sqlmesh.utils.windows import IS_WINDOWS from web.server.console import api_console from web.server.exceptions import ApiException from web.server.settings import Settings, get_context, get_settings -from sqlmesh.utils.windows import IS_WINDOWS if t.TYPE_CHECKING: import pandas as pd @@ -66,7 +66,9 @@ def validate_path(path: str, settings: Settings = Depends(get_settings)) -> str: if any( full_path.match(pattern) for pattern in ( - context.config_for_path(Path(path))[0].ignore_patterns if context else c.IGNORE_PATTERNS + context.config_for_path(Path(path))[0].ignore_patterns + if context + else c.IGNORE_PATTERNS ) ): raise HTTPException(status_code=HTTP_404_NOT_FOUND) diff --git a/web/server/watcher.py b/web/server/watcher.py index 8bc87c8719..2048e848f8 100644 --- a/web/server/watcher.py +++ b/web/server/watcher.py @@ -6,15 +6,12 @@ from sqlmesh.core import constants as c from sqlmesh.core.context import Context from web.server import models -from web.server.api.endpoints.files import _get_directory, _get_file_with_content +from web.server.api.endpoints.files import (_get_directory, + _get_file_with_content) from web.server.console import api_console from web.server.exceptions import ApiException -from web.server.settings import ( - Settings, - get_context, - get_settings, - invalidate_context_cache, -) +from web.server.settings import (Settings, get_context, get_settings, + invalidate_context_cache) from web.server.utils import is_relative_to @@ -30,10 +27,14 @@ async def watch_project() -> None: ] ignore_dirs = [".env"] cache_path = ( - context.cache_dir.resolve() if context else (settings.project_path / c.CACHE).resolve() + context.cache_dir.resolve() + if context + else (settings.project_path / c.CACHE).resolve() ) ignore_paths: t.List[t.Union[str, Path]] = [cache_path] - ignore_entity_patterns = context.config.ignore_patterns if context else c.IGNORE_PATTERNS + ignore_entity_patterns = ( + context.config.ignore_patterns if context else c.IGNORE_PATTERNS + ) ignore_entity_patterns.append("^.*\\.db(\\.wal)?$") async for entries in awatch( @@ -71,7 +72,8 @@ async def watch_project() -> None: change=change, path=str(relative_path), file=_get_file_with_content( - settings.project_path / relative_path, str(relative_path) + settings.project_path / relative_path, + str(relative_path), ), ) )