mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-30 08:55:37 +00:00
fix(python): convert objects for JSON fields (#4211)
## Summary Closes: #4060 - Convert Python dict/list objects to JSON strings during ingestion. - Support JSON fields nested inside structs and lists. - Cover `add` and `merge_insert` JSON ingestion paths. - Add regression coverage for nested struct and list JSON fields. ## Testing - `git diff --check` - Python compilation passed. - Focused pytest was attempted but could not complete because the native extension build stalled during `uv` bootstrap.
This commit is contained in:
@@ -5,6 +5,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
import deprecation
|
||||
import warnings
|
||||
from abc import ABC, abstractmethod
|
||||
@@ -266,12 +267,14 @@ def _into_pyarrow_reader(
|
||||
|
||||
# convert to list of dict if data is a bunch of LanceModels
|
||||
if isinstance(data[0], LanceModel):
|
||||
schema = data[0].__class__.to_arrow_schema()
|
||||
model_schema = data[0].__class__.to_arrow_schema()
|
||||
data = [model_to_dict(d) for d in data]
|
||||
return pa.Table.from_pylist(data, schema=schema).to_reader()
|
||||
data = _serialize_json_values(data, schema or model_schema)
|
||||
return pa.Table.from_pylist(data, schema=model_schema).to_reader()
|
||||
elif isinstance(data[0], pa.RecordBatch):
|
||||
return pa.Table.from_batches(data).to_reader()
|
||||
else:
|
||||
data = _serialize_json_values(data, schema)
|
||||
return pa.Table.from_pylist(data).to_reader()
|
||||
elif _check_for_pandas(data) and isinstance(data, pd.DataFrame):
|
||||
table = pa.Table.from_pandas(data, preserve_index=False)
|
||||
@@ -619,6 +622,108 @@ def _field_extension_name(field: pa.Field) -> Optional[str]:
|
||||
return extension_name
|
||||
|
||||
|
||||
def _is_json_field(field: pa.Field) -> bool:
|
||||
return _field_extension_name(field) in ("arrow.json", "lance.json")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _JsonSerializationPlan:
|
||||
arrow_field: pa.Field
|
||||
children: Optional[Dict[str, "_JsonSerializationPlan"]] = None
|
||||
item: Optional["_JsonSerializationPlan"] = None
|
||||
|
||||
|
||||
def _json_serialization_plan(field: pa.Field) -> Optional[_JsonSerializationPlan]:
|
||||
if _is_json_field(field):
|
||||
return _JsonSerializationPlan(field)
|
||||
|
||||
if pa.types.is_struct(field.type):
|
||||
children: Dict[str, _JsonSerializationPlan] = {}
|
||||
for child_field in field.type:
|
||||
child_plan = _json_serialization_plan(child_field)
|
||||
if child_plan is not None:
|
||||
children[child_field.name] = child_plan
|
||||
if children:
|
||||
return _JsonSerializationPlan(field, children=children)
|
||||
|
||||
if _is_list_like(field.type):
|
||||
item_plan = _json_serialization_plan(field.type.value_field)
|
||||
if item_plan is not None:
|
||||
return _JsonSerializationPlan(field, item=item_plan)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _json_serialization_plans(
|
||||
schema: pa.Schema,
|
||||
) -> Dict[str, _JsonSerializationPlan]:
|
||||
plans: Dict[str, _JsonSerializationPlan] = {}
|
||||
for field in schema:
|
||||
plan = _json_serialization_plan(field)
|
||||
if plan is not None:
|
||||
plans[field.name] = plan
|
||||
return plans
|
||||
|
||||
|
||||
def _serialize_json_value(value: Any, plan: _JsonSerializationPlan) -> Any:
|
||||
if value is None or isinstance(value, str):
|
||||
return value
|
||||
if _is_json_field(plan.arrow_field):
|
||||
if isinstance(value, (dict, list)):
|
||||
return json.dumps(value)
|
||||
return value
|
||||
|
||||
if plan.children is not None and isinstance(value, dict):
|
||||
serialized = None
|
||||
for child_name, child_plan in plan.children.items():
|
||||
if child_name not in value:
|
||||
continue
|
||||
child_value = _serialize_json_value(value[child_name], child_plan)
|
||||
if child_value is not value[child_name]:
|
||||
if serialized is None:
|
||||
serialized = dict(value)
|
||||
serialized[child_name] = child_value
|
||||
return serialized if serialized is not None else value
|
||||
|
||||
if plan.item is not None and isinstance(value, list):
|
||||
serialized = None
|
||||
for index, item in enumerate(value):
|
||||
serialized_item = _serialize_json_value(item, plan.item)
|
||||
if serialized_item is not item:
|
||||
if serialized is None:
|
||||
serialized = list(value)
|
||||
serialized[index] = serialized_item
|
||||
return serialized if serialized is not None else value
|
||||
|
||||
return value
|
||||
|
||||
|
||||
def _serialize_json_values(data: Any, target_schema: Optional[pa.Schema]) -> Any:
|
||||
if target_schema is None or not isinstance(data, list):
|
||||
return data
|
||||
|
||||
plans = _json_serialization_plans(target_schema)
|
||||
if not plans:
|
||||
return data
|
||||
|
||||
serialized_rows = []
|
||||
for row in data:
|
||||
if not isinstance(row, dict):
|
||||
serialized_rows.append(row)
|
||||
continue
|
||||
serialized_row = None
|
||||
for field_name, plan in plans.items():
|
||||
if field_name not in row:
|
||||
continue
|
||||
value = _serialize_json_value(row[field_name], plan)
|
||||
if value is not row[field_name]:
|
||||
if serialized_row is None:
|
||||
serialized_row = dict(row)
|
||||
serialized_row[field_name] = value
|
||||
serialized_rows.append(serialized_row if serialized_row is not None else row)
|
||||
return serialized_rows
|
||||
|
||||
|
||||
def _align_field_types(
|
||||
fields: List[pa.Field],
|
||||
target_fields: List[pa.Field],
|
||||
@@ -5720,6 +5825,7 @@ class AsyncTable:
|
||||
|
||||
"""
|
||||
schema = await self.schema()
|
||||
data = _serialize_json_values(data, schema)
|
||||
if on_bad_vectors is None:
|
||||
on_bad_vectors = "error"
|
||||
if fill_value is None:
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
|
||||
import ctypes
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
@@ -17,6 +18,7 @@ from typing import List
|
||||
from unittest.mock import patch
|
||||
|
||||
import lancedb
|
||||
from lancedb import table as table_module
|
||||
from lancedb.dependencies import _PANDAS_AVAILABLE
|
||||
from lancedb.index import BTree, FTS, HnswFlat, HnswPq, HnswSq, IvfPq
|
||||
import numpy as np
|
||||
@@ -3027,6 +3029,236 @@ async def test_add_sanitization_encodes_json(mem_db_async: AsyncConnection):
|
||||
assert rows == [{"id": "c", "j": '{"k":3}'}]
|
||||
|
||||
|
||||
@pytest.mark.skipif(not hasattr(pa, "json_"), reason="requires PyArrow JSON type")
|
||||
def test_add_python_objects_to_json_column(mem_db: DBConnection):
|
||||
schema = pa.schema(
|
||||
[
|
||||
pa.field("id", pa.int64()),
|
||||
pa.field("payload", pa.json_()),
|
||||
pa.field("metadata", pa.struct([("label", pa.string())])),
|
||||
pa.field("tags", pa.list_(pa.string())),
|
||||
]
|
||||
)
|
||||
table = mem_db.create_table("json_python_objects_add", schema=schema)
|
||||
table.add(
|
||||
[
|
||||
{
|
||||
"id": 1,
|
||||
"payload": {"foo": "bar", "count": 2},
|
||||
"metadata": {"label": "dict"},
|
||||
"tags": ["a", "b"],
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"payload": {
|
||||
"name": "alice",
|
||||
"tags": ["x", "y"],
|
||||
"nested": {"enabled": True},
|
||||
},
|
||||
"metadata": {"label": "nested"},
|
||||
"tags": ["c"],
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"payload": ["x", "y", "z"],
|
||||
"metadata": {"label": "list"},
|
||||
"tags": ["d", "e"],
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"payload": '{"foo": "bar"}',
|
||||
"metadata": {"label": "string"},
|
||||
"tags": ["f"],
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"payload": None,
|
||||
"metadata": {"label": "null"},
|
||||
"tags": ["g"],
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
rows = {row["id"]: row for row in table.to_arrow().to_pylist()}
|
||||
assert json.loads(rows[1]["payload"]) == {"foo": "bar", "count": 2}
|
||||
assert json.loads(rows[2]["payload"]) == {
|
||||
"name": "alice",
|
||||
"tags": ["x", "y"],
|
||||
"nested": {"enabled": True},
|
||||
}
|
||||
assert json.loads(rows[3]["payload"]) == ["x", "y", "z"]
|
||||
assert json.loads(rows[4]["payload"]) == {"foo": "bar"}
|
||||
assert rows[5]["payload"] is None
|
||||
assert rows[1]["metadata"] == {"label": "dict"}
|
||||
assert rows[1]["tags"] == ["a", "b"]
|
||||
|
||||
matched = table.search().where("json_extract(payload, '$.count') = '2'").to_list()
|
||||
assert [row["id"] for row in matched] == [1]
|
||||
|
||||
|
||||
@pytest.mark.skipif(not hasattr(pa, "json_"), reason="requires PyArrow JSON type")
|
||||
def test_merge_insert_python_objects_to_json_column(mem_db: DBConnection):
|
||||
schema = pa.schema(
|
||||
[
|
||||
pa.field("id", pa.int64()),
|
||||
pa.field("payload", pa.json_()),
|
||||
pa.field("metadata", pa.struct([("label", pa.string())])),
|
||||
pa.field("tags", pa.list_(pa.string())),
|
||||
]
|
||||
)
|
||||
table = mem_db.create_table("json_python_objects_merge", schema=schema)
|
||||
table.add(
|
||||
[
|
||||
{
|
||||
"id": 1,
|
||||
"payload": {"old": True},
|
||||
"metadata": {"label": "old"},
|
||||
"tags": ["old"],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
table.merge_insert(
|
||||
"id"
|
||||
).when_matched_update_all().when_not_matched_insert_all().execute(
|
||||
[
|
||||
{
|
||||
"id": 1,
|
||||
"payload": {"foo": "bar", "count": 2},
|
||||
"metadata": {"label": "dict"},
|
||||
"tags": ["a", "b"],
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"payload": {
|
||||
"name": "alice",
|
||||
"tags": ["x", "y"],
|
||||
"nested": {"enabled": True},
|
||||
},
|
||||
"metadata": {"label": "nested"},
|
||||
"tags": ["c"],
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"payload": ["x", "y", "z"],
|
||||
"metadata": {"label": "list"},
|
||||
"tags": ["d", "e"],
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"payload": '{"foo": "bar"}',
|
||||
"metadata": {"label": "string"},
|
||||
"tags": ["f"],
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"payload": None,
|
||||
"metadata": {"label": "null"},
|
||||
"tags": ["g"],
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
rows = {row["id"]: row for row in table.to_arrow().to_pylist()}
|
||||
assert json.loads(rows[1]["payload"]) == {"foo": "bar", "count": 2}
|
||||
assert json.loads(rows[2]["payload"]) == {
|
||||
"name": "alice",
|
||||
"tags": ["x", "y"],
|
||||
"nested": {"enabled": True},
|
||||
}
|
||||
assert json.loads(rows[3]["payload"]) == ["x", "y", "z"]
|
||||
assert json.loads(rows[4]["payload"]) == {"foo": "bar"}
|
||||
assert rows[5]["payload"] is None
|
||||
assert rows[1]["metadata"] == {"label": "dict"}
|
||||
assert rows[1]["tags"] == ["a", "b"]
|
||||
|
||||
matched = table.search().where("json_extract(payload, '$.count') = '2'").to_list()
|
||||
assert [row["id"] for row in matched] == [1]
|
||||
|
||||
|
||||
@pytest.mark.skipif(not hasattr(pa, "json_"), reason="requires PyArrow JSON type")
|
||||
def test_add_python_objects_to_nested_json_fields(mem_db: DBConnection):
|
||||
schema = pa.schema(
|
||||
[
|
||||
pa.field("id", pa.int64()),
|
||||
pa.field(
|
||||
"info",
|
||||
pa.struct([pa.field("payload", pa.json_())]),
|
||||
),
|
||||
pa.field("documents", pa.list_(pa.field("item", pa.json_()))),
|
||||
]
|
||||
)
|
||||
table = mem_db.create_table("nested_json_python_objects_add", schema=schema)
|
||||
|
||||
table.add(
|
||||
[
|
||||
{
|
||||
"id": 1,
|
||||
"info": {"payload": {"kind": "struct", "value": 1}},
|
||||
"documents": [{"kind": "list", "value": 2}],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
row = table.to_arrow().to_pylist()[0]
|
||||
assert json.loads(row["info"]["payload"]) == {"kind": "struct", "value": 1}
|
||||
assert json.loads(row["documents"][0]) == {"kind": "list", "value": 2}
|
||||
assert [
|
||||
row["id"]
|
||||
for row in table.search()
|
||||
.where("json_extract(info.payload, '$.value') = '1'")
|
||||
.to_list()
|
||||
] == [1]
|
||||
assert [
|
||||
row["id"]
|
||||
for row in table.search()
|
||||
.where("json_extract(documents[1], '$.value') = '2'")
|
||||
.to_list()
|
||||
] == [1]
|
||||
|
||||
|
||||
@pytest.mark.skipif(not hasattr(pa, "json_"), reason="requires PyArrow JSON type")
|
||||
def test_json_serialization_plan_skips_non_json_branches():
|
||||
schema = pa.schema(
|
||||
[
|
||||
pa.field("id", pa.int64()),
|
||||
pa.field("vector", pa.list_(pa.float32(), 3)),
|
||||
pa.field("payload", pa.json_()),
|
||||
pa.field(
|
||||
"metadata",
|
||||
pa.struct(
|
||||
[
|
||||
pa.field("label", pa.string()),
|
||||
pa.field("payload", pa.json_()),
|
||||
]
|
||||
),
|
||||
),
|
||||
pa.field(
|
||||
"documents",
|
||||
pa.list_(
|
||||
pa.struct(
|
||||
[
|
||||
pa.field("title", pa.string()),
|
||||
pa.field("payload", pa.json_()),
|
||||
]
|
||||
)
|
||||
),
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
plans = table_module._json_serialization_plans(schema)
|
||||
|
||||
assert set(plans) == {"payload", "metadata", "documents"}
|
||||
metadata_children = plans["metadata"].children
|
||||
assert metadata_children is not None
|
||||
assert set(metadata_children) == {"payload"}
|
||||
documents_item = plans["documents"].item
|
||||
assert documents_item is not None
|
||||
assert documents_item.children is not None
|
||||
assert set(documents_item.children) == {"payload"}
|
||||
|
||||
|
||||
@pytest.mark.skipif(not hasattr(pa, "json_"), reason="requires PyArrow JSON type")
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_all_null_json_batch(mem_db_async: AsyncConnection):
|
||||
|
||||
Reference in New Issue
Block a user