Skip to content
Merged
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
11 changes: 8 additions & 3 deletions mplang/backends/table_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -609,7 +609,12 @@ def _unwrap(val: TableValue | pa.Table | pd.DataFrame) -> pa.Table:

@table.run_sql_p.def_impl
def run_sql_impl(interpreter: Interpreter, op: Operation, *args: Any) -> TableValue:
"""Execute SQL query on input tables."""
"""Execute SQL query on input tables.

Query relations from a DuckDB state other than the selected execution state
are materialized as Arrow tables before registration. This may use
substantial memory for large intermediate relations.
"""
query = op.attrs["query"]
dialect = op.attrs.get("dialect", "duckdb")
table_names = op.attrs["table_names"]
Expand All @@ -626,8 +631,6 @@ def run_sql_impl(interpreter: Interpreter, op: Operation, *args: Any) -> TableVa
if isinstance(data, QueryTableSource):
if state is None:
state = data.state
elif state != data.state:
raise ValueError("All tables must belong to the same DuckDB connection")

if state is None:
conn = duckdb.connect()
Expand All @@ -638,6 +641,8 @@ def run_sql_impl(interpreter: Interpreter, op: Operation, *args: Any) -> TableVa
# register tables or create view
for name, tbl in zip(table_names, tables, strict=True):
data = tbl.unwrap()
if isinstance(data, QueryTableSource) and data.state is not state:
data = tbl.data
if name in state.tables:
if state.tables[name] is not data:
# TODO: rename and rewrite sql??
Expand Down
179 changes: 178 additions & 1 deletion tests/backends/test_table_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,12 +15,20 @@
"""Tests for Table Runtime Implementation."""

import os
from types import SimpleNamespace

import duckdb
import pyarrow as pa
import pytest

import mplang.edsl.typing as elt
from mplang.backends.table_impl import ParquetReader, TableValue
from mplang.backends.table_impl import (
DuckDBState,
ParquetReader,
QueryTableSource,
TableValue,
run_sql_impl,
)
from mplang.dialects import table


Expand Down Expand Up @@ -81,6 +89,175 @@ def workload():
os.remove(path)


def test_run_sql_materializes_query_source_from_different_state():
"""Foreign DuckDB relations are materialized into the selected state."""
local_conn = duckdb.connect()
foreign_conn = duckdb.connect()
try:
local_state = DuckDBState(local_conn)
foreign_state = DuckDBState(foreign_conn)
local = TableValue(
QueryTableSource(
local_conn.sql("SELECT 1 AS id, 'local' AS local_value"),
local_state,
)
)
foreign = TableValue(
QueryTableSource(
foreign_conn.sql("SELECT 1 AS id, 'foreign' AS foreign_value"),
foreign_state,
)
)
op = SimpleNamespace(
attrs={
"query": (
"SELECT local_t.id, local_value, foreign_value "
"FROM local_t JOIN foreign_t USING (id)"
),
"dialect": "duckdb",
"table_names": ["local_t", "foreign_t"],
}
)

result = run_sql_impl(None, op, local, foreign)

assert result.data == pa.table({
"id": pa.array([1], type=pa.int32()),
"local_value": ["local"],
"foreign_value": ["foreign"],
})
assert isinstance(local.unwrap(), QueryTableSource)
assert isinstance(foreign.unwrap(), pa.Table)
assert local_state.tables["foreign_t"] is foreign.data
finally:
local_conn.close()
foreign_conn.close()


def test_run_sql_keeps_same_state_query_sources_lazy():
"""Relations from the execution state keep the lazy view path."""
conn = duckdb.connect()
try:
state = DuckDBState(conn)
left = TableValue(
QueryTableSource(conn.sql("SELECT 1 AS id, 'left' AS left_value"), state)
)
right = TableValue(
QueryTableSource(conn.sql("SELECT 1 AS id, 'right' AS right_value"), state)
)
op = SimpleNamespace(
attrs={
"query": (
"SELECT left_t.id, left_value, right_value "
"FROM left_t JOIN right_t USING (id)"
),
"dialect": "duckdb",
"table_names": ["left_t", "right_t"],
}
)

result = run_sql_impl(None, op, left, right)

assert result.data == pa.table({
"id": pa.array([1], type=pa.int32()),
"left_value": ["left"],
"right_value": ["right"],
})
assert isinstance(left.unwrap(), QueryTableSource)
assert isinstance(right.unwrap(), QueryTableSource)
assert state.tables["left_t"] is left.unwrap()
assert state.tables["right_t"] is right.unwrap()
finally:
conn.close()


def test_run_sql_reuses_materialized_foreign_table_with_same_name():
"""The registered Arrow object preserves same-name identity validation."""
local_conn = duckdb.connect()
foreign_conn = duckdb.connect()
try:
local_state = DuckDBState(local_conn)
foreign_state = DuckDBState(foreign_conn)
local = TableValue(
QueryTableSource(
local_conn.sql("SELECT 1 AS id, 'local' AS local_value"),
local_state,
)
)
foreign = TableValue(
QueryTableSource(
foreign_conn.sql("SELECT 1 AS id, 'foreign' AS foreign_value"),
foreign_state,
)
)
op = SimpleNamespace(
attrs={
"query": "SELECT * FROM local_t JOIN foreign_t USING (id)",
"dialect": "duckdb",
"table_names": ["local_t", "foreign_t"],
}
)

run_sql_impl(None, op, local, foreign)
result = run_sql_impl(None, op, local, foreign)

assert result.data.to_pylist() == [
{
"id": 1,
"local_value": "local",
"foreign_value": "foreign",
}
]
finally:
local_conn.close()
foreign_conn.close()


def test_run_sql_rejects_different_table_with_registered_name():
"""A different object cannot replace a registered foreign table name."""
local_conn = duckdb.connect()
foreign_conn = duckdb.connect()
try:
local_state = DuckDBState(local_conn)
foreign_state = DuckDBState(foreign_conn)
local = TableValue(
QueryTableSource(
local_conn.sql("SELECT 1 AS id, 'local' AS local_value"),
local_state,
)
)
foreign = TableValue(
QueryTableSource(
foreign_conn.sql("SELECT 1 AS id, 'foreign' AS foreign_value"),
foreign_state,
)
)
op = SimpleNamespace(
attrs={
"query": "SELECT * FROM local_t JOIN foreign_t USING (id)",
"dialect": "duckdb",
"table_names": ["local_t", "foreign_t"],
}
)

run_sql_impl(None, op, local, foreign)
replacement = TableValue(
pa.table({
"id": pa.array([1], type=pa.int32()),
"foreign_value": ["replacement"],
})
)

with pytest.raises(RuntimeError) as error:
run_sql_impl(None, op, local, replacement)

assert isinstance(error.value.__cause__, ValueError)
assert str(error.value.__cause__) == "foreign_t has been registered."
finally:
local_conn.close()
foreign_conn.close()


def test_table_constant_dataframe():
"""Test creating constant table from DataFrame."""
import pandas as pd
Expand Down
Loading