From dd2539c2cae5e57a1fec0c23944a413177738476 Mon Sep 17 00:00:00 2001 From: Rudra Prasad Bhuyan Date: Wed, 23 Sep 2026 02:36:11 +0530 Subject: [PATCH] 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. --- python/python/lancedb/table.py | 110 +++++++++++++- python/python/tests/test_table.py | 232 ++++++++++++++++++++++++++++++ 2 files changed, 340 insertions(+), 2 deletions(-) diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index d9c27ba81..297b4237f 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -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: diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index 0b3dc01c3..5c56fea38 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -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):