fix(python): accept null extension column batches

This commit is contained in:
Gatefixer
2026-08-26 22:55:54 +00:00
parent 79f626b09e
commit 2426f275cd
3 changed files with 78 additions and 0 deletions
+47
View File
@@ -452,6 +452,50 @@ def _field_extension_name(field: pa.Field) -> Optional[str]:
return extension_name
def _prepare_extension_list(data: DATA, target_schema: pa.Schema) -> DATA:
"""Give inferred list columns the logical type required by extensions."""
if not isinstance(data, list) or not data or not isinstance(data[0], dict):
return data
target_fields = {field.name: field for field in target_schema}
extension_names = {
name: _field_extension_name(field) for name, field in target_fields.items()
}
if not any(
name in {"arrow.json", "lance.json", "lance.blob.v2"}
for name in extension_names.values()
):
return data
inferred = pa.Table.from_pylist(data)
fields = []
changed = False
for field in inferred.schema:
extension_name = extension_names.get(field.name)
if extension_name in {"arrow.json", "lance.json"}:
json_factory = getattr(pa, "json_", None)
if json_factory is not None:
field = pa.field(field.name, json_factory(), nullable=field.nullable)
else:
field = pa.field(
field.name,
pa.string(),
nullable=field.nullable,
metadata={b"ARROW:extension:name": b"arrow.json"},
)
changed = True
elif extension_name == "lance.blob.v2" and pa.types.is_null(field.type):
field = pa.field(field.name, pa.large_binary(), nullable=field.nullable)
changed = True
fields.append(field)
if not changed:
return inferred
insert_schema = pa.schema(fields, metadata=inferred.schema.metadata)
return pa.Table.from_pylist(data, schema=insert_schema)
def _align_field_types(
fields: List[pa.Field],
target_fields: List[pa.Field],
@@ -5442,6 +5486,9 @@ class AsyncTable:
if fill_value is None:
fill_value = 0.0
if mode != "overwrite":
data = _prepare_extension_list(data, schema)
# _santitize_data is an old code path, but we will use it until the
# new code path is ready.
if mode == "overwrite":
+9
View File
@@ -203,6 +203,15 @@ def test_fetch_blobs_preserves_null_and_empty_values():
assert blobs[3].as_py() == b"present"
def test_add_all_null_list_to_blob_column():
table = _blob_table("all_null_add", [{"id": 1, "image": None}])
hits = table.search().to_arrow()
blobs = table.fetch_blobs("image", hits)
assert len(blobs) == 1
assert blobs[0].as_py() is None
def test_fetch_blob_ranges_aligns_repeated_ranges_and_nulls():
table = _blob_table(
"range_alignment",
+22
View File
@@ -770,6 +770,28 @@ async def test_add_async(mem_db_async: AsyncConnection):
assert await table.count_rows() == 3
@pytest.mark.skipif(not hasattr(pa, "json_"), reason="requires PyArrow JSON type")
@pytest.mark.asyncio
@pytest.mark.parametrize(
("values", "expected"),
[
([None], [None]),
([None, '{"k": 1}'], [None, '{"k":1}']),
(['{"k": 2}'], ['{"k":2}']),
],
)
async def test_add_list_of_dicts_to_json_column(
mem_db_async: AsyncConnection, values, expected
):
schema = pa.schema([pa.field("id", pa.int64()), pa.field("value", pa.json_())])
table = await mem_db_async.create_table("json_list_add", schema=schema)
await table.add([{"id": idx, "value": value} for idx, value in enumerate(values)])
rows = (await table.to_arrow()).sort_by("id").to_pylist()
assert [row["value"] for row in rows] == expected
def test_add_overwrite_infers_vector_schema(mem_db: DBConnection):
"""Overwrite should infer vector columns the same way create_table does.