Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 28 additions & 7 deletions sqlmesh/core/engine_adapter/duckdb.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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)
Expand All @@ -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,
Expand Down
29 changes: 29 additions & 0 deletions tests/core/engine_adapter/test_duckdb.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import typing as t
from datetime import date

import pandas as pd # noqa: TID253
import pytest
Expand Down Expand Up @@ -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)