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:
Rudra Prasad Bhuyan
2026-09-22 14:06:11 -07:00
committed by GitHub
parent a6a6617ee2
commit dd2539c2ca
2 changed files with 340 additions and 2 deletions
+108 -2
View File
@@ -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:
+232
View File
@@ -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):