diff --git a/sqlmesh/core/engine_adapter/duckdb.py b/sqlmesh/core/engine_adapter/duckdb.py index 3666a801dd..6cf2c1eb50 100644 --- a/sqlmesh/core/engine_adapter/duckdb.py +++ b/sqlmesh/core/engine_adapter/duckdb.py @@ -173,9 +173,30 @@ def _create_table( track_rows_processed: bool = True, **kwargs: t.Any, ) -> None: + table_name = ( + table_name_or_schema.this + if isinstance(table_name_or_schema, exp.Schema) + else table_name_or_schema + ) + catalog = exp.to_table(table_name).catalog or self.get_current_catalog() + partitioned_by_exps = None - if self._get_catalog_type(self.get_current_catalog()) == "ducklake": + insert_expression = None + if self._get_catalog_type(catalog) == "ducklake": partitioned_by_exps = kwargs.pop("partitioned_by", None) + if ( + partitioned_by_exps + and expression is not None + and (replace or not exists or not self.table_exists(table_name)) + ): + insert_expression = expression.copy() + query = t.cast(exp.Query, expression) + expression = ( + exp.select("*") + .from_(query.subquery("_sqlmesh_schema_only", copy=False)) + .where(exp.false()) + .limit(0) + ) super()._create_table( table_name_or_schema, @@ -191,12 +212,6 @@ def _create_table( ) if partitioned_by_exps: - # Schema object contains column definitions, so we extract Table - table_name = ( - table_name_or_schema.this - if isinstance(table_name_or_schema, exp.Schema) - else table_name_or_schema - ) table_name_str = ( table_name.sql(dialect=self.dialect) if isinstance(table_name, exp.Table) @@ -207,6 +222,12 @@ def _create_table( ) self.execute(f"ALTER TABLE {table_name_str} SET PARTITIONED BY ({partitioned_by_str});") + if insert_expression is not None: + self.execute( + exp.insert(insert_expression, exp.to_table(table_name)), + track_rows_processed=track_rows_processed, + ) + def _drop_object( self, name: TableName | SchemaName, diff --git a/tests/core/engine_adapter/test_duckdb.py b/tests/core/engine_adapter/test_duckdb.py index 696689219b..5eb8e3d54b 100644 --- a/tests/core/engine_adapter/test_duckdb.py +++ b/tests/core/engine_adapter/test_duckdb.py @@ -1,4 +1,5 @@ import typing as t +from datetime import date import pandas as pd # noqa: TID253 import pytest @@ -226,3 +227,31 @@ def test_drop_object_cascade_by_catalog_type(make_mocked_engine_adapter: t.Calla 'DROP SCHEMA IF EXISTS "lake"."virt" CASCADE', 'DROP TABLE IF EXISTS "native"."phys"."t" CASCADE', ] + + +def test_ducklake_partitioning_on_initial_query(adapter: EngineAdapter, duck_conn, tmp_path): + catalog = "ducklake_initial_partition_db" + + duck_conn.install_extension("ducklake") + duck_conn.load_extension("ducklake") + duck_conn.execute( + f"ATTACH 'ducklake:{tmp_path}/{catalog}.ducklake' AS {catalog} " + f"(DATA_PATH '{tmp_path}', DATA_INLINING_ROW_LIMIT 0);" + ) + + adapter.create_schema(f"{catalog}.test_schema") + adapter.replace_query( + f"{catalog}.test_schema.test_table", + parse_one("SELECT 1 AS id, DATE '2000-01-01' AS ds UNION ALL SELECT 2, DATE '2000-01-02'"), + partitioned_by=[exp.to_column("ds")], + ) + + assert adapter.fetchall(f"SELECT * FROM {catalog}.test_schema.test_table ORDER BY id") == [ + (1, date(2000, 1, 1)), + (2, date(2000, 1, 2)), + ] + partition_ids = duck_conn.execute( + f"SELECT partition_id FROM __ducklake_metadata_{catalog}.main.ducklake_data_file" + ).fetchall() + assert partition_ids + assert all(partition_id is not None for (partition_id,) in partition_ids)