diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index c354a944e..54d91d6cd 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -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": diff --git a/python/python/tests/test_blob.py b/python/python/tests/test_blob.py index 5d7682f24..d2b1bf053 100644 --- a/python/python/tests/test_blob.py +++ b/python/python/tests/test_blob.py @@ -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", diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index 56e0eacfd..25b7323b7 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -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.