mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-03 12:08:52 +00:00
fix(python): accept null extension column batches
This commit is contained in:
@@ -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":
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user