diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index 4ad239d13..78a9e12a2 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -489,6 +489,38 @@ def _with_field_type( ) +def _with_list_value_field( + data_type: pa.DataType, value_field: pa.Field +) -> pa.DataType: + if pa.types.is_list(data_type): + return pa.list_(value_field) + if pa.types.is_large_list(data_type): + return pa.large_list(value_field) + return pa.list_(value_field, data_type.list_size) + + +def _extension_storage_field(field: pa.Field) -> pa.Field: + """Return a from-pylist-compatible field for nested write extensions.""" + extension_name = _field_extension_name(field) + if extension_name in _JSON_EXTENSION_NAMES: + metadata = dict(field.metadata or {}) + metadata[b"ARROW:extension:name"] = b"arrow.json" + return _with_field_type(field, pa.string(), metadata=metadata) + if extension_name == _BLOB_EXTENSION_NAME: + metadata = dict(field.metadata or {}) + metadata[b"ARROW:extension:name"] = _BLOB_EXTENSION_NAME.encode() + metadata[b"ARROW:extension:metadata"] = b"" + storage_type = getattr(field.type, "storage_type", field.type) + return _with_field_type(field, storage_type, metadata=metadata) + if pa.types.is_struct(field.type): + children = [_extension_storage_field(child) for child in field.type] + return _with_field_type(field, pa.struct(children)) + if _is_list_like(field.type): + value_field = _extension_storage_field(field.type.value_field) + return _with_field_type(field, _with_list_value_field(field.type, value_field)) + return field + + def _prepare_extension_field( field: pa.Field, target_field: pa.Field ) -> Tuple[pa.Field, bool]: @@ -516,24 +548,52 @@ def _prepare_extension_field( return _with_field_type(field, pa.struct(children)), True if _is_list_like(field.type) and _is_list_like(target_field.type): - prepared, changed = _prepare_extension_field( - field.type.value_field, target_field.type.value_field - ) - if changed: - prepared = _with_field_type( - prepared, prepared.type, name=target_field.type.value_field.name - ) - if pa.types.is_list(field.type): - data_type = pa.list_(prepared) - elif pa.types.is_large_list(field.type): - data_type = pa.large_list(prepared) - else: - data_type = pa.list_(prepared, field.type.list_size) + target_value_field = target_field.type.value_field + if _field_contains_write_extension(target_value_field): + prepared = _extension_storage_field(target_value_field) + data_type = _with_list_value_field(target_field.type, prepared) return _with_field_type(field, data_type), True return field, False +def _prepare_extension_value( + value: Any, target_field: pa.Field, *, within_list: bool = False +) -> Any: + """Shape raw nested blob values for PyArrow's struct construction.""" + if value is None: + return None + + extension_name = _field_extension_name(target_field) + if extension_name == _BLOB_EXTENSION_NAME and within_list: + if isinstance(value, (bytes, bytearray, memoryview)): + return {"data": value} + if isinstance(value, str): + return {"uri": value} + return value + + if pa.types.is_struct(target_field.type) and isinstance(value, dict): + target_children = {child.name: child for child in target_field.type} + return { + name: _prepare_extension_value( + child_value, target_children[name], within_list=within_list + ) + if name in target_children + else child_value + for name, child_value in value.items() + } + + if _is_list_like(target_field.type) and isinstance(value, (list, tuple)): + return [ + _prepare_extension_value( + item, target_field.type.value_field, within_list=True + ) + for item in value + ] + + return value + + 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): @@ -561,7 +621,16 @@ def _prepare_extension_list(data: DATA, target_schema: pa.Schema) -> DATA: return inferred insert_schema = pa.schema(fields, metadata=inferred.schema.metadata) - return pa.Table.from_pylist(data, schema=insert_schema) + prepared_data = [ + { + name: _prepare_extension_value(value, target_fields[name]) + if name in target_fields + else value + for name, value in row.items() + } + for row in data + ] + return pa.Table.from_pylist(prepared_data, schema=insert_schema) def _is_blob_source_field(field: pa.Field) -> bool: diff --git a/python/python/tests/test_blob.py b/python/python/tests/test_blob.py index 4c453b532..edff6d255 100644 --- a/python/python/tests/test_blob.py +++ b/python/python/tests/test_blob.py @@ -247,6 +247,36 @@ def test_add_all_null_list_to_nested_blob_column(): assert all(blob.as_py() is None for blob in blobs) +@pytest.mark.parametrize("large_list", [False, True], ids=["list", "large_list"]) +def test_add_list_of_dicts_to_blob_list_column(large_list): + db = lancedb.connect("memory:///") + blob_field = lancedb.blob("image") + blob_values = _blob_array("image", [b"seed"]) + if large_list: + items_field = pa.field("items", pa.large_list(blob_field)) + items = pa.LargeListArray.from_arrays( + pa.array([0, 1], type=pa.int64()), blob_values + ) + else: + items_field = pa.field("items", pa.list_(blob_field)) + items = pa.ListArray.from_arrays(pa.array([0, 1], type=pa.int32()), blob_values) + seed = pa.Table.from_arrays( + [pa.array([0], type=pa.int64()), items], + schema=pa.schema([pa.field("id", pa.int64()), items_field]), + ) + table = db.create_table(f"blob_{large_list}_list_add", data=seed) + + table.add([{"id": 1, "items": [None]}]) + table.add( + [{"id": 2, "items": [b"a", None]}], + on_bad_vectors="fill", + ) + + ids = table.search().select(["id"]).to_arrow()["id"].to_pylist() + assert sorted(ids) == [0, 1, 2] + assert pa.types.is_large_list(table.schema.field("items").type) is large_list + + def test_fetch_blob_ranges_aligns_repeated_ranges_and_nulls(): table = _blob_table( "range_alignment",