# Licensed to the Apache Software Foundation (ASF) under one # or more contributor license agreements. See the NOTICE file # distributed with this work for additional information # regarding copyright ownership. The ASF licenses this file # to you under the Apache License, Version 2.0 (the # "License"); you may not use this file except in compliance # with the License. You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, # software distributed under the License is distributed on an # "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. import ctypes import datetime as dt import gc import gzip import pathlib import shutil from dataclasses import fields import pyarrow as pa import pyarrow.compute as pc import pyarrow.dataset as ds import pytest from datafusion import ( Accumulator, CsvReadOptions, DataFrame, RuntimeEnvBuilder, SessionConfig, SessionContext, SessionExtensionComponents, SQLOptions, Table, column, literal, udaf, udf, udwf, ) from datafusion.user_defined import WindowEvaluator def test_create_context_no_args(): SessionContext() def test_create_context_session_config_only(): SessionContext(config=SessionConfig()) def test_create_context_runtime_config_only(): SessionContext(runtime=RuntimeEnvBuilder()) @pytest.mark.parametrize("path_to_str", [True, False]) def test_runtime_configs(tmp_path, path_to_str): path1 = tmp_path / "dir1" path2 = tmp_path / "dir2" path1 = str(path1) if path_to_str else path1 path2 = str(path2) if path_to_str else path2 runtime = RuntimeEnvBuilder().with_disk_manager_specified(path1, path2) config = SessionConfig().with_default_catalog_and_schema("foo", "bar") ctx = SessionContext(config, runtime) assert ctx is not None db = ctx.catalog("foo").schema("bar") assert db is not None @pytest.mark.parametrize("path_to_str", [True, False]) def test_temporary_files(tmp_path, path_to_str): path = str(tmp_path) if path_to_str else tmp_path runtime = RuntimeEnvBuilder().with_temp_file_path(path) config = SessionConfig().with_default_catalog_and_schema("foo", "bar") ctx = SessionContext(config, runtime) assert ctx is not None db = ctx.catalog("foo").schema("bar") assert db is not None def test_create_context_with_all_valid_args(): runtime = RuntimeEnvBuilder().with_disk_manager_os().with_fair_spill_pool(10000000) config = ( SessionConfig() .with_create_default_catalog_and_schema(enabled=True) .with_default_catalog_and_schema("foo", "bar") .with_target_partitions(1) .with_information_schema(enabled=True) .with_repartition_joins(enabled=False) .with_repartition_aggregations(enabled=False) .with_repartition_windows(enabled=False) .with_parquet_pruning(enabled=False) ) ctx = SessionContext(config, runtime) # verify that at least some of the arguments worked ctx.catalog("foo").schema("bar") with pytest.raises(KeyError): ctx.catalog("datafusion") def test_session_config_set_rejects_an_unknown_namespace(): """A bad config key raises rather than aborting through a Rust panic. `datafusion.runtime.*` appears in `information_schema.df_settings` but has no `ConfigOptions` namespace, so it is the key a naive "read the settings back and replay them on the worker" loop hits first. """ # `ValueError`, not a bare `Exception`: a panic would arrive as # `PanicException`, which derives from `BaseException` and so would not be # caught here at all. Both this and the constructor cases below rely on it. with pytest.raises(ValueError, match="runtime"): SessionConfig().set("datafusion.runtime.memory_limit", "unlimited") def test_session_config_set_rejects_an_unparsable_value(): """A well-known key with a value of the wrong type raises too.""" with pytest.raises(ValueError, match="batch_size"): SessionConfig().set("datafusion.execution.batch_size", "not_an_int") def test_session_config_constructor_applies_options(): """A dict passed to the constructor reaches the session's options.""" config = SessionConfig( { "datafusion.execution.batch_size": "1024", "datafusion.execution.target_partitions": "3", } ) ctx = SessionContext(config.with_information_schema(True)) settings = ctx.sql( "select name, value from information_schema.df_settings" " where name in ('datafusion.execution.batch_size'," " 'datafusion.execution.target_partitions')" ).to_pydict() assert dict(zip(settings["name"], settings["value"], strict=True)) == { "datafusion.execution.batch_size": "1024", "datafusion.execution.target_partitions": "3", } def test_session_config_constructor_rejects_an_unknown_namespace(): """A bad key in the constructor's dict raises rather than panicking. The same defect as `SessionConfig.set` had, reached through the argument that a replayed `information_schema.df_settings` dictionary arrives in. """ with pytest.raises(ValueError, match="runtime"): SessionConfig({"datafusion.runtime.memory_limit": "unlimited"}) def test_session_config_constructor_rejects_an_unparsable_value(): """A well-known key with a value of the wrong type raises too.""" with pytest.raises(ValueError, match="batch_size"): SessionConfig({"datafusion.execution.batch_size": "not_an_int"}) def test_register_record_batches(ctx): # create a RecordBatch and register it as memtable batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) ctx.register_record_batches("t", [[batch]]) assert ctx.catalog().schema().names() == {"t"} result = ctx.sql("SELECT a+b, a-b FROM t").collect() assert result[0].column(0) == pa.array([5, 7, 9]) assert result[0].column(1) == pa.array([-3, -3, -3]) def test_register_record_batches_empty(ctx): # A partition list with no record batches carries no schema, so this used to # panic on unchecked `[0][0]` indexing. It should now raise a clear error. with pytest.raises(ValueError, match="no record batches"): ctx.register_record_batches("t", [[]]) # An empty outer partition list carries no schema either, and raises the same error. with pytest.raises(ValueError, match="no record batches"): ctx.register_record_batches("t", []) # The schema is still recovered from a later non-empty partition. batch = pa.RecordBatch.from_arrays([pa.array([1, 2, 3])], names=["a"]) ctx.register_record_batches("t2", [[], [batch]]) assert ctx.sql("SELECT a FROM t2").collect()[0].column(0) == pa.array([1, 2, 3]) def test_create_dataframe_registers_unique_table_name(ctx): # create a RecordBatch and register it as memtable batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) df = ctx.create_dataframe([[batch]]) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert len(tables[0]) == 33 assert tables[0].startswith("c") # ensure that the rest of the table name contains # only hexadecimal numbers for c in tables[0][1:]: assert c in "0123456789abcdef" def test_create_dataframe_registers_with_defined_table_name(ctx): # create a RecordBatch and register it as memtable batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) df = ctx.create_dataframe([[batch]], name="tbl") tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert tables[0] == "tbl" def test_from_arrow_table(ctx): # create a PyArrow table data = {"a": [1, 2, 3], "b": [4, 5, 6]} table = pa.Table.from_pydict(data) # convert to DataFrame df = ctx.from_arrow(table) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert isinstance(df, DataFrame) assert set(df.schema().names) == {"a", "b"} assert df.collect()[0].num_rows == 3 def record_batch_generator(num_batches: int): schema = pa.schema([("a", pa.int64()), ("b", pa.int64())]) for _i in range(num_batches): yield pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], schema=schema ) @pytest.mark.parametrize( "source", [ # __arrow_c_array__ sources pa.array([{"a": 1, "b": 4}, {"a": 2, "b": 5}, {"a": 3, "b": 6}]), # __arrow_c_stream__ sources pa.RecordBatch.from_pydict({"a": [1, 2, 3], "b": [4, 5, 6]}), pa.RecordBatchReader.from_batches( pa.schema([("a", pa.int64()), ("b", pa.int64())]), record_batch_generator(1) ), pa.Table.from_pydict({"a": [1, 2, 3], "b": [4, 5, 6]}), ], ) def test_from_arrow_sources(ctx, source) -> None: df = ctx.from_arrow(source) assert df assert isinstance(df, DataFrame) assert df.schema().names == ["a", "b"] assert df.count() == 3 def test_from_arrow_table_with_name(ctx): # create a PyArrow table data = {"a": [1, 2, 3], "b": [4, 5, 6]} table = pa.Table.from_pydict(data) # convert to DataFrame with optional name df = ctx.from_arrow(table, name="tbl") tables = list(ctx.catalog().schema().names()) assert df assert tables[0] == "tbl" def test_from_arrow_table_empty(ctx): data = {"a": [], "b": []} schema = pa.schema([("a", pa.int32()), ("b", pa.string())]) table = pa.Table.from_pydict(data, schema=schema) # convert to DataFrame df = ctx.from_arrow(table) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert isinstance(df, DataFrame) assert set(df.schema().names) == {"a", "b"} assert len(df.collect()) == 0 def test_from_arrow_table_empty_no_schema(ctx): data = {"a": [], "b": []} table = pa.Table.from_pydict(data) # convert to DataFrame df = ctx.from_arrow(table) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert isinstance(df, DataFrame) assert set(df.schema().names) == {"a", "b"} assert len(df.collect()) == 0 def test_from_pylist(ctx): # create a dataframe from Python list data = [ {"a": 1, "b": 4}, {"a": 2, "b": 5}, {"a": 3, "b": 6}, ] df = ctx.from_pylist(data) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert isinstance(df, DataFrame) assert set(df.schema().names) == {"a", "b"} assert df.collect()[0].num_rows == 3 def test_from_pydict(ctx): # create a dataframe from Python dictionary data = {"a": [1, 2, 3], "b": [4, 5, 6]} df = ctx.from_pydict(data) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert isinstance(df, DataFrame) assert set(df.schema().names) == {"a", "b"} assert df.collect()[0].num_rows == 3 def test_from_pandas(ctx): # create a dataframe from pandas dataframe pd = pytest.importorskip("pandas") data = {"a": [1, 2, 3], "b": [4, 5, 6]} pandas_df = pd.DataFrame(data) df = ctx.from_pandas(pandas_df) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert isinstance(df, DataFrame) assert set(df.schema().names) == {"a", "b"} assert df.collect()[0].num_rows == 3 def test_from_polars(ctx): # create a dataframe from Polars dataframe pd = pytest.importorskip("polars") data = {"a": [1, 2, 3], "b": [4, 5, 6]} polars_df = pd.DataFrame(data) df = ctx.from_polars(polars_df) tables = list(ctx.catalog().schema().names()) assert df assert len(tables) == 1 assert isinstance(df, DataFrame) assert set(df.schema().names) == {"a", "b"} assert df.collect()[0].num_rows == 3 def test_register_table(ctx, database): default = ctx.catalog() public = default.schema("public") assert public.names() == {"csv", "csv1", "csv2"} table = public.table("csv") ctx.register_table("csv3", table) assert public.names() == {"csv", "csv1", "csv2", "csv3"} def test_read_table_from_catalog(ctx, database): default = ctx.catalog() public = default.schema("public") assert public.names() == {"csv", "csv1", "csv2"} table = public.table("csv") table_df = ctx.read_table(table) table_df.show() def test_read_table_from_df(ctx): df = ctx.from_pydict({"a": [1, 2]}) result = ctx.read_table(df).collect() assert [b.to_pydict() for b in result] == [{"a": [1, 2]}] def test_read_table_from_dataset(ctx): batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) dataset = ds.dataset([batch]) result = ctx.read_table(dataset).collect() assert result[0].column(0) == pa.array([1, 2, 3]) assert result[0].column(1) == pa.array([4, 5, 6]) def test_deregister_table(ctx, database): default = ctx.catalog() public = default.schema("public") assert public.names() == {"csv", "csv1", "csv2"} ctx.deregister_table("csv") assert public.names() == {"csv1", "csv2"} def test_deregister_udf(): ctx = SessionContext() is_null = udf( lambda x: x.is_null(), [pa.float64()], pa.bool_(), volatility="immutable", name="my_is_null", ) ctx.register_udf(is_null) # Verify it works df = ctx.from_pydict({"a": [1.0, None]}) ctx.register_table("t", df.into_view()) result = ctx.sql("SELECT my_is_null(a) FROM t").collect() assert result[0].column(0) == pa.array([False, True]) # Deregister and verify it's gone ctx.deregister_udf("my_is_null") with pytest.raises(ValueError): ctx.sql("SELECT my_is_null(a) FROM t").collect() def test_deregister_udaf(): import pyarrow.compute as pc ctx = SessionContext() from datafusion import Accumulator, udaf class MySum(Accumulator): def __init__(self): self._sum = 0.0 def update(self, values: pa.Array) -> None: self._sum += pc.sum(values).as_py() def merge(self, states: list[pa.Array]) -> None: self._sum += pc.sum(states[0]).as_py() def state(self) -> list: return [self._sum] def evaluate(self) -> pa.Scalar: return self._sum my_sum = udaf( MySum, [pa.float64()], pa.float64(), [pa.float64()], volatility="immutable", name="my_sum", ) ctx.register_udaf(my_sum) df = ctx.from_pydict({"a": [1.0, 2.0, 3.0]}) ctx.register_table("t", df.into_view()) result = ctx.sql("SELECT my_sum(a) FROM t").collect() assert result[0].column(0) == pa.array([6.0]) ctx.deregister_udaf("my_sum") with pytest.raises(ValueError): ctx.sql("SELECT my_sum(a) FROM t").collect() def test_deregister_udwf(): ctx = SessionContext() from datafusion import udwf from datafusion.user_defined import WindowEvaluator class MyRowNumber(WindowEvaluator): def __init__(self): self._row = 0 def evaluate_all(self, values, num_rows): return pa.array(list(range(1, num_rows + 1)), type=pa.uint64()) my_row_number = udwf( MyRowNumber, [pa.float64()], pa.uint64(), volatility="immutable", name="my_row_number", ) ctx.register_udwf(my_row_number) df = ctx.from_pydict({"a": [1.0, 2.0, 3.0]}) ctx.register_table("t", df.into_view()) result = ctx.sql("SELECT my_row_number(a) OVER () FROM t").collect() assert result[0].column(0) == pa.array([1, 2, 3], type=pa.uint64()) ctx.deregister_udwf("my_row_number") with pytest.raises(ValueError): ctx.sql("SELECT my_row_number(a) OVER () FROM t").collect() def test_deregister_udtf(): import pyarrow.dataset as ds ctx = SessionContext() from datafusion import Table, udtf class MyTable: def __call__(self): batch = pa.RecordBatch.from_pydict({"x": [1, 2, 3]}) return Table(ds.dataset([batch])) my_table = udtf(MyTable(), "my_table") ctx.register_udtf(my_table) result = ctx.sql("SELECT * FROM my_table()").collect() assert result[0].column(0) == pa.array([1, 2, 3]) ctx.deregister_udtf("my_table") with pytest.raises(ValueError): ctx.sql("SELECT * FROM my_table()").collect() def test_register_table_from_dataframe(ctx): df = ctx.from_pydict({"a": [1, 2]}) ctx.register_table("df_tbl", df) result = ctx.sql("SELECT * FROM df_tbl").collect() assert [b.to_pydict() for b in result] == [{"a": [1, 2]}] @pytest.mark.parametrize("temporary", [True, False]) def test_register_table_from_dataframe_into_view(ctx, temporary): df = ctx.from_pydict({"a": [1, 2]}) table = df.into_view(temporary=temporary) assert isinstance(table, Table) if temporary: assert table.kind == "temporary" else: assert table.kind == "view" ctx.register_table("view_tbl", table) result = ctx.sql("SELECT * FROM view_tbl").collect() assert [b.to_pydict() for b in result] == [{"a": [1, 2]}] def test_table_from_dataframe(ctx): df = ctx.from_pydict({"a": [1, 2]}) table = Table(df) assert isinstance(table, Table) ctx.register_table("from_dataframe_tbl", table) result = ctx.sql("SELECT * FROM from_dataframe_tbl").collect() assert [b.to_pydict() for b in result] == [{"a": [1, 2]}] def test_table_from_dataframe_internal(ctx): df = ctx.from_pydict({"a": [1, 2]}) table = Table(df.df) assert isinstance(table, Table) ctx.register_table("from_internal_dataframe_tbl", table) result = ctx.sql("SELECT * FROM from_internal_dataframe_tbl").collect() assert [b.to_pydict() for b in result] == [{"a": [1, 2]}] def test_register_dataset(ctx): # create a RecordBatch and register it as a pyarrow.dataset.Dataset batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) dataset = ds.dataset([batch]) ctx.register_dataset("t", dataset) assert ctx.catalog().schema().names() == {"t"} result = ctx.sql("SELECT a+b, a-b FROM t").collect() assert result[0].column(0) == pa.array([5, 7, 9]) assert result[0].column(1) == pa.array([-3, -3, -3]) def test_dataset_filter(ctx, capfd): # create a RecordBatch and register it as a pyarrow.dataset.Dataset batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) dataset = ds.dataset([batch]) ctx.register_dataset("t", dataset) assert ctx.catalog().schema().names() == {"t"} df = ctx.sql("SELECT a+b, a-b FROM t WHERE a BETWEEN 2 and 3 AND b > 5") # Make sure the filter was pushed down in Physical Plan df.explain() captured = capfd.readouterr() assert "filter_expr=(((a >= 2) and (a <= 3)) and (b > 5))" in captured.out result = df.collect() assert result[0].column(0) == pa.array([9]) assert result[0].column(1) == pa.array([-3]) def test_dataset_count(ctx): # `datafusion-python` issue: https://github.com/apache/datafusion-python/issues/800 batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) dataset = ds.dataset([batch]) ctx.register_dataset("t", dataset) # Testing the dataframe API df = ctx.table("t") assert df.count() == 3 # Testing the SQL API count = ctx.sql("SELECT COUNT(*) FROM t") count = count.collect() assert count[0].column(0) == pa.array([3]) def test_pyarrow_predicate_pushdown_is_null(ctx, capfd): """Ensure that pyarrow filter gets pushed down for `IsNull`""" # create a RecordBatch and register it as a pyarrow.dataset.Dataset batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6]), pa.array([7, None, 9])], names=["a", "b", "c"], ) dataset = ds.dataset([batch]) ctx.register_dataset("t", dataset) # Make sure the filter was pushed down in Physical Plan df = ctx.sql("SELECT a FROM t WHERE c is NULL") df.explain() captured = capfd.readouterr() assert "filter_expr=is_null(c, {nan_is_null=false})" in captured.out result = df.collect() assert result[0].column(0) == pa.array([2]) def test_pyarrow_predicate_pushdown_timestamp(ctx, tmpdir, capfd): """Ensure that pyarrow filter gets pushed down for timestamp""" # Ref: https://github.com/apache/datafusion-python/issues/703 # create pyarrow dataset with no actual files col_type = pa.timestamp("ns", "+00:00") nyd_2000 = pa.scalar(dt.datetime(2000, 1, 1, tzinfo=dt.timezone.utc), col_type) pa_dataset_fs = pa.fs.SubTreeFileSystem(str(tmpdir), pa.fs.LocalFileSystem()) pa_dataset_format = pa.dataset.ParquetFileFormat() pa_dataset_partition = pa.dataset.field("a") <= nyd_2000 fragments = [ # NOTE: we never actually make this file. # Working predicate pushdown means it never gets accessed pa_dataset_format.make_fragment( "1.parquet", filesystem=pa_dataset_fs, partition_expression=pa_dataset_partition, ) ] pa_dataset = pa.dataset.FileSystemDataset( fragments, pa.schema([pa.field("a", col_type)]), pa_dataset_format, pa_dataset_fs, ) ctx.register_dataset("t", pa_dataset) # the partition for our only fragment is for a < 2000-01-01. # so querying for a > 2024-01-01 should not touch any files df = ctx.sql("SELECT * FROM t WHERE a > '2024-01-01T00:00:00+00:00'") assert df.collect() == [] def test_dataset_filter_nested_data(ctx): # create Arrow StructArrays to test nested data types data = pa.StructArray.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) batch = pa.RecordBatch.from_arrays( [data], names=["nested_data"], ) dataset = ds.dataset([batch]) ctx.register_dataset("t", dataset) assert ctx.catalog().schema().names() == {"t"} df = ctx.table("t") # This filter will not be pushed down to DatasetExec since it # isn't supported df = df.filter(column("nested_data")["b"] > literal(5)).select( column("nested_data")["a"] + column("nested_data")["b"], column("nested_data")["a"] - column("nested_data")["b"], ) result = df.collect() assert result[0].column(0) == pa.array([9]) assert result[0].column(1) == pa.array([-3]) def test_table_exist(ctx): batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) dataset = ds.dataset([batch]) ctx.register_dataset("t", dataset) assert ctx.table_exist("t") is True def test_table_not_found(ctx): from uuid import uuid4 with pytest.raises(KeyError): ctx.table(f"not-found-{uuid4()}") def test_session_start_time(ctx): import datetime import re st = ctx.session_start_time() assert isinstance(st, str) # Truncate nanoseconds to microseconds for Python 3.10 compat st = re.sub(r"(\.\d{6})\d+", r"\1", st) dt = datetime.datetime.fromisoformat(st) assert dt.isoformat() def test_enable_ident_normalization(ctx): assert ctx.enable_ident_normalization() is True ctx.sql("SET datafusion.sql_parser.enable_ident_normalization = false") assert ctx.enable_ident_normalization() is False def test_parse_sql_expr(ctx): from datafusion.common import DFSchema schema = DFSchema.empty() expr = ctx.parse_sql_expr("1 + 2", schema) assert str(expr) == "Expr(Int64(1) + Int64(2))" def test_execute_logical_plan(ctx): df = ctx.from_pydict({"a": [1, 2, 3]}) plan = df.logical_plan() df2 = ctx.execute_logical_plan(plan) result = df2.collect() assert result[0].column(0) == pa.array([1, 2, 3]) def test_refresh_catalogs(ctx): ctx.refresh_catalogs() def test_remove_optimizer_rule(ctx): assert ctx.remove_optimizer_rule("push_down_filter") is True assert ctx.remove_optimizer_rule("nonexistent_rule") is False def test_set_query_planner_rejects_wrong_capsule(ctx): with pytest.raises(ValueError, match="datafusion_query_planner"): ctx.set_query_planner(ctx.__datafusion_task_context_provider__()) def test_with_extension_rejects_wrong_capsule(ctx): """The extension options hook names the capsule it was handed. Like the rest of the capsule family, this reports which capsule turned up rather than CPython's fixed "called with incorrect name" string. """ class WrongCapsule: def __datafusion_extension_options__(self): return ctx.__datafusion_task_context_provider__() with pytest.raises(ValueError, match="datafusion_extension_options"): SessionConfig().with_extension(WrongCapsule()) def test_pre_55_codec_signature_reports_an_upgrade(ctx): """A getter that refuses the session is named, not left as a bare TypeError. Extension libraries implement these getters, so the pre-55.0.0 signature is what an out-of-date one still has. The original error stays reachable as ``__cause__`` rather than being replaced outright. """ class PreSessionCodec: def __datafusion_logical_extension_codec__(self): msg = "should never be called" raise AssertionError(msg) with pytest.raises(ImportError, match="__datafusion_logical_extension_codec__"): ctx.with_logical_extension_codec(PreSessionCodec()) with pytest.raises(ImportError) as excinfo: ctx.with_logical_extension_codec(PreSessionCodec()) assert isinstance(excinfo.value.__cause__, TypeError) assert "positional argument" in str(excinfo.value.__cause__) def test_type_error_inside_a_getter_is_not_reported_as_an_upgrade(ctx): """A correctly-signed getter's own TypeError must survive unchanged. Only the call machinery's arity error means the library is out of date. Rewriting every TypeError would send an author debugging their own getter off to upgrade a library that is already correct. """ class RaisesTypeError: def __datafusion_logical_extension_codec__(self, session): msg = "bad cast inside the getter" raise TypeError(msg) with pytest.raises(TypeError, match="bad cast inside the getter"): ctx.with_logical_extension_codec(RaisesTypeError()) def test_non_type_errors_from_a_getter_propagate(ctx): """Anything that is not a TypeError was never a signature problem.""" class RaisesValueError: def __datafusion_logical_extension_codec__(self, session): msg = "something else entirely" raise ValueError(msg) with pytest.raises(ValueError, match="something else entirely"): ctx.with_logical_extension_codec(RaisesValueError()) def test_set_query_planner_capsule(ctx): capsule = ctx.__datafusion_query_planner__() get_name = ctypes.pythonapi.PyCapsule_GetName get_name.argtypes = [ctypes.py_object] get_name.restype = ctypes.c_char_p assert get_name(capsule) == b"datafusion_query_planner" ctx.register_record_batches( "query_planner_test", [[pa.RecordBatch.from_pydict({"value": [1, 2, 3]})]], ) ctx.set_query_planner(capsule) assert ctx.table_exist("query_planner_test") batches = ctx.sql("SELECT 1 AS value").collect() assert batches[0].column(0) == pa.array([1]) def test_installing_a_planner_leaves_the_session_intact(ctx): """The planner is written into the existing session, not a copy of it. Registrations made before the install are still visible afterwards, and ones made after are visible too -- there is a single session throughout, so neither the catalogs nor the function registry are snapshotted. """ ctx.register_record_batches( "registered_before", [[pa.RecordBatch.from_pydict({"value": [1]})]], ) before = udf( lambda arr: arr, [pa.int64()], pa.int64(), volatility="immutable", name="registered_before", ) ctx.register_udf(before) ctx.set_query_planner(ctx.__datafusion_query_planner__()) ctx.register_record_batches( "registered_after", [[pa.RecordBatch.from_pydict({"value": [2]})]], ) after = udf( lambda arr: arr, [pa.int64()], pa.int64(), volatility="immutable", name="registered_after", ) ctx.register_udf(after) assert ctx.table_exist("registered_before") assert ctx.table_exist("registered_after") assert ctx.sql("SELECT registered_before(1)").collect() assert ctx.sql("SELECT registered_after(1)").collect() def test_contexts_sharing_a_session_share_the_planner(ctx): """A context derived before the install still plans through the planner. ``with_python_udf_inlining`` returns a handle on the same session, and the query planner lives in that session's state. """ sibling = ctx.with_python_udf_inlining(enabled=False) ctx.register_record_batches( "shared_planner_test", [[pa.RecordBatch.from_pydict({"value": [1, 2, 3]})]], ) ctx.set_query_planner(ctx.__datafusion_query_planner__()) assert sibling.table_exist("shared_planner_test") assert sibling.session_id() == ctx.session_id() class _NamedCodec: """Wraps a codec capsule in an object that can name itself. ``with_extensions`` requires objects rather than bare capsules, because a codec's wire id is read off the object it is handed over as. This is the shape a library holding a raw capsule hands over. """ def __init__(self, capsule, codec_id): self._capsule = capsule self.__datafusion_codec_id__ = codec_id def __datafusion_logical_extension_codec__(self, session=None): return self._capsule def __datafusion_physical_extension_codec__(self, session=None): return self._capsule class _CodecOnlyExtension: """Contributes decline-all codecs exported from an unrelated session. Retaining ``ctx`` is what the protocol tells real extensions not to do — a bundle is reusable, so a cached context belongs to whichever session it was last installed on. It is kept here only so a test can assert *which* context the factory was handed. """ def __init__(self, prefix="my_library"): self.exporter = SessionContext() self.prefix = prefix self.bound_ctx = None def __datafusion_session_components__(self, ctx): self.bound_ctx = ctx return SessionExtensionComponents( logical_extension_codecs=( _NamedCodec( self.exporter.__datafusion_logical_extension_codec__(), f"{self.prefix}.logical", ), ), physical_extension_codecs=( _NamedCodec( self.exporter.__datafusion_physical_extension_codec__(), f"{self.prefix}.physical", ), ), ) class _PlannerExtension: """Contributes a planner, recording the fallback it was handed. Passing ``fallback`` straight back through is the degenerate wrap: it plans the same queries to the same plans, which is what lets a pure-Python test assert the threading without a real layering planner. It is not a no-op — the capsule gets installed, so the session ends up planning through a foreign planner — but nothing here depends on that either way. """ def __init__(self, calls=None): self.fallbacks = [] self.planner_ctx = None # Shared list the hooks append themselves to, so a test can assert the # order they ran in rather than only that each ran. self.calls = [] if calls is None else calls def __datafusion_session_planner__(self, ctx, fallback): self.calls.append(self) self.planner_ctx = ctx self.fallbacks.append(fallback) return fallback def test_with_extensions_accepts_no_extensions(ctx): """No extensions installs nothing and returns a handle on this session. A caller assembling the list at runtime — from a plugin registry, say — should not have to special-case it being empty, and every sibling varargs method on ``DataFrame`` accepts zero arguments the same way. """ ctx.register_record_batches( "empty_extensions_test", [[pa.RecordBatch.from_pydict({"value": [1]})]], ) result = ctx.with_extensions() assert result.session_id() == ctx.session_id() assert result.table_exist("empty_extensions_test") assert result.logical_extension_codec_ids() == ctx.logical_extension_codec_ids() def test_with_extensions_no_extensions_keeps_an_installed_planner(ctx): """The empty case must not disturb a planner the session already has. A call that installs no codec and no planner skips the planner commit entirely, so the installed planner keeps the chains it was bound to and the session still plans through it. """ extension = _CodecOnlyExtension() installed = ctx.with_extensions(extension, _PlannerExtension()) result = installed.with_extensions() assert result.logical_extension_codec_ids() == ( installed.logical_extension_codec_ids() ) batches = result.sql("SELECT 1 AS value").collect() assert batches[0].column(0) == pa.array([1]) def test_with_extensions_rejects_non_extension(ctx): with pytest.raises(TypeError, match="__datafusion_session_planner__"): ctx.with_extensions(object()) def test_with_extensions_rejects_bad_components(ctx): class BadExtension: def __datafusion_session_components__(self, ctx): return 42 with pytest.raises(TypeError, match="SessionExtensionComponents"): ctx.with_extensions(BadExtension()) @pytest.mark.parametrize( "field", ["logical_extension_codecs", "physical_extension_codecs"] ) def test_session_extension_components_rejects_a_single_codec(field): """A lone codec is not an iterable of codecs. Dropping the trailing comma is the easy way to write one by accident. The check lives on the value type so the error lands in the extension library's own frame, naming the field it got wrong, rather than surfacing later inside ``with_extensions`` as ``'_NamedCodec' object is not iterable``. """ codec = _NamedCodec( SessionContext().__datafusion_logical_extension_codec__(), "my_library.logical", ) with pytest.raises(TypeError, match=r"must be an iterable of codec objects"): SessionExtensionComponents(**{field: codec}) def test_session_extension_components_rejects_a_string(): """A str is iterable, so it needs refusing on its own. Left alone it would normalize into a tuple of characters and fail much later as that many bogus codecs. """ with pytest.raises(TypeError, match=r"not a single str"): SessionExtensionComponents(logical_extension_codecs="my_library.logical") def test_with_extensions_accepts_a_planner_only_extension(ctx): """An extension may implement the planner hook alone. A library that ships an optimizing planner and no codecs — nothing to contribute in phase one — should not have to return empty components. """ extension = _PlannerExtension() result = ctx.with_extensions(extension) assert len(extension.fallbacks) == 1 assert result.session_id() == ctx.session_id() def test_with_extensions_threads_the_planner_through_in_order(ctx): """Each planner hook receives what the previous one returned. Planners nest rather than chain, so the host hands each bundle the planner built so far. Argument order is nesting order, last one outermost. """ calls = [] first, second = _PlannerExtension(calls), _PlannerExtension(calls) ctx.with_extensions(first, second) # Argument order, once each. Nothing else pins the order: both hooks # return capsules, and a host that ran them backwards would still leave # each with one fallback recorded. assert calls == [first, second] # `first` returned its fallback unchanged, but the host re-exports every # hook's return value before handing it on, so `second` receives a capsule # of its own rather than the object `first` was handed. assert second.fallbacks[0] is not first.fallbacks[0] # That is as far as pure Python reaches: a capsule is opaque, so this # cannot tell a re-export of `first`'s planner from a fresh read of the # session's. `test_with_extensions_nests_planners_in_argument_order` in # examples/datafusion-ffi-query-planner-example is what pins the nesting, # by asserting the outer planner delegated to the inner one. def test_with_extensions_planner_hook_sees_the_new_handle(ctx): """Phase two runs against the handle carrying the final codec chains. A planner captured against the pre-install handle would encode through a chain missing every codec this call installed. """ codecs = _CodecOnlyExtension() planner = _PlannerExtension() result = ctx.with_extensions(codecs, planner) assert planner.planner_ctx.logical_extension_codec_ids() == ["my_library.logical"] assert result.logical_extension_codec_ids() == ["my_library.logical"] def test_with_extensions_skips_a_planner_hook_returning_none(ctx): """Returning ``None`` contributes no planner and keeps the fallback. The skip is what lets the call succeed at all: a host that treated the ``None`` as a contribution would hand it to the export step and fail with ``'None' is not an instance of 'PyCapsule'`` before ``downstream`` ran. """ class NoPlanner: def __init__(self): self.fallbacks = [] def __datafusion_session_planner__(self, ctx, fallback): self.fallbacks.append(fallback) # Spelled out rather than left to fall off the end: `None` is the # protocol's "contribute no planner", which is what this test is # about, and an implicit one would read as an oversight. return None # noqa: RET501, PLR1711 skipped = NoPlanner() downstream = _PlannerExtension() result = ctx.with_extensions(skipped, downstream) # Both hooks ran, and `downstream` was handed a planner rather than the # `None` in front of it. It is a fresh read of the session's planner, not # the object `skipped` was given, so a host that fell back by reusing the # previous hook's *input* is ruled out too. assert len(skipped.fallbacks) == 1 assert len(downstream.fallbacks) == 1 assert downstream.fallbacks[0] is not skipped.fallbacks[0] batches = result.sql("SELECT 1 AS value").collect() assert batches[0].column(0) == pa.array([1]) def test_with_extensions_rejects_bad_codec_capsule(ctx): """A correctly shaped object still has to return the right capsule.""" class BadCodecExtension: def __datafusion_session_components__(self, ctx): wrong_capsule = ctx.__datafusion_task_context_provider__() return SessionExtensionComponents( logical_extension_codecs=( _NamedCodec(wrong_capsule, "my_library.logical"), ), ) with pytest.raises( ValueError, match="Expected name 'datafusion_logical_extension_codec'" ): ctx.with_extensions(BadCodecExtension()) def test_with_extensions_rejects_a_bare_capsule_codec(ctx): """A codec must be an object that can name itself, not a bare capsule. An id is read off the object a codec is handed over as, and a capsule has no type to read one from. ``with_extensions`` takes no ``codec_id=``, so the capsule is refused here rather than given an id derived from something that is not the codec. """ class BareCapsuleExtension: def __init__(self): self.exporter = SessionContext() def __datafusion_session_components__(self, ctx): return SessionExtensionComponents( logical_extension_codecs=( self.exporter.__datafusion_logical_extension_codec__(), ), ) with pytest.raises( TypeError, match="must be an object exposing `__datafusion_logical_extension_codec__`", ): ctx.with_extensions(BareCapsuleExtension()) def test_with_extensions_rejects_a_bare_physical_capsule_codec(ctx): """The physical getter is named in its own diagnostic.""" class BareCapsuleExtension: def __init__(self): self.exporter = SessionContext() def __datafusion_session_components__(self, ctx): return SessionExtensionComponents( physical_extension_codecs=( self.exporter.__datafusion_physical_extension_codec__(), ), ) with pytest.raises( TypeError, match="must be an object exposing `__datafusion_physical_extension_codec__`", ): ctx.with_extensions(BareCapsuleExtension()) def test_with_extensions_codec_ids_survive_composition(ctx): """A codec keeps its id when its extension is nested inside another one. Wire ids have to mean the same thing in whichever process decodes, so packaging one extension inside another — the natural way for an application to present several libraries as one — must not re-tag the inner library's payloads. Reading the id off the handed-over object rather than off the contributing extension is what guarantees that. """ class ComposedExtension: """Presents another extension's components as its own.""" def __init__(self, inner): self.inner = inner def __datafusion_session_components__(self, ctx): return self.inner.__datafusion_session_components__(ctx) direct = ctx.with_extensions(_CodecOnlyExtension()) wrapped = SessionContext().with_extensions(ComposedExtension(_CodecOnlyExtension())) assert direct.logical_extension_codec_ids() == ["my_library.logical"] assert wrapped.logical_extension_codec_ids() == ["my_library.logical"] assert direct.physical_extension_codec_ids() == ["my_library.physical"] assert wrapped.physical_extension_codec_ids() == ["my_library.physical"] def test_with_extensions_uses_ids_declared_on_the_codec(ctx): """``__datafusion_codec_id__`` on the handed-over object names the codec. This is how an extension contributing more than one codec of a kind tells them apart. """ class TwoNamedCodecs: def __init__(self): self.exporter = SessionContext() def __datafusion_session_components__(self, ctx): return SessionExtensionComponents( logical_extension_codecs=( _NamedCodec( self.exporter.__datafusion_logical_extension_codec__(), "my_library.first", ), _NamedCodec( self.exporter.__datafusion_logical_extension_codec__(), "my_library.second", ), ), ) result = ctx.with_extensions(TwoNamedCodecs()) assert result.logical_extension_codec_ids() == [ "my_library.first", "my_library.second", ] def test_with_extensions_rejects_two_codecs_of_one_class(ctx): """Two wrappers of one class claim one class-derived id, so they collide. Numbering them by position would be an id another library can mint the same value from, and would break stored plans the first time the extension reordered what it returns, so the ambiguity is refused instead. """ class UnnamedCodec: def __init__(self, capsule): self._capsule = capsule def __datafusion_logical_extension_codec__(self, session=None): return self._capsule class TwoUnnamedCodecs: def __init__(self): self.exporter = SessionContext() def __datafusion_session_components__(self, ctx): capsule = self.exporter.__datafusion_logical_extension_codec__ return SessionExtensionComponents( logical_extension_codecs=( UnnamedCodec(capsule()), UnnamedCodec(capsule()), ), ) with pytest.raises(ValueError, match="__datafusion_codec_id__"): ctx.with_extensions(TwoUnnamedCodecs()) def test_with_extensions_leaves_an_exporting_object_its_own_id(ctx): """A codec handed over as an object keeps the identity it declares.""" exporter = SessionContext() class ObjectCodecExtension: def __datafusion_session_components__(self, ctx): return SessionExtensionComponents(logical_extension_codecs=(exporter,)) result = ctx.with_extensions(ObjectCodecExtension()) assert result.logical_extension_codec_ids() == [exporter.__datafusion_codec_id__] def test_with_extensions_installs_codecs_and_planner(ctx): ctx.register_record_batches( "extensions_test", [[pa.RecordBatch.from_pydict({"value": [1, 2, 3]})]], ) extension = _CodecOnlyExtension() result = ctx.with_extensions(extension, _PlannerExtension()) assert result.table_exist("extensions_test") # In-memory tables need a real extension codec to round-trip through the # FFI planner, so query plans that don't serialize a table provider. batches = result.sql("SELECT 1 AS value").collect() assert batches[0].column(0) == pa.array([1]) def test_with_extensions_binds_to_the_receiving_session(ctx): extension = _CodecOnlyExtension() result = ctx.with_extensions(extension) # Factories are handed the receiver itself, so a component bound during # installation targets the session the returned handle also wraps. There # is no intermediate context that could be collected out from under it. assert extension.bound_ctx is ctx assert result.session_id() == ctx.session_id() # One session: a registration through either handle is visible to both. ctx.register_record_batches( "bound_test", [[pa.RecordBatch.from_pydict({"value": [1]})]], ) assert result.table_exist("bound_test") def test_with_extensions_survives_source_collection(): extension = _CodecOnlyExtension() result = SessionContext().with_extensions(extension, _PlannerExtension()) gc.collect() batches = result.sql("SELECT 1 AS value").collect() assert batches[0].column(0) == pa.array([1]) def test_with_extensions_failure_leaves_source_usable(ctx): class BoomExtension: def __datafusion_session_components__(self, ctx): msg = "boom" raise RuntimeError(msg) with pytest.raises(RuntimeError, match="boom"): ctx.with_extensions(_CodecOnlyExtension(), BoomExtension()) batches = ctx.sql("SELECT 1 AS value").collect() assert batches[0].column(0) == pa.array([1]) def _doubler(name="double"): """A scalar function under a name the caller picks.""" return udf( lambda arr: pa.array([v.as_py() * 2 for v in arr]), [pa.int64()], pa.int64(), volatility="stable", name=name, ) class _Total(Accumulator): """The smallest accumulator that survives a partial/final split.""" def __init__(self): self._sum = 0 def state(self) -> list[pa.Scalar]: return [pa.scalar(self._sum)] def update(self, values: pa.Array) -> None: self._sum += pc.sum(values).as_py() or 0 def merge(self, states: list[pa.Array]) -> None: self._sum += pc.sum(states[0]).as_py() or 0 def evaluate(self) -> pa.Scalar: return pa.scalar(self._sum) class _First(WindowEvaluator): """Repeats the first value of the partition across every row.""" def evaluate_all(self, values: list[pa.Array], num_rows: int) -> pa.Array: first = values[0][0].as_py() return pa.array([first] * num_rows) def _total(name="total"): """An aggregate function under a name the caller picks.""" return udaf( _Total, pa.int64(), pa.int64(), [pa.int64()], volatility="stable", name=name, ) def _first(name="first_value_of"): """A window function under a name the caller picks.""" return udwf( _First, pa.int64(), pa.int64(), volatility="immutable", name=name, ) class _FunctionExtension: """Contributes functions and nothing else. The shape a library shipping only functions has: no codecs, no planner, so the whole of its installation is what it declares here. """ def __init__(self, udfs=(), udafs=(), udwfs=()): self._udfs = udfs self._udafs = udafs self._udwfs = udwfs def __datafusion_session_components__(self, ctx): return SessionExtensionComponents( udfs=self._udfs, udafs=self._udafs, udwfs=self._udwfs ) def test_with_extensions_registers_a_declared_udf(ctx): """A declared scalar function is callable from SQL on the returned handle.""" result = ctx.with_extensions(_FunctionExtension(udfs=(_doubler(),))) result.from_pydict({"a": [1, 2, 3]}, name="nums") batches = result.sql("SELECT double(a) AS doubled FROM nums").collect() assert batches[0].column(0) == pa.array([2, 4, 6]) def test_with_extensions_registers_udafs_and_udwfs(ctx): """The other two function kinds install the same way.""" result = ctx.with_extensions( _FunctionExtension(udafs=(_total(),), udwfs=(_first(),)) ) result.from_pydict({"a": [1, 2, 3]}, name="nums") assert result.sql("SELECT total(a) FROM nums").collect()[0].column(0) == pa.array( [6] ) batches = result.sql("SELECT first_value_of(a) OVER () FROM nums").collect() assert batches[0].column(0) == pa.array([1, 1, 1]) def test_with_extensions_registers_on_the_shared_session(ctx): """Registrations land on the session, which the source context also holds. Only the codec chains belong to the returned handle. Pinned deliberately: a future change that made registrations private to the handle would be a behaviour change, not a fix. """ ctx.with_extensions(_FunctionExtension(udfs=(_doubler(),))) assert ctx.udf("double").name == "double" @pytest.mark.parametrize( ("field", "make", "label", "lookup"), [ ("udfs", _doubler, "scalar function", "udf"), ("udafs", _total, "aggregate function", "udaf"), ("udwfs", _first, "window function", "udwf"), ], ) def test_with_extensions_rejects_a_name_two_extensions_claim( ctx, field, make, label, lookup ): """Registrations have no fall-through, so a clash cannot be resolved by order. Unlike codecs, which dispatch by id, a second function under one name would silently replace the first. Parametrized over the kinds to pin each ``_FUNCTION_KINDS`` row's field and label wiring, not just the machinery. A codec-carrying bundle rides along to pin the other half of the transaction: resolution runs *after* the codec chains are built, so this failure lands between the two steps, and the chains must not reach the session either. """ name = make().name with pytest.raises(ValueError, match=rf"{label} named '{name}'"): ctx.with_extensions( _CodecOnlyExtension(), _FunctionExtension(**{field: (make(),)}), _FunctionExtension(**{field: (make(),)}), ) with pytest.raises(KeyError): getattr(ctx, lookup)(name) assert ctx.logical_extension_codec_ids() == [] assert ctx.physical_extension_codec_ids() == [] def test_with_extensions_rejects_one_extension_passed_twice(ctx): """A bundle object listed twice is the caller's duplicate, not a naming bug. The remedy has to match the mistake, and nothing the bundle author renames helps here — both claims come from the one declaration. Collisions are therefore keyed on argument position rather than on object identity, which would read a repeat as a bundle colliding with itself and offer a rename that cannot be made. """ extension = _FunctionExtension(udfs=(_doubler(),)) with pytest.raises(ValueError, match=r"argument 0 .* and argument 1 ") as excinfo: ctx.with_extensions(extension, extension) assert "rename" not in str(excinfo.value) with pytest.raises(KeyError): ctx.udf("double") def test_with_extensions_rejects_a_name_one_extension_claims_twice(ctx): """A bundle colliding with itself is its own bug, not a clash of libraries. Separated from the two-extension case because the remedy differs: a bundle author can rename their own function, and a caller cannot rename someone else's. """ with pytest.raises( ValueError, match=r"declares two scalar functions named 'double'" ): ctx.with_extensions(_FunctionExtension(udfs=(_doubler(), _doubler()))) with pytest.raises(KeyError): ctx.udf("double") def test_with_extensions_allows_shadowing_an_existing_function(ctx): """Claiming a name the session already has is legal. ``ctx.udfs()`` holds every built-in, and ``enable_spark_functions`` overrides built-ins by design, so refusing this would refuse a supported use rather than catch a mistake. """ result = ctx.with_extensions(_FunctionExtension(udfs=(_doubler(name="abs"),))) result.from_pydict({"a": [1, 2, 3]}, name="nums") batches = result.sql("SELECT abs(a) AS shadowed FROM nums").collect() assert batches[0].column(0) == pa.array([2, 4, 6]) def test_with_extensions_registers_nothing_when_a_components_hook_raises(ctx): """A failure in phase one leaves the first extension's functions uninstalled.""" class BoomExtension: def __datafusion_session_components__(self, ctx): msg = "boom" raise RuntimeError(msg) with pytest.raises(RuntimeError, match="boom"): ctx.with_extensions( _FunctionExtension(udfs=(_doubler(),)), BoomExtension(), ) with pytest.raises(KeyError): ctx.udf("double") def test_with_extensions_registers_nothing_when_a_planner_hook_raises(ctx): """A failure in phase two does too, which is what pins the ordering. By the time the planner hooks run, the functions have been resolved and their names checked. Committing them at that point rather than after would pass every other test here and still leave this one registered. """ class BoomPlanner: def __datafusion_session_planner__(self, ctx, fallback): msg = "boom" raise RuntimeError(msg) with pytest.raises(RuntimeError, match="boom"): ctx.with_extensions( _FunctionExtension(udfs=(_doubler(),)), BoomPlanner(), ) with pytest.raises(KeyError): ctx.udf("double") def test_with_extensions_rejects_an_unusable_declaration(ctx): """Something that is neither a wrapper nor an exportable names both sides.""" with pytest.raises(TypeError, match=r"__datafusion_scalar_udf__"): ctx.with_extensions(_FunctionExtension(udfs=(object(),))) @pytest.mark.parametrize("field", ["udfs", "udafs", "udwfs"]) def test_session_extension_components_rejects_a_single_function(field): """A lone function is not an iterable of them, as for codecs.""" with pytest.raises(TypeError, match=r"must be an iterable of function objects"): SessionExtensionComponents(**{field: _doubler()}) def test_session_extension_components_rejects_a_single_optimizer_rule(): """The same for rules, naming what that field holds.""" with pytest.raises( TypeError, match=r"must be an iterable of optimizer rule objects" ): SessionExtensionComponents(physical_optimizer_rules=object()) def test_with_extensions_rejects_a_rule_that_is_not_a_rule(ctx): """A declaration that is not a rule at all names the bundle that made it. Which of several bundles is at fault is the whole content of the message, and the only place it is still known is here. A failure also has to leave the session alone even though the extension ahead of it declared a function that was perfectly good. """ class RuleExtension: def __datafusion_session_components__(self, ctx): return SessionExtensionComponents(physical_optimizer_rules=(object(),)) with pytest.raises(TypeError, match=r"got .* from .*RuleExtension"): ctx.with_extensions( _FunctionExtension(udfs=(_doubler(),)), RuleExtension(), ) with pytest.raises(KeyError): ctx.udf("double") def test_with_extensions_rejects_a_rule_whose_getter_returns_a_non_capsule(ctx): """A rule shaped right but returning junk is refused by the importer. The bundle is past the point where it can be named — it declared the right shape — so this is the one rule failure that surfaces as a ``RuntimeError`` from the import rather than a ``TypeError`` from the resolve. """ class NotACapsule: def __datafusion_physical_optimizer_rule__(self): return object() class RuleExtension: def __datafusion_session_components__(self, ctx): return SessionExtensionComponents(physical_optimizer_rules=(NotACapsule(),)) with pytest.raises(RuntimeError, match="datafusion_physical_optimizer_rule"): ctx.with_extensions(_FunctionExtension(udfs=(_doubler(),)), RuleExtension()) with pytest.raises(KeyError): ctx.udf("double") def test_with_extensions_declaring_no_rules_leaves_the_session_id(ctx): """A call declaring no rules leaves the session id alone. This is the control for the no-op path: with nothing to install the state rebuild is skipped, so the id is untouched rather than carried over. The carry-over itself is not reachable from here — the rebuild needs a real rule capsule, which only a compiled extension can hand over. That half is pinned by ``test_rules_install_without_changing_the_session_id`` in ``datafusion-ffi-example``, where a fresh id would leave ``session_id()`` disagreeing with every ``TaskContext`` the session has handed out. """ before = ctx.session_id() result = ctx.with_extensions(_FunctionExtension(udfs=(_doubler(),))) assert result.session_id() == before assert ctx.session_id() == before def test_every_component_field_has_an_installer(): """A field added to the components dataclass must be wired into the install. ``SessionExtensionComponents`` normalizes any field carrying the component metadata, so one added without an installer would be accepted from a bundle and then quietly dropped — the failure this pins is a contributed component going nowhere, with no error to say so. Reaching into private names on purpose: the two sides answer different questions. The metadata says which fields are collections to normalize; ``_FUNCTION_KINDS``, the codec pair, and the rules say which of them ``with_extensions`` knows how to install. Nothing observable from outside can tell you they have drifted, because the symptom is silence. """ from datafusion.context import _FUNCTION_KINDS by_noun: dict[str, set[str]] = {} for spec in fields(SessionExtensionComponents): noun = spec.metadata.get("datafusion_component") if noun is not None: by_noun.setdefault(noun, set()).add(spec.name) assert by_noun == { "codec": {"logical_extension_codecs", "physical_extension_codecs"}, "function": {kind.field for kind in _FUNCTION_KINDS}, "optimizer rule": {"physical_optimizer_rules"}, } def test_every_function_kind_names_something_real(): """The names on a ``_FunctionKind`` row resolve to what it says they do. They are held as strings and looked up during ``with_extensions``, to keep the ``user_defined`` import out of this module's import cycle. The cost is that a typo in a row surfaces as an ``AttributeError`` part-way through an install rather than at import. Every row that exists today is covered by a behaviour test above; this is what covers the next one, which may be added before its own test is. """ from datafusion import user_defined from datafusion.context import _FUNCTION_KINDS for kind in _FUNCTION_KINDS: assert isinstance(getattr(user_defined, kind.wrapper), type) assert callable(getattr(user_defined, kind.factory)) assert kind.field in {spec.name for spec in fields(SessionExtensionComponents)} def test_table_provider(ctx): batch = pa.RecordBatch.from_pydict({"x": [10, 20, 30]}) ctx.register_record_batches("provider_test", [[batch]]) tbl = ctx.table_provider("provider_test") assert tbl.schema == pa.schema([("x", pa.int64())]) def test_table_provider_not_found(ctx): with pytest.raises(KeyError): ctx.table_provider("nonexistent_table") def test_read_json(ctx): path = pathlib.Path(__file__).parent.resolve() # Default test_data_path = path / "data_test_context" / "data.json" df = ctx.read_json(test_data_path) result = df.collect() assert result[0].column(0) == pa.array(["a", "b", "c"]) assert result[0].column(1) == pa.array([1, 2, 3]) # Schema schema = pa.schema( [ pa.field("A", pa.string(), nullable=True), ] ) df = ctx.read_json(test_data_path, schema=schema) result = df.collect() assert result[0].column(0) == pa.array(["a", "b", "c"]) assert result[0].schema == schema # File extension test_data_path = path / "data_test_context" / "data.json" df = ctx.read_json(test_data_path, file_extension=".json") result = df.collect() assert result[0].column(0) == pa.array(["a", "b", "c"]) assert result[0].column(1) == pa.array([1, 2, 3]) def test_read_json_compressed(ctx, tmp_path): path = pathlib.Path(__file__).parent.resolve() test_data_path = path / "data_test_context" / "data.json" # File compression type gzip_path = tmp_path / "data.json.gz" with ( pathlib.Path.open(test_data_path, "rb") as csv_file, gzip.open(gzip_path, "wb") as gzipped_file, ): gzipped_file.writelines(csv_file) df = ctx.read_json(gzip_path, file_extension=".gz", file_compression_type="gz") result = df.collect() assert result[0].column(0) == pa.array(["a", "b", "c"]) assert result[0].column(1) == pa.array([1, 2, 3]) def test_read_csv(ctx): csv_df = ctx.read_csv(path="testing/data/csv/aggregate_test_100.csv") csv_df.select(column("c1")).show() def test_read_csv_list(ctx, tmp_path): source_path = pathlib.Path("testing/data/csv/aggregate_test_100.csv") copied_path = tmp_path / source_path.name shutil.copy(source_path, copied_path) csv_df = ctx.read_csv(path=[source_path]) expected = csv_df.count() * 2 double_csv_df = ctx.read_csv(path=[source_path, copied_path]) actual = double_csv_df.count() double_csv_df.select(column("c1")).show() assert actual == expected def test_read_csv_compressed(ctx, tmp_path): test_data_path = pathlib.Path("testing/data/csv/aggregate_test_100.csv") expected = ctx.read_csv(test_data_path).collect() # File compression type gzip_path = tmp_path / "aggregate_test_100.csv.gz" with ( pathlib.Path.open(test_data_path, "rb") as csv_file, gzip.open(gzip_path, "wb") as gzipped_file, ): gzipped_file.writelines(csv_file) csv_df = ctx.read_csv(gzip_path, file_extension=".gz", file_compression_type="gz") assert csv_df.collect() == expected csv_df = ctx.read_csv( gzip_path, options=CsvReadOptions(file_extension=".gz", file_compression_type="gz"), ) assert csv_df.collect() == expected def test_read_parquet(ctx): parquet_df = ctx.read_parquet(path="parquet/data/alltypes_plain.parquet") parquet_df.show() assert parquet_df is not None path = pathlib.Path.cwd() / "parquet/data/alltypes_plain.parquet" parquet_df = ctx.read_parquet(path=path) assert parquet_df is not None def test_read_avro(ctx): avro_df = ctx.read_avro(path="testing/data/avro/alltypes_plain.avro") avro_df.show() assert avro_df is not None path = pathlib.Path.cwd() / "testing/data/avro/alltypes_plain.avro" avro_df = ctx.read_avro(path=path) assert avro_df is not None def test_read_arrow(ctx, tmp_path): # Write an Arrow IPC file, then read it back table = pa.table({"a": [1, 2, 3], "b": ["x", "y", "z"]}) arrow_path = tmp_path / "test.arrow" with pa.ipc.new_file(str(arrow_path), table.schema) as writer: writer.write_table(table) df = ctx.read_arrow(str(arrow_path)) result = df.collect() assert result[0].column(0) == pa.array([1, 2, 3]) assert result[0].column(1) == pa.array(["x", "y", "z"]) # Also verify pathlib.Path works df = ctx.read_arrow(arrow_path) result = df.collect() assert result[0].column(0) == pa.array([1, 2, 3]) def test_read_empty(ctx): df = ctx.read_empty() result = df.collect() assert len(result) == 1 assert result[0].num_columns == 0 df = ctx.empty_table() result = df.collect() assert len(result) == 1 assert result[0].num_columns == 0 def test_register_arrow(ctx, tmp_path): # Write an Arrow IPC file, then register and query it table = pa.table({"x": [10, 20, 30]}) arrow_path = tmp_path / "test.arrow" with pa.ipc.new_file(str(arrow_path), table.schema) as writer: writer.write_table(table) ctx.register_arrow("arrow_tbl", str(arrow_path)) result = ctx.sql("SELECT * FROM arrow_tbl").collect() assert result[0].column(0) == pa.array([10, 20, 30]) # Also verify pathlib.Path works ctx.register_arrow("arrow_tbl_path", arrow_path) result = ctx.sql("SELECT * FROM arrow_tbl_path").collect() assert result[0].column(0) == pa.array([10, 20, 30]) def test_register_batch(ctx): batch = pa.RecordBatch.from_pydict({"a": [1, 2, 3], "b": [4, 5, 6]}) ctx.register_batch("batch_tbl", batch) result = ctx.sql("SELECT * FROM batch_tbl").collect() assert result[0].column(0) == pa.array([1, 2, 3]) assert result[0].column(1) == pa.array([4, 5, 6]) def test_register_batch_empty(ctx): batch = pa.RecordBatch.from_pydict({"a": pa.array([], type=pa.int64())}) ctx.register_batch("empty_batch_tbl", batch) result = ctx.sql("SELECT * FROM empty_batch_tbl").collect() assert result[0].num_rows == 0 def test_read_batch_returns_dataframe(ctx): batch = pa.RecordBatch.from_pydict({"a": [1, 2, 3], "b": [4, 5, 6]}) df = ctx.read_batch(batch) assert df.to_pydict() == {"a": [1, 2, 3], "b": [4, 5, 6]} # read_batch should not register a named table. assert ctx.catalog().schema().names() == set() def test_read_batches_concatenates(ctx): b1 = pa.RecordBatch.from_pydict({"a": [1, 2]}) b2 = pa.RecordBatch.from_pydict({"a": [3, 4]}) df = ctx.read_batches([b1, b2]) assert df.to_pydict() == {"a": [1, 2, 3, 4]} def test_read_batches_accepts_iterable(ctx): b1 = pa.RecordBatch.from_pydict({"a": [1, 2]}) b2 = pa.RecordBatch.from_pydict({"a": [3, 4]}) # Generator: ensures non-list iterables are materialized before FFI. df = ctx.read_batches(b for b in (b1, b2)) assert df.to_pydict() == {"a": [1, 2, 3, 4]} # Tuple: same. df = ctx.read_batches((b1, b2)) assert df.to_pydict() == {"a": [1, 2, 3, 4]} def test_create_sql_options(): SQLOptions() def test_sql_with_options_no_ddl(ctx): sql = "CREATE TABLE IF NOT EXISTS valuetable AS VALUES(1,'HELLO'),(12,'DATAFUSION')" ctx.sql(sql) options = SQLOptions().with_allow_ddl(allow=False) with pytest.raises(Exception, match="DDL"): ctx.sql_with_options(sql, options=options) def test_sql_with_options_no_dml(ctx): table_name = "t" batch = pa.RecordBatch.from_arrays( [pa.array([1, 2, 3]), pa.array([4, 5, 6])], names=["a", "b"], ) dataset = ds.dataset([batch]) ctx.register_dataset(table_name, dataset) sql = f'INSERT INTO "{table_name}" VALUES (1, 2), (2, 3);' ctx.sql(sql) options = SQLOptions().with_allow_dml(allow=False) with pytest.raises(Exception, match="DML"): ctx.sql_with_options(sql, options=options) def test_sql_with_options_no_statements(ctx): sql = "SET time zone = 1;" ctx.sql(sql) options = SQLOptions().with_allow_statements(allow=False) with pytest.raises(Exception, match="SetVariable"): ctx.sql_with_options(sql, options=options) @pytest.fixture def batch(): return pa.RecordBatch.from_arrays( [pa.array([4, 5, 6])], names=["a"], ) def test_create_dataframe_with_global_ctx(batch): ctx = SessionContext.global_ctx() df = ctx.create_dataframe([[batch]]) result = df.collect()[0].column(0) assert result == pa.array([4, 5, 6]) def test_csv_read_options_builder_pattern(): """Test CsvReadOptions builder pattern.""" from datafusion import CsvReadOptions options = ( CsvReadOptions() .with_has_header(False) .with_delimiter("|") .with_quote("'") .with_schema_infer_max_records(2000) .with_truncated_rows(True) .with_newlines_in_values(True) .with_file_extension(".tsv") ) assert options.has_header is False assert options.delimiter == "|" assert options.quote == "'" assert options.schema_infer_max_records == 2000 assert options.truncated_rows is True assert options.newlines_in_values is True assert options.file_extension == ".tsv" def read_csv_with_options_inner( tmp_path: pathlib.Path, csv_content: str, options: CsvReadOptions, expected: pa.RecordBatch, as_read: bool, global_ctx: bool, ) -> None: from datafusion import SessionContext # Create a test CSV file group_dir = tmp_path / "group=a" group_dir.mkdir(exist_ok=True) csv_path = group_dir / "test.csv" csv_path.write_text(csv_content, newline="\n") ctx = SessionContext() if as_read: if global_ctx: from datafusion.io import read_csv df = read_csv(str(tmp_path), options=options) else: df = ctx.read_csv(str(tmp_path), options=options) else: ctx.register_csv("test_table", str(tmp_path), options=options) df = ctx.sql("SELECT * FROM test_table") df.show() # Verify the data result = df.collect() assert len(result) == 1 assert result[0] == expected @pytest.mark.parametrize( ("as_read", "global_ctx"), [ (True, True), (True, False), (False, False), ], ) def test_read_csv_with_options(tmp_path, as_read, global_ctx): """Test reading CSV with CsvReadOptions.""" csv_content = "Alice;30;|New York; NY|\nBob;25\n#Charlie;35;Paris\nPhil;75;Detroit' MI\nKarin;50;|Stockholm\nSweden|" # noqa: E501 # Some of the read options are difficult to test in combination # such as schema and schema_infer_max_records so run multiple tests # file_sort_order doesn't impact reading, but included here to ensure # all options parse correctly options = CsvReadOptions( has_header=False, delimiter=";", quote="|", terminator="\n", escape="\\", comment="#", newlines_in_values=True, schema_infer_max_records=1, null_regex="[pP]+aris", truncated_rows=True, file_sort_order=[[column("column_1").sort(), column("column_2")], ["column_3"]], ) expected = pa.RecordBatch.from_arrays( [ pa.array(["Alice", "Bob", "Phil", "Karin"]), pa.array([30, 25, 75, 50]), pa.array(["New York; NY", None, "Detroit' MI", "Stockholm\nSweden"]), ], names=["column_1", "column_2", "column_3"], ) read_csv_with_options_inner( tmp_path, csv_content, options, expected, as_read, global_ctx ) schema = pa.schema( [ pa.field("name", pa.string(), nullable=False), pa.field("age", pa.float32(), nullable=False), pa.field("location", pa.string(), nullable=True), ] ) options.with_schema(schema) expected = pa.RecordBatch.from_arrays( [ pa.array(["Alice", "Bob", "Phil", "Karin"]), pa.array([30.0, 25.0, 75.0, 50.0]), pa.array(["New York; NY", None, "Detroit' MI", "Stockholm\nSweden"]), ], schema=schema, ) read_csv_with_options_inner( tmp_path, csv_content, options, expected, as_read, global_ctx ) csv_content = "name,age\nAlice,30\nBob,25\nCharlie,35\nDiego,40\nEmily,15" expected = pa.RecordBatch.from_arrays( [ pa.array(["Alice", "Bob", "Charlie", "Diego", "Emily"]), pa.array([30, 25, 35, 40, 15]), pa.array(["a", "a", "a", "a", "a"]), ], schema=pa.schema( [ pa.field("name", pa.string(), nullable=True), pa.field("age", pa.int64(), nullable=True), pa.field("group", pa.string(), nullable=False), ] ), ) options = CsvReadOptions( table_partition_cols=[("group", pa.string())], ) read_csv_with_options_inner( tmp_path, csv_content, options, expected, as_read, global_ctx ) def test_pre_52_table_provider_signature_reports_an_upgrade(ctx): """The table provider hook reports an upgrade the same way codecs do. The 52.0.0 signature change added the session argument. This path had its own copy of the error mapping and so missed later corrections to it. """ class PreSessionProvider: def __datafusion_table_provider__(self): msg = "should never be called" raise AssertionError(msg) with pytest.raises(ImportError, match="__datafusion_table_provider__") as excinfo: ctx.register_table("old_sig", PreSessionProvider()) assert isinstance(excinfo.value.__cause__, TypeError) def test_catalog_provider_getter_arity_error_names_the_codec(ctx): """A getter refusing the codec capsule is diagnosed, not left as a TypeError. `__datafusion_catalog_provider__` is handed the host's logical extension codec, not the session. It used to call `getattr(...).call1(...)` directly and so produced a bare `TypeError` for an out-of-date library; it now routes through `call_capsule_getter` like the rest of the family. The message has to name the codec rather than the SessionContext, or it would send the author to change the wrong parameter. """ class PreCodecCatalogProvider: def __datafusion_catalog_provider__(self): msg = "should never be called" raise AssertionError(msg) with pytest.raises(ImportError, match="__datafusion_catalog_provider__") as excinfo: ctx.register_catalog_provider("old_sig", PreCodecCatalogProvider()) assert "logical extension codec" in str(excinfo.value) assert "SessionContext" not in str(excinfo.value) assert isinstance(excinfo.value.__cause__, TypeError) def test_type_error_inside_a_catalog_provider_getter_propagates(ctx): """A correctly-signed catalog getter's own TypeError survives unchanged.""" class RaisesTypeError: def __datafusion_catalog_provider__(self, codec): msg = "bad cast inside the getter" raise TypeError(msg) with pytest.raises(TypeError, match="bad cast inside the getter"): ctx.register_catalog_provider("raises", RaisesTypeError()) def test_type_error_inside_a_table_provider_getter_propagates(ctx): """A correctly-signed provider getter's own TypeError survives unchanged.""" class RaisesTypeError: def __datafusion_table_provider__(self, session): msg = "bad cast inside the getter" raise TypeError(msg) with pytest.raises(TypeError, match="bad cast inside the getter"): ctx.register_table("raises", RaisesTypeError())