This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git
The following commit(s) were added to refs/heads/master by this push:
new 13b23098fb [java][python] Add VARIANT path upsert via variant_set
(#9253)
13b23098fb is described below
commit 13b23098fbffff073e7a1899373b784528a326c3
Author: XiaoHongbo <[email protected]>
AuthorDate: Tue Aug 18 15:58:13 2026 +0800
[java][python] Add VARIANT path upsert via variant_set (#9253)
---
docs/docs/pypaimon/python-api.mdx | 20 +-
.../apache/paimon/data/variant/GenericVariant.java | 98 ++-
.../paimon/data/variant/GenericVariantBuilder.java | 12 +-
.../paimon/data/variant/GenericVariantUtil.java | 21 +
.../paimon/data/variant/GenericVariantTest.java | 35 +
paimon-python/pypaimon/data/__init__.py | 8 +-
paimon-python/pypaimon/data/generic_variant.py | 12 +-
paimon-python/pypaimon/data/variant_path.py | 760 +++++++++++++++--
paimon-python/pypaimon/tests/variant_path_test.py | 41 +
paimon-python/pypaimon/tests/variant_set_test.py | 898 +++++++++++++++++++++
10 files changed, 1800 insertions(+), 105 deletions(-)
diff --git a/docs/docs/pypaimon/python-api.mdx
b/docs/docs/pypaimon/python-api.mdx
index 23a27a1e13..b611b2d42e 100644
--- a/docs/docs/pypaimon/python-api.mdx
+++ b/docs/docs/pypaimon/python-api.mdx
@@ -1136,7 +1136,7 @@ without decoding unrelated fields:
import pyarrow as pa
import pyarrow.compute as pc
-from pypaimon.data import variant_get, variant_replace
+from pypaimon.data import variant_get, variant_replace, variant_set
current = variant_get(payload, '$.velocity.y', pa.float64())
updated_payload = variant_replace(
@@ -1149,6 +1149,24 @@ string-keyed map types. Replacement supports scalar
types. Missing paths read
as NULL and remain unchanged unless `strict=True` is specified. Pass mappings
to process multiple paths in one pass.
+`variant_set` upserts paths: existing paths are replaced like
+`variant_replace`, and a missing final key is inserted when its parent path
+exists and is an OBJECT:
+
+```python
+updated_payload = variant_set(payload, {
+ '$.velocity.y': pc.negate(current),
+ '$.processed': pa.scalar(True, type=pa.bool_()),
+})
+```
+
+Values may be a `pa.Scalar` (broadcast to every row) or a `pa.Array` /
+`pa.ChunkedArray` with one value per row; Arrow NULL values are stored as
+VARIANT NULL and SQL NULL rows are preserved. `variant_set` raises
+`ValueError` when an intermediate path is missing, when the parent of a
+missing key is not an OBJECT, or for a missing array index — it never
+creates intermediate objects or extends arrays.
+
**`GenericVariant` API:**
diff --git
a/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariant.java
b/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariant.java
index 89181c72e0..9424e061c4 100644
---
a/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariant.java
+++
b/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariant.java
@@ -18,8 +18,10 @@
package org.apache.paimon.data.variant;
+import org.apache.paimon.data.BinaryString;
import org.apache.paimon.data.variant.VariantPathSegment.ArrayExtraction;
import org.apache.paimon.data.variant.VariantPathSegment.ObjectExtraction;
+import org.apache.paimon.memory.MemorySegment;
import org.apache.paimon.types.DataType;
import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.core.JsonFactory;
@@ -52,6 +54,7 @@ import static
org.apache.paimon.data.variant.GenericVariantUtil.getMetadataKey;
import static org.apache.paimon.data.variant.GenericVariantUtil.handleArray;
import static org.apache.paimon.data.variant.GenericVariantUtil.handleObject;
import static
org.apache.paimon.data.variant.GenericVariantUtil.malformedVariant;
+import static
org.apache.paimon.data.variant.GenericVariantUtil.pointToMetadataKey;
import static org.apache.paimon.data.variant.GenericVariantUtil.readUnsigned;
import static org.apache.paimon.data.variant.GenericVariantUtil.valueSize;
import static
org.apache.paimon.data.variant.GenericVariantUtil.variantConstructorSizeLimit;
@@ -242,12 +245,13 @@ public final class GenericVariant implements Variant,
Serializable {
value,
pos,
(size, idSize, offsetSize, idStart, offsetStart, dataStart) ->
{
+ MetadataKeyLookup keyLookup = new
MetadataKeyLookup(metadata, key);
// Use linear search for a short list. Switch to binary
search when the length
// reaches `BINARY_SEARCH_THRESHOLD`.
if (size < BINARY_SEARCH_THRESHOLD) {
for (int i = 0; i < size; ++i) {
int id = readUnsigned(value, idStart + idSize * i,
idSize);
- if (key.equals(getMetadataKey(metadata, id))) {
+ if (keyLookup.compareUtf8(id) == 0) {
int offset =
readUnsigned(
value, offsetStart +
offsetSize * i, offsetSize);
@@ -255,31 +259,83 @@ public final class GenericVariant implements Variant,
Serializable {
}
}
} else {
- int low = 0;
- int high = size - 1;
- while (low <= high) {
- // Use unsigned right shift to compute the middle
of `low` and `high`.
- // This is not only a performance optimization,
because it can properly
- // handle the case where `low + high` overflows
int.
- int mid = (low + high) >>> 1;
- int id = readUnsigned(value, idStart + idSize *
mid, idSize);
- int cmp = getMetadataKey(metadata,
id).compareTo(key);
- if (cmp < 0) {
- low = mid + 1;
- } else if (cmp > 0) {
- high = mid - 1;
- } else {
- int offset =
- readUnsigned(
- value, offsetStart +
offsetSize * mid, offsetSize);
- return new GenericVariant(value, metadata,
dataStart + offset);
- }
- }
+ GenericVariant result =
+ binarySearchObjectField(
+ size,
+ idSize,
+ offsetSize,
+ idStart,
+ offsetStart,
+ dataStart,
+ keyLookup,
+ true);
+ return result != null
+ ? result
+ : binarySearchObjectField(
+ size,
+ idSize,
+ offsetSize,
+ idStart,
+ offsetStart,
+ dataStart,
+ keyLookup,
+ false);
}
return null;
});
}
+ private GenericVariant binarySearchObjectField(
+ int size,
+ int idSize,
+ int offsetSize,
+ int idStart,
+ int offsetStart,
+ int dataStart,
+ MetadataKeyLookup keyLookup,
+ boolean utf8Order) {
+ int low = 0;
+ int high = size - 1;
+ while (low <= high) {
+ int mid = (low + high) >>> 1;
+ int id = readUnsigned(value, idStart + idSize * mid, idSize);
+ int comparison =
+ utf8Order
+ ? keyLookup.compareUtf8(id)
+ : getMetadataKey(metadata,
id).compareTo(keyLookup.key);
+ if (comparison < 0) {
+ low = mid + 1;
+ } else if (comparison > 0) {
+ high = mid - 1;
+ } else {
+ int offset = readUnsigned(value, offsetStart + offsetSize *
mid, offsetSize);
+ return new GenericVariant(value, metadata, dataStart + offset);
+ }
+ }
+ return null;
+ }
+
+ private static final class MetadataKeyLookup {
+ private final byte[] metadata;
+ private final String key;
+ private final MemorySegment[] metadataSegments;
+ private final BinaryString binaryKey;
+ private final BinaryString candidate;
+
+ private MetadataKeyLookup(byte[] metadata, String key) {
+ this.metadata = metadata;
+ this.key = key;
+ this.metadataSegments = new MemorySegment[]
{MemorySegment.wrap(metadata)};
+ this.binaryKey = BinaryString.fromString(key);
+ this.candidate = BinaryString.fromAddress(metadataSegments, 0, 0);
+ }
+
+ private int compareUtf8(int id) {
+ pointToMetadataKey(metadata, metadataSegments, id, candidate);
+ return candidate.compareTo(binaryKey);
+ }
+ }
+
/** Variant object field. */
public static final class ObjectField {
public final String key;
diff --git
a/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantBuilder.java
b/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantBuilder.java
index 921832a0ae..fa752cdb06 100644
---
a/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantBuilder.java
+++
b/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantBuilder.java
@@ -18,6 +18,8 @@
package org.apache.paimon.data.variant;
+import org.apache.paimon.data.BinaryString;
+
import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.core.JsonFactory;
import
org.apache.paimon.shade.jackson2.com.fasterxml.jackson.core.JsonParseException;
import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.core.JsonParser;
@@ -528,22 +530,28 @@ public class GenericVariantBuilder {
*/
public static final class FieldEntry implements Comparable<FieldEntry> {
final String key;
+ final BinaryString binaryKey;
final int id;
final int offset;
public FieldEntry(String key, int id, int offset) {
+ this(key, BinaryString.fromString(key), id, offset);
+ }
+
+ private FieldEntry(String key, BinaryString binaryKey, int id, int
offset) {
this.key = key;
+ this.binaryKey = binaryKey;
this.id = id;
this.offset = offset;
}
FieldEntry withNewOffset(int newOffset) {
- return new FieldEntry(key, id, newOffset);
+ return new FieldEntry(key, binaryKey, id, newOffset);
}
@Override
public int compareTo(FieldEntry other) {
- return key.compareTo(other.key);
+ return binaryKey.compareTo(other.binaryKey);
}
}
diff --git
a/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantUtil.java
b/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantUtil.java
index c99428b9d0..6711738743 100644
---
a/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantUtil.java
+++
b/paimon-common/src/main/java/org/apache/paimon/data/variant/GenericVariantUtil.java
@@ -18,6 +18,9 @@
package org.apache.paimon.data.variant;
+import org.apache.paimon.data.BinaryString;
+import org.apache.paimon.memory.MemorySegment;
+
import java.math.BigDecimal;
import java.math.BigInteger;
import java.nio.ByteBuffer;
@@ -651,4 +654,22 @@ public class GenericVariantUtil {
checkIndex(stringStart + nextOffset - 1, metadata.length);
return new String(metadata, stringStart + offset, nextOffset - offset);
}
+
+ static void pointToMetadataKey(
+ byte[] metadata, MemorySegment[] metadataSegments, int id,
BinaryString result) {
+ checkIndex(0, metadata.length);
+ int offsetSize = ((metadata[0] >> 6) & 0x3) + 1;
+ int dictSize = readUnsigned(metadata, 1, offsetSize);
+ if (id >= dictSize) {
+ throw malformedVariant();
+ }
+ int stringStart = 1 + (dictSize + 2) * offsetSize;
+ int offset = readUnsigned(metadata, 1 + (id + 1) * offsetSize,
offsetSize);
+ int nextOffset = readUnsigned(metadata, 1 + (id + 2) * offsetSize,
offsetSize);
+ if (offset > nextOffset) {
+ throw malformedVariant();
+ }
+ checkIndex(stringStart + nextOffset - 1, metadata.length);
+ result.pointTo(metadataSegments, stringStart + offset, nextOffset -
offset);
+ }
}
diff --git
a/paimon-common/src/test/java/org/apache/paimon/data/variant/GenericVariantTest.java
b/paimon-common/src/test/java/org/apache/paimon/data/variant/GenericVariantTest.java
index e7d81b57a2..227eca2da0 100644
---
a/paimon-common/src/test/java/org/apache/paimon/data/variant/GenericVariantTest.java
+++
b/paimon-common/src/test/java/org/apache/paimon/data/variant/GenericVariantTest.java
@@ -138,6 +138,41 @@ public class GenericVariantTest {
assertThat(variant.variantGet("$.nullField", DataTypes.BOOLEAN(),
castArgs)).isNull();
}
+ @Test
+ public void testObjectFieldOrderingCompatibility() {
+ String bmpKey = "\uE000";
+ String supplementaryKey = new String(Character.toChars(0x10000));
+ StringBuilder json = new StringBuilder("{");
+ for (int i = 0; i < 30; i++) {
+ if (i > 0) {
+ json.append(',');
+ }
+ json.append("\"k").append(i < 10 ? "0" :
"").append(i).append("\":null");
+ }
+ json.append(",\"").append(bmpKey).append("\":null");
+ json.append(",\"").append(supplementaryKey).append("\":null}");
+
+ GenericVariant specOrdered = GenericVariant.fromJson(json.toString());
+ byte[] specValue = specOrdered.value();
+ int idStart = 2;
+ int offsetStart = idStart + 32;
+ assertThat(specValue[idStart + 30] & 0xFF).isEqualTo(30);
+ assertThat(specValue[idStart + 31] & 0xFF).isEqualTo(31);
+ assertThat(specOrdered.getFieldByKey(bmpKey)).isNotNull();
+ assertThat(specOrdered.getFieldByKey(supplementaryKey)).isNotNull();
+
+ byte[] legacyValue = specValue.clone();
+ byte temporary = legacyValue[idStart + 30];
+ legacyValue[idStart + 30] = legacyValue[idStart + 31];
+ legacyValue[idStart + 31] = temporary;
+ temporary = legacyValue[offsetStart + 30];
+ legacyValue[offsetStart + 30] = legacyValue[offsetStart + 31];
+ legacyValue[offsetStart + 31] = temporary;
+ GenericVariant legacyOrdered = new GenericVariant(legacyValue,
specOrdered.metadata());
+ assertThat(legacyOrdered.getFieldByKey(bmpKey)).isNotNull();
+ assertThat(legacyOrdered.getFieldByKey(supplementaryKey)).isNotNull();
+ }
+
@Test
public void testShredding() {
GenericVariant variant = GenericVariant.fromJson("{\"a\": 1, \"b\":
\"hello\"}");
diff --git a/paimon-python/pypaimon/data/__init__.py
b/paimon-python/pypaimon/data/__init__.py
index c0211f0460..692a8ac925 100644
--- a/paimon-python/pypaimon/data/__init__.py
+++ b/paimon-python/pypaimon/data/__init__.py
@@ -22,13 +22,18 @@ if sys.version_info[:2] < (3, 7):
# Module-level __getattr__ is unavailable before Python 3.7.
from pypaimon.data.timestamp import Timestamp
from pypaimon.data.decimal import Decimal
- from pypaimon.data.variant_path import variant_get, variant_replace
+ from pypaimon.data.variant_path import (
+ variant_get,
+ variant_replace,
+ variant_set,
+ )
__all__ = [
'Timestamp',
'Decimal',
'variant_get',
'variant_replace',
+ 'variant_set',
]
_MODULE_BY_EXPORT = {
@@ -36,6 +41,7 @@ _MODULE_BY_EXPORT = {
'Decimal': 'pypaimon.data.decimal',
'variant_get': 'pypaimon.data.variant_path',
'variant_replace': 'pypaimon.data.variant_path',
+ 'variant_set': 'pypaimon.data.variant_path',
}
diff --git a/paimon-python/pypaimon/data/generic_variant.py
b/paimon-python/pypaimon/data/generic_variant.py
index 94b37210c6..491f8dfae1 100644
--- a/paimon-python/pypaimon/data/generic_variant.py
+++ b/paimon-python/pypaimon/data/generic_variant.py
@@ -91,6 +91,11 @@ _EPOCH_DT_UTC = datetime.datetime(1970, 1, 1,
tzinfo=datetime.timezone.utc)
_EPOCH_DT_NTZ = datetime.datetime(1970, 1, 1)
+def _check_variant_sizes(value_size, metadata_size):
+ if value_size > _SIZE_LIMIT or metadata_size > _SIZE_LIMIT:
+ raise ValueError('VARIANT_CONSTRUCTOR_SIZE_LIMIT')
+
+
class _Type(enum.Enum):
"""Internal high-level variant value types (many-to-one from wire
types)."""
OBJECT = 'OBJECT'
@@ -264,8 +269,9 @@ class _GenericVariantBuilder:
def _ensure(self, n):
needed = self._pos + n
+ _check_variant_sizes(needed, 0)
if needed > len(self._buf):
- new_cap = max(needed, len(self._buf) * 2)
+ new_cap = min(_SIZE_LIMIT, max(needed, len(self._buf) * 2))
new_buf = bytearray(new_cap)
new_buf[:self._pos] = self._buf[:self._pos]
self._buf = new_buf
@@ -388,7 +394,7 @@ class _GenericVariantBuilder:
self._write_le(micros_since_epoch & 0xFFFFFFFFFFFFFFFF, 8)
def _finish_writing_object(self, start, fields):
- fields.sort(key=lambda f: f[0])
+ fields.sort(key=lambda f: f[0].encode('utf-8'))
for i in range(1, len(fields)):
if fields[i][0] == fields[i - 1][0]:
raise ValueError('Duplicate key in variant object')
@@ -516,6 +522,7 @@ class _GenericVariantBuilder:
offset_start = 1 + offset_size
string_start = offset_start + (n_keys + 1) * offset_size
metadata_size = string_start + total_str_size
+ _check_variant_sizes(self._pos, metadata_size)
metadata = bytearray(metadata_size)
metadata[0] = _VERSION | ((offset_size - 1) << 6)
@@ -570,6 +577,7 @@ class GenericVariant:
__slots__ = ('_value', '_metadata', '_pos')
def __init__(self, value: bytes, metadata: bytes, _pos: int = 0):
+ _check_variant_sizes(len(value), len(metadata))
self._value = bytes(value)
self._metadata = bytes(metadata)
self._pos = _pos
diff --git a/paimon-python/pypaimon/data/variant_path.py
b/paimon-python/pypaimon/data/variant_path.py
index e7ce3d0adb..17fcbab08b 100644
--- a/paimon-python/pypaimon/data/variant_path.py
+++ b/paimon-python/pypaimon/data/variant_path.py
@@ -14,24 +14,30 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-"""Read and replace paths in Arrow VARIANT columns."""
+"""Read, replace, and upsert paths in Arrow VARIANT columns."""
import functools
import re
import struct
+import threading
from typing import Dict, Mapping, Optional, Sequence, Tuple
import numpy as np
import pyarrow as pa
+from cachetools import LRUCache
from pypaimon.data._variant_binary import (
_ARRAY,
_OBJECT,
_PRIMITIVE,
_SHORT_STR,
+ _U8_MAX,
_U32_SIZE,
_VERSION,
_VERSION_MASK,
+ _array_header,
+ _get_int_size,
+ _object_header,
_primitive_header,
_read_unsigned,
)
@@ -50,6 +56,7 @@ from pypaimon.data.generic_variant import (
_PRIMITIVE_FIXED_SIZES,
GenericVariant,
_Type,
+ _check_variant_sizes,
_variant_get_type,
)
from pypaimon.data.variant_shredding import (
@@ -132,6 +139,91 @@ def _validate_metadata_version(metadata):
_malformed("invalid metadata version")
+_metadata_cache = threading.local()
+_NO_METADATA_CACHE = object()
+
+
+def _with_metadata_cache(func):
+ @functools.wraps(func)
+ def wrapper(*args, **kwargs):
+ previous = getattr(_metadata_cache, 'value', _NO_METADATA_CACHE)
+ _metadata_cache.value = LRUCache(maxsize=256)
+ try:
+ return func(*args, **kwargs)
+ finally:
+ if previous is _NO_METADATA_CACHE:
+ del _metadata_cache.value
+ else:
+ _metadata_cache.value = previous
+
+ return wrapper
+
+
+def _cached_metadata_key_ids(metadata: bytes) -> Dict[str, int]:
+ cache = getattr(_metadata_cache, 'value', None)
+ if cache is not None:
+ hit = cache.get(metadata)
+ if hit is not None:
+ return hit
+ key_ids = _metadata_key_ids(metadata)
+ if cache is not None:
+ cache[metadata] = key_ids
+ return key_ids
+
+
+def _metadata_with_keys(metadata: bytes, new_keys: Tuple[str, ...]):
+ """Append missing keys to the metadata dictionary, keeping ids.
+
+ Returns read-only ``(new_metadata, key_ids, names_by_id)``;
+ ``new_metadata`` is None when every key already exists.
+ """
+ cache_key = (metadata, new_keys)
+ cache = getattr(_metadata_cache, 'value', None)
+ if cache is not None:
+ hit = cache.get(cache_key)
+ if hit is not None:
+ return hit
+ key_ids = dict(_cached_metadata_key_ids(metadata))
+ names = [None] * len(key_ids)
+ for key, key_id in key_ids.items():
+ names[key_id] = key
+ missing = list(dict.fromkeys(
+ key for key in new_keys if key not in key_ids))
+ for key in missing:
+ key_ids[key] = len(names)
+ names.append(key)
+ names_by_id = {key_id: key for key, key_id in key_ids.items()}
+ if not missing:
+ result = (None, key_ids, names_by_id)
+ else:
+ encoded = [name.encode('utf-8') for name in names]
+ total_size = sum(len(name) for name in encoded)
+ max_size = max(total_size, len(encoded))
+ offset_size = _get_int_size(max_size) if max_size > 0 else 1
+ offset_start = 1 + offset_size
+ string_start = offset_start + (len(encoded) + 1) * offset_size
+ metadata_size = string_start + total_size
+ _check_variant_sizes(0, metadata_size)
+ rebuilt = bytearray(metadata_size)
+ rebuilt[0] = _VERSION | ((offset_size - 1) << 6)
+ rebuilt[1:1 + offset_size] = len(encoded).to_bytes(
+ offset_size, 'little')
+ current = 0
+ for index, name in enumerate(encoded):
+ rebuilt[offset_start + index * offset_size:
+ offset_start + (index + 1) * offset_size] = (
+ current.to_bytes(offset_size, 'little'))
+ rebuilt[string_start + current:
+ string_start + current + len(name)] = name
+ current += len(name)
+ rebuilt[offset_start + len(encoded) * offset_size:string_start] = (
+ current.to_bytes(offset_size, 'little'))
+ result = (bytes(rebuilt), key_ids, names_by_id)
+ if cache is not None:
+ cache[cache_key] = result
+ return result
+
+
def _malformed(message):
raise ValueError(f"MALFORMED_VARIANT: {message}")
@@ -163,6 +255,10 @@ def _checked_object_layout(value, pos, limit):
or len(set(offsets[:-1])) != size))
or any(offset >= sentinel for offset in offsets[:-1])):
_malformed("invalid object offsets")
+ if size and len({
+ _read_unsigned(value, id_start + i * id_width, id_width)
+ for i in range(size)}) != size:
+ _malformed("duplicate object field id")
_require_range(data_start, sentinel, limit)
return (
size, id_width, id_start, data_start, offsets,
@@ -249,6 +345,43 @@ def _checked_value_size(value, pos, limit=None):
return end - pos
+def _validate_value_field_ids(value, pos, limit, metadata_size):
+ """Validate object field ids in one unedited value subtree."""
+ stack = [(pos, limit)]
+ while stack:
+ current_pos, current_limit = stack.pop()
+ value_end = current_pos + _checked_value_size(
+ value, current_pos, current_limit)
+ if value_end != current_limit:
+ _malformed("child size does not match container offsets")
+ basic_type = value[current_pos] & 0x3
+ if basic_type == _OBJECT:
+ size, id_size, id_start, data_start, offsets, _ = (
+ _checked_object_layout(value, current_pos, value_end))
+ ids = [
+ _read_unsigned(value, id_start + i * id_size, id_size)
+ for i in range(size)
+ ]
+ if any(field_id >= metadata_size for field_id in ids):
+ _malformed("object field id is missing from metadata")
+ ordered_offsets = sorted(offsets)
+ end_by_offset = dict(zip(
+ ordered_offsets, ordered_offsets[1:]))
+ for slot in range(size):
+ child_start, child_end = _checked_object_child_bounds(
+ value, data_start, offsets, slot, end_by_offset)
+ if (value[child_start] & 0x3) in (_OBJECT, _ARRAY):
+ stack.append((child_start, child_end))
+ elif basic_type == _ARRAY:
+ size, data_start, offsets, _ = _checked_array_layout(
+ value, current_pos, value_end)
+ for index in range(size):
+ stack.append((
+ data_start + offsets[index],
+ data_start + offsets[index + 1],
+ ))
+
+
def _field_slot(id_table: bytes, id_size: int, key_id: int) -> Optional[int]:
for slot in range(len(id_table) // id_size):
if _read_unsigned(id_table, slot * id_size, id_size) == key_id:
@@ -259,15 +392,19 @@ def _field_slot(id_table: bytes, id_size: int, key_id:
int) -> Optional[int]:
@functools.lru_cache(maxsize=256)
def _compile_paths(paths: Tuple[_Path, ...]):
nodes = [(None, None, None)]
- node_by_prefix = {(): 0}
+ node_by_edge = {}
results = []
for path in paths:
- for length in range(1, len(path) + 1):
- prefix = path[:length]
- if prefix not in node_by_prefix:
- node_by_prefix[prefix] = len(nodes)
- nodes.append((node_by_prefix[prefix[:-1]],) + prefix[-1])
- results.append(node_by_prefix[path])
+ parent = 0
+ for segment in path:
+ edge = (parent, segment)
+ node = node_by_edge.get(edge)
+ if node is None:
+ node = len(nodes)
+ node_by_edge[edge] = node
+ nodes.append((parent,) + segment)
+ parent = node
+ results.append(parent)
return tuple(nodes), tuple(results)
@@ -282,7 +419,7 @@ def _path_positions(
nodes, result_nodes = _compile_paths(tuple(paths))
_validate_metadata_version(metadata)
key_ids = (
- _metadata_key_ids(metadata)
+ _cached_metadata_key_ids(bytes(metadata))
if any(kind == 'key' for _, kind, _ in nodes[1:]) else {}
)
bounds = [(0, len(value))]
@@ -596,7 +733,7 @@ def _vectorized_path_positions(
nodes, result_nodes = _compile_paths(tuple(paths))
_validate_metadata_version(first_metadata)
key_ids = (
- _metadata_key_ids(first_metadata)
+ _cached_metadata_key_ids(first_metadata)
if any(kind == 'key' for _, kind, _ in nodes[1:]) else {}
)
row_offsets = values.numpy_offsets()
@@ -640,15 +777,6 @@ def _vectorized_path_positions(
id_table = bytes(
first_value[id_start:id_start + size * id_size])
slot = _field_slot(id_table, id_size, key_id)
- if slot is None:
- return None
- successor_slot = min(
- (
- index for index in range(size + 1)
- if first_offsets[index] > first_offsets[slot]
- ),
- key=lambda index: first_offsets[index],
- )
size_widths = np.where(
((type_info >> 4) & 0x1) != 0, _U32_SIZE, 1)
@@ -665,10 +793,24 @@ def _vectorized_path_positions(
data_starts = offset_starts + (sizes + 1) * offset_widths
if np.any(data_starts > parent_ends):
return None
- ids = _take_unsigned(
- data, id_starts + slot * id_widths, id_widths)
- if np.any(ids != key_id):
- return None
+ for index in range(size):
+ expected_id = _read_unsigned(
+ first_value, id_start + index * id_size, id_size)
+ ids = _take_unsigned(
+ data, id_starts + index * id_widths, id_widths)
+ if np.any(ids != expected_id):
+ return None
+ if slot is None:
+ positions.append(None)
+ limits.append(None)
+ continue
+ successor_slot = min(
+ (
+ index for index in range(size + 1)
+ if first_offsets[index] > first_offsets[slot]
+ ),
+ key=lambda index: first_offsets[index],
+ )
else:
if np.any((headers & 0x3) != _ARRAY):
return None
@@ -1042,8 +1184,7 @@ def _rebuilt_offsets(lengths, value_format):
return offsets
-def _sparse_rebuilt_chunk(
- chunk, values, data, data_start, rebuilt_rows):
+def _sparse_rebuilt_binary(values, data, data_start, rebuilt_rows):
old_offsets = values.numpy_offsets()
lengths = old_offsets[1:] - old_offsets[:-1]
for row, rebuilt in rebuilt_rows.items():
@@ -1068,12 +1209,18 @@ def _sparse_rebuilt_chunk(
None if values.array.null_count == 0
else values.array.is_valid().buffers()[1]
)
- rebuilt_values = pa.Array.from_buffers(
+ return pa.Array.from_buffers(
values.array.type,
- len(chunk),
+ len(values.array),
[validity, pa.py_buffer(offsets), pa.py_buffer(output)],
null_count=values.array.null_count,
)
+
+
+def _sparse_rebuilt_chunk(
+ chunk, values, data, data_start, rebuilt_rows):
+ rebuilt_values = _sparse_rebuilt_binary(
+ values, data, data_start, rebuilt_rows)
return pa.StructArray.from_arrays(
[rebuilt_values, chunk.field(1)],
fields=list(chunk.type),
@@ -1162,6 +1309,70 @@ class _Replacement:
f"VARIANT path type does not match {self.type}")
+class _PatchState:
+ """Lazy copy-on-write buffer shared by per-group in-place patches."""
+
+ def __init__(self, values: _BinaryValues):
+ self._values = values
+ self.data = None
+ self.data_start = 0
+ self.output_data = None
+
+ def ensure(self):
+ if self.data is None:
+ self.data, self.data_start = self._values.copy_used_data()
+ self.output_data = np.frombuffer(self.data, dtype=np.uint8)
+ return self.output_data
+
+
+def _patch_planned_group(
+ planned, parsed, chunk_length, global_row, state, slow_rows, strict):
+ rows, row_starts, source_data, positions, limits = planned
+ replacements = []
+ compatible = np.ones(len(rows), dtype=bool)
+ has_replacement = False
+ for (path, _, provider), pos, limit in zip(parsed, positions, limits):
+ if pos is None:
+ if strict:
+ raise ValueError(
+ f"VARIANT path does not exist: {path}")
+ replacements.append(None)
+ continue
+ has_replacement = True
+ replacement, replacement_valid = provider.numpy_values(
+ global_row,
+ chunk_length,
+ None if len(rows) == chunk_length else rows,
+ )
+ absolute = row_starts + pos
+ compatible &= (
+ replacement_valid
+ & (source_data[absolute] == provider._type_header)
+ & (absolute + provider._fixed_size == row_starts + limit)
+ )
+ replacements.append((pos, provider, replacement))
+ if not has_replacement:
+ return
+ slow_rows.update(int(row) for row in rows[~compatible])
+ if not np.any(compatible):
+ return
+ output_data = state.ensure()
+ relative_starts = row_starts - state.data_start
+ compatible_rows = rows[compatible]
+ for item in replacements:
+ if item is None:
+ continue
+ pos, provider, replacement = item
+ absolute = (relative_starts + pos)[compatible]
+ output_data[absolute] = provider._type_header
+ value_size = provider._fixed_size - 1
+ replacement_bytes = np.ascontiguousarray(
+ replacement[compatible]).view(np.uint8).reshape(
+ len(compatible_rows), value_size)
+ indices = absolute[:, None] + 1 + np.arange(value_size)
+ output_data[indices] = replacement_bytes
+
+
def _vectorized_replace_chunk(
chunk,
values,
@@ -1179,57 +1390,13 @@ def _vectorized_replace_chunk(
plans, slow_rows = _partition_path_plans(
values, chunk.field(1), valid_rows, parsed_paths)
slow_rows = set(int(row) for row in slow_rows)
- data = None
- data_start = 0
- output_data = None
+ state = _PatchState(values)
for planned in plans:
- rows, row_starts, source_data, positions, limits = planned
- replacements = []
- compatible = np.ones(len(rows), dtype=bool)
- has_replacement = False
- for (path, _, provider), pos, limit in zip(
- parsed, positions, limits):
- if pos is None:
- if strict:
- raise ValueError(
- f"VARIANT path does not exist: {path}")
- replacements.append(None)
- continue
- has_replacement = True
- replacement, replacement_valid = provider.numpy_values(
- global_row,
- len(chunk),
- None if len(rows) == len(chunk) else rows,
- )
- absolute = row_starts + pos
- compatible &= (
- replacement_valid
- & (source_data[absolute] == provider._type_header)
- & (absolute + provider._fixed_size == row_starts + limit)
- )
- replacements.append((pos, provider, replacement))
- if not has_replacement:
- continue
- slow_rows.update(int(row) for row in rows[~compatible])
- if not np.any(compatible):
- continue
- if data is None:
- data, data_start = values.copy_used_data()
- output_data = np.frombuffer(data, dtype=np.uint8)
- relative_starts = row_starts - data_start
- compatible_rows = rows[compatible]
- for item in replacements:
- if item is None:
- continue
- pos, provider, replacement = item
- absolute = (relative_starts + pos)[compatible]
- output_data[absolute] = provider._type_header
- value_size = provider._fixed_size - 1
- replacement_bytes = np.ascontiguousarray(
- replacement[compatible]).view(np.uint8).reshape(
- len(compatible_rows), value_size)
- indices = absolute[:, None] + 1 + np.arange(value_size)
- output_data[indices] = replacement_bytes
+ _patch_planned_group(
+ planned, parsed, len(chunk), global_row, state,
+ slow_rows, strict)
+ data = state.data
+ data_start = state.data_start
metadata = _BinaryValues(chunk.field(1))
rebuilt_rows = {}
@@ -1278,7 +1445,8 @@ def _vectorized_replace_chunk(
def _supported_replacement_type(data_type: pa.DataType) -> bool:
return (
- pa.types.is_boolean(data_type)
+ pa.types.is_null(data_type)
+ or pa.types.is_boolean(data_type)
or pa.types.is_signed_integer(data_type)
or pa.types.is_float32(data_type)
or pa.types.is_float64(data_type)
@@ -1293,6 +1461,8 @@ def _supported_replacement_type(data_type: pa.DataType)
-> bool:
def _replacement_type_matches(value, pos, data_type):
+ if pa.types.is_null(data_type):
+ return True
variant_type = _variant_get_type(value, pos)
if variant_type == _Type.NULL:
return True
@@ -1395,6 +1565,7 @@ def _variant_get(column, paths: Mapping[str,
pa.DataType]):
}
+@_with_metadata_cache
def variant_get(column, path, data_type=None):
"""Read one or more VARIANT paths without implicit casts."""
if isinstance(path, Mapping):
@@ -1420,6 +1591,7 @@ def _validate_distinct_paths(parsed) -> None:
"VARIANT replacement paths must not overlap")
+@_with_metadata_cache
def variant_replace(
column,
path,
@@ -1473,3 +1645,435 @@ def variant_replace(
if not chunked:
return result_chunks[0]
return pa.chunked_array(result_chunks, type=data_type)
+
+
+class _ValueParts:
+
+ __slots__ = ('parts', 'size')
+
+ def __init__(self, parts, size):
+ _check_variant_sizes(size, 0)
+ self.parts = tuple(parts)
+ self.size = size
+
+ def __len__(self):
+ return self.size
+
+
+def _part_size(part):
+ return part.size if isinstance(part, _ValueParts) else len(part)
+
+
+def _materialize_value(part):
+ if not isinstance(part, _ValueParts):
+ _check_variant_sizes(len(part), 0)
+ return bytes(part)
+ output = bytearray(part.size)
+ output_pos = 0
+ stack = list(reversed(part.parts))
+ while stack:
+ current = stack.pop()
+ if isinstance(current, _ValueParts):
+ stack.extend(reversed(current.parts))
+ continue
+ size = len(current)
+ output[output_pos:output_pos + size] = current
+ output_pos += size
+ return bytes(output)
+
+
+def _build_object_value_parts(fields):
+ """Build an object without copying child values."""
+ size = len(fields)
+ data_size = sum(_part_size(child) for _, child in fields)
+ large_size = size > _U8_MAX
+ size_bytes = _U32_SIZE if large_size else 1
+ max_id = max((field_id for field_id, _ in fields), default=0)
+ id_size = _get_int_size(max_id) if max_id > 0 else 1
+ offset_size = _get_int_size(data_size) if data_size > 0 else 1
+ buf = bytearray()
+ buf.append(_object_header(large_size, id_size, offset_size))
+ buf += size.to_bytes(size_bytes, 'little')
+ for field_id, _ in fields:
+ buf += field_id.to_bytes(id_size, 'little')
+ offset = 0
+ for _, child in fields:
+ buf += offset.to_bytes(offset_size, 'little')
+ offset += _part_size(child)
+ buf += offset.to_bytes(offset_size, 'little')
+ header = bytes(buf)
+ return _ValueParts(
+ [header] + [child for _, child in fields],
+ len(header) + data_size,
+ )
+
+
+def _build_array_value_parts(children):
+ """Build an array without copying child values."""
+ size = len(children)
+ data_size = sum(_part_size(child) for child in children)
+ large_size = size > _U8_MAX
+ size_bytes = _U32_SIZE if large_size else 1
+ offset_size = _get_int_size(data_size) if data_size > 0 else 1
+ buf = bytearray()
+ buf.append(_array_header(large_size, offset_size))
+ buf += size.to_bytes(size_bytes, 'little')
+ offset = 0
+ for child in children:
+ buf += offset.to_bytes(offset_size, 'little')
+ offset += _part_size(child)
+ buf += offset.to_bytes(offset_size, 'little')
+ header = bytes(buf)
+ return _ValueParts([header] + children, len(header) + data_size)
+
+
+def _build_object_value_ordered(fields):
+ """Build object value bytes keeping the given field order."""
+ return _materialize_value(_build_object_value_parts(fields))
+
+
+def _apply_edits(
+ value,
+ pos,
+ limit,
+ edits,
+ key_ids,
+ names_by_id,
+ source_metadata_size=None,
+):
+ """Apply edits and validate source ids before metadata extension."""
+ source = value if isinstance(value, memoryview) else memoryview(value)
+ results = {}
+ next_token = 1
+ stack = [('visit', 0, pos, limit, edits)]
+ while stack:
+ action = stack.pop()
+ kind = action[0]
+ if kind == 'finish_object':
+ _, token, ids, children, inserts, child_tokens = action
+ for slot, child_token in child_tokens:
+ children[slot] = results.pop(child_token)
+ fields = list(zip(ids, children)) + inserts
+ if len({field_id for field_id, _ in fields}) != len(fields):
+ _malformed("duplicate object field id")
+ if inserts:
+ try:
+ fields.sort(
+ key=lambda field: names_by_id[
+ field[0]].encode('utf-8'))
+ except KeyError:
+ _malformed("object key is missing from metadata")
+ results[token] = _build_object_value_parts(fields)
+ continue
+ if kind == 'finish_array':
+ _, token, children, child_tokens = action
+ for index, child_token in child_tokens:
+ children[index] = results.pop(child_token)
+ results[token] = _build_array_value_parts(children)
+ continue
+
+ _, token, node_pos, node_limit, node_edits = action
+ value_end = node_pos + _checked_value_size(
+ value, node_pos, node_limit)
+ if value_end != node_limit:
+ _malformed("child size does not match container offsets")
+ inserts = []
+ descend = {}
+ replacement = None
+ for segments, op, key_id, payload in node_edits:
+ if op == 'replace' and not segments:
+ replacement = payload
+ break
+ if op == 'insert' and len(segments) == 1:
+ inserts.append((key_id, payload))
+ else:
+ descend.setdefault(segments[0], []).append(
+ (segments[1:], op, key_id, payload))
+ if replacement is not None:
+ results[token] = replacement
+ continue
+
+ basic_type = value[node_pos] & 0x3
+ child_actions = []
+ if basic_type == _OBJECT:
+ size, id_size, id_start, data_start, offsets, _ = (
+ _checked_object_layout(value, node_pos, value_end))
+ ids = [
+ _read_unsigned(value, id_start + i * id_size, id_size)
+ for i in range(size)
+ ]
+ if (source_metadata_size is not None
+ and any(field_id >= source_metadata_size
+ for field_id in ids)):
+ _malformed("object field id is missing from metadata")
+ ordered_offsets = sorted(offsets)
+ end_by_offset = dict(zip(
+ ordered_offsets, ordered_offsets[1:]))
+ slot_by_id = {
+ field_id: index for index, field_id in enumerate(ids)
+ }
+ edits_by_slot = {}
+ for (_, segment), child_edits in descend.items():
+ slot = slot_by_id[key_ids[segment]]
+ edits_by_slot[slot] = child_edits
+ children = []
+ child_tokens = []
+ for slot in range(size):
+ child_pos, child_end = _checked_object_child_bounds(
+ value, data_start, offsets, slot, end_by_offset)
+ child_edits = edits_by_slot.get(slot)
+ if child_edits is not None:
+ child_token = next_token
+ next_token += 1
+ children.append(None)
+ child_tokens.append((slot, child_token))
+ child_actions.append((
+ 'visit', child_token, child_pos, child_end,
+ child_edits,
+ ))
+ else:
+ if (source_metadata_size is not None
+ and (value[child_pos] & 0x3)
+ in (_OBJECT, _ARRAY)):
+ _validate_value_field_ids(
+ value, child_pos, child_end,
+ source_metadata_size)
+ children.append(source[child_pos:child_end])
+ stack.append((
+ 'finish_object', token, ids, children, inserts,
+ child_tokens,
+ ))
+ elif basic_type == _ARRAY:
+ size, data_start, offsets, _ = _checked_array_layout(
+ value, node_pos, value_end)
+ edits_by_index = {
+ segment: child_edits
+ for (_, segment), child_edits in descend.items()
+ }
+ children = []
+ child_tokens = []
+ for index in range(size):
+ child_pos = data_start + offsets[index]
+ child_end = data_start + offsets[index + 1]
+ child_edits = edits_by_index.get(index)
+ if child_edits is not None:
+ child_token = next_token
+ next_token += 1
+ children.append(None)
+ child_tokens.append((index, child_token))
+ child_actions.append((
+ 'visit', child_token, child_pos, child_end,
+ child_edits,
+ ))
+ else:
+ if source_metadata_size is not None:
+ _validate_value_field_ids(
+ value, child_pos, child_end,
+ source_metadata_size)
+ children.append(source[child_pos:child_end])
+ stack.append((
+ 'finish_array', token, children, child_tokens,
+ ))
+ else:
+ _malformed("path segment does not match the value type")
+ stack.extend(child_actions)
+ return _materialize_value(results[0])
+
+
+def _set_chunk(chunk, values, parsed, global_row):
+ parsed_paths = [parsed_path for _, parsed_path, _ in parsed]
+ parent_paths = [parsed_path[:-1] for parsed_path in parsed_paths]
+ query_paths = tuple(parsed_paths) + tuple(parent_paths)
+ count = len(parsed)
+ metadata_column = chunk.field(1)
+ valid_rows = _valid_row_indices(chunk, values, metadata_column)
+ if not len(valid_rows):
+ return chunk
+ plans, slow_rows = _partition_path_plans(
+ values, metadata_column, valid_rows, query_paths)
+ slow_rows = set(int(row) for row in slow_rows)
+ metadata_values = _BinaryValues(metadata_column)
+ state = _PatchState(values)
+ rebuilt_rows = {}
+ rebuilt_metadata = {}
+ scalar_payloads = {}
+
+ def payload_for(index, provider, row):
+ if provider._array is not None:
+ return provider.encode(provider.scalar_at(global_row + row))
+ if index not in scalar_payloads:
+ scalar_payloads[index] = provider.encode(provider.scalar_at(0))
+ return scalar_payloads[index]
+
+ def rebuild_row(row, view, insert_set, key_ids, names_by_id,
+ original_metadata, new_metadata,
+ source_metadata_size=None,
+ validated_positions=None):
+ edits = []
+ for index, (path, parsed_path, provider) in enumerate(parsed):
+ payload = payload_for(index, provider, row)
+ if index in insert_set:
+ edits.append((
+ parsed_path, 'insert',
+ key_ids[parsed_path[-1][1]], payload))
+ else:
+ if validated_positions is not None:
+ provider.validate_source(
+ view, validated_positions[index])
+ edits.append((parsed_path, 'replace', None, payload))
+ rebuilt = _apply_edits(
+ view, 0, len(view), edits, key_ids, names_by_id,
+ source_metadata_size)
+ _check_variant_sizes(
+ len(rebuilt),
+ len(new_metadata if new_metadata is not None
+ else original_metadata),
+ )
+ if new_metadata is not None or rebuilt != view:
+ rebuilt_rows[row] = rebuilt
+ if new_metadata is not None:
+ rebuilt_metadata[row] = new_metadata
+
+ for planned in plans:
+ rows, row_starts, source_data, positions, limits = planned
+ target_positions = positions[:count]
+ target_limits = limits[:count]
+ parent_positions = positions[count:]
+ insert_indices = []
+ for index, (path, parsed_path, provider) in enumerate(parsed):
+ if target_positions[index] is not None:
+ continue
+ parent_pos = parent_positions[index]
+ if parent_pos is None:
+ raise ValueError(
+ f"VARIANT parent path does not exist: {path}")
+ if not parsed_path or parsed_path[-1][0] != 'key':
+ raise ValueError(
+ "VARIANT array index insertion is not supported: "
+ + path)
+ parent_headers = source_data[row_starts + parent_pos]
+ if np.any((parent_headers & 0x3) != _OBJECT):
+ raise ValueError(
+ f"VARIANT parent path is not an object: {path}")
+ insert_indices.append(index)
+ if not insert_indices and all(
+ provider._fixed_size is not None
+ for _, _, provider in parsed):
+ _patch_planned_group(
+ (rows, row_starts, source_data,
+ target_positions, target_limits),
+ parsed, len(chunk), global_row, state, slow_rows, False)
+ continue
+ first_metadata = bytes(metadata_values.view(int(rows[0])))
+ insert_keys = tuple(
+ parsed[index][1][-1][1] for index in insert_indices)
+ metadata_key_ids = _cached_metadata_key_ids(first_metadata)
+ source_metadata_size = (
+ len(metadata_key_ids)
+ if any(key not in metadata_key_ids for key in insert_keys)
+ else None
+ )
+ new_metadata, key_ids, names_by_id = _metadata_with_keys(
+ first_metadata, insert_keys)
+ insert_set = set(insert_indices)
+ for offset_index, row in enumerate(rows):
+ row = int(row)
+ rebuild_row(
+ row, values.view(row), insert_set, key_ids, names_by_id,
+ first_metadata, new_metadata, source_metadata_size,
+ [
+ None if target_positions[index] is None
+ else int(target_positions[index][offset_index])
+ for index in range(count)
+ ],
+ )
+
+ for row in sorted(slow_rows):
+ view = values.view(row)
+ row_metadata = bytes(metadata_values.view(row))
+ positions = _path_positions(view, row_metadata, query_paths)
+ target_positions = positions[:count]
+ parent_positions = positions[count:]
+ insert_keys = []
+ insert_set = set()
+ for index, (path, parsed_path, provider) in enumerate(parsed):
+ if target_positions[index] is not None:
+ provider.validate_source(view, target_positions[index])
+ continue
+ parent_pos = parent_positions[index]
+ if parent_pos is None:
+ raise ValueError(
+ f"VARIANT parent path does not exist: {path}")
+ if not parsed_path or parsed_path[-1][0] != 'key':
+ raise ValueError(
+ "VARIANT array index insertion is not supported: "
+ + path)
+ if (view[parent_pos] & 0x3) != _OBJECT:
+ raise ValueError(
+ f"VARIANT parent path is not an object: {path}")
+ insert_set.add(index)
+ insert_keys.append(parsed_path[-1][1])
+ metadata_key_ids = _cached_metadata_key_ids(row_metadata)
+ source_metadata_size = (
+ len(metadata_key_ids)
+ if any(key not in metadata_key_ids for key in insert_keys)
+ else None
+ )
+ new_metadata, key_ids, names_by_id = _metadata_with_keys(
+ row_metadata, tuple(insert_keys))
+ rebuild_row(
+ row, view, insert_set, key_ids, names_by_id, row_metadata,
+ new_metadata, source_metadata_size)
+
+ if rebuilt_rows:
+ state.ensure()
+ new_values = _sparse_rebuilt_binary(
+ values, state.data, state.data_start, rebuilt_rows)
+ elif state.data is not None:
+ return _patched_chunk(chunk, values, state.data, state.data_start)
+ else:
+ return chunk
+ if rebuilt_metadata:
+ new_metadata_column = _sparse_rebuilt_binary(
+ metadata_values, metadata_values.data, 0, rebuilt_metadata)
+ else:
+ new_metadata_column = metadata_column
+ return pa.StructArray.from_arrays(
+ [new_values, new_metadata_column],
+ fields=list(chunk.type),
+ mask=chunk.is_null(),
+ )
+
+
+@_with_metadata_cache
+def variant_set(column, path, value=None):
+ """Replace existing VARIANT paths or insert missing final OBJECT keys.
+
+ Missing intermediate paths and non-OBJECT parents raise ValueError.
+ """
+ if isinstance(path, Mapping):
+ if value is not None:
+ raise TypeError(
+ "VARIANT value must be omitted for path mappings")
+ updates = path
+ else:
+ updates = {path: value}
+ parsed = [
+ (target, _parse_path(target), _Replacement(item, len(column)))
+ for target, item in updates.items()
+ ]
+ _validate_distinct_paths(parsed)
+ if not parsed:
+ return column
+
+ chunks, chunked, data_type = _variant_chunks(column)
+ result_chunks = []
+ global_row = 0
+ for chunk in chunks:
+ values = _BinaryValues(chunk.field(0))
+ result_chunks.append(_set_chunk(chunk, values, parsed, global_row))
+ global_row += len(chunk)
+
+ if not chunked:
+ return result_chunks[0]
+ return pa.chunked_array(result_chunks, type=data_type)
diff --git a/paimon-python/pypaimon/tests/variant_path_test.py
b/paimon-python/pypaimon/tests/variant_path_test.py
index 9ddf75c700..b3a569ea68 100644
--- a/paimon-python/pypaimon/tests/variant_path_test.py
+++ b/paimon-python/pypaimon/tests/variant_path_test.py
@@ -26,6 +26,8 @@ import pyarrow.compute as pc
from pypaimon.data._variant_binary import _primitive_header
from pypaimon.data.generic_variant import _DOUBLE, GenericVariant
from pypaimon.data.variant_path import (
+ _compile_paths,
+ _metadata_cache,
_metadata_key_ids,
_path_positions,
_rebuilt_offsets,
@@ -77,6 +79,45 @@ def _typed_object(fields):
class TestVariantGet(unittest.TestCase):
+ def test_compile_paths_builds_trie_without_prefix_slices(self):
+ class NoSlicePath(tuple):
+ def __getitem__(self, item):
+ if isinstance(item, slice):
+ raise AssertionError("path prefix was materialized")
+ return super().__getitem__(item)
+
+ paths = (
+ NoSlicePath((('key', 'root'), ('index', 0), ('key', 'left'))),
+ NoSlicePath((('key', 'root'), ('index', 0), ('key', 'right'))),
+ )
+
+ nodes, results = _compile_paths(paths)
+
+ self.assertEqual(len(nodes), 5)
+ self.assertEqual(results, (3, 4))
+ self.assertEqual(nodes[3], (2, 'key', 'left'))
+ self.assertEqual(nodes[4], (2, 'key', 'right'))
+
+ def test_metadata_cache_is_bounded_and_released(self):
+ column = _variants([
+ {'value': float(index), 'key_%d' % index: index}
+ for index in range(300)
+ ])
+ cache_sizes = []
+
+ def parse_metadata(metadata):
+ cache_sizes.append(len(_metadata_cache.value))
+ return _metadata_key_ids(metadata)
+
+ with patch(
+ 'pypaimon.data.variant_path._metadata_key_ids',
+ side_effect=parse_metadata):
+ result = variant_get(column, '$.value', pa.float64())
+
+ self.assertEqual(result.to_pylist(), [float(i) for i in range(300)])
+ self.assertLessEqual(max(cache_sizes), 256)
+ self.assertFalse(hasattr(_metadata_cache, 'value'))
+
def test_nested_paths_and_missing_values(self):
column = pa.chunked_array([
_variants([{'a.b': [{'value': 1.5}]}, None]),
diff --git a/paimon-python/pypaimon/tests/variant_set_test.py
b/paimon-python/pypaimon/tests/variant_set_test.py
new file mode 100644
index 0000000000..45352e1ef8
--- /dev/null
+++ b/paimon-python/pypaimon/tests/variant_set_test.py
@@ -0,0 +1,898 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+from unittest.mock import patch
+
+import numpy as np
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pypaimon.data import variant_replace, variant_set
+from pypaimon.data.generic_variant import GenericVariant, _check_variant_sizes
+from pypaimon.data.variant_path import (
+ _apply_edits,
+ _build_object_value_ordered,
+ _materialize_value,
+ _metadata_key_ids,
+ _metadata_with_keys,
+ _path_positions,
+ _rebuilt_offsets,
+ _validate_value_field_ids,
+ variant_get,
+)
+from pypaimon.data.variant_shredding import (
+ _build_array_value,
+ _build_object_value,
+ _encode_scalar_to_value_bytes,
+)
+
+# Bytes built by the Java GenericVariantBuilder for
+# {"angular_velocity":{"y":1.5,"z":-2.5},
+# "linear_acceleration":{"y":0.25,"z":4.0},"processed":true,"seq":7}.
+_JAVA_VALUE = bytes.fromhex(
+ '0204000304050019323335020201020009121c000000000000f83f1c000000000000'
+ '04c0020201020009121c000000000000d03f1c0000000000001040040c07')
+_JAVA_METADATA = bytes.fromhex(
+ '010600101112252e31616e67756c61725f76656c6f63697479797a6c696e6561725f'
+ '616363656c65726174696f6e70726f636573736564736571')
+_JAVA_PYTHON_VALUE = {
+ 'angular_velocity': {'y': 1.5, 'z': -2.5},
+ 'linear_acceleration': {'y': 0.25, 'z': 4.0},
+ 'processed': True,
+ 'seq': 7,
+}
+
+
+def _variants(values):
+ return GenericVariant.to_arrow_array([
+ GenericVariant.from_python(value) if value is not None else None
+ for value in values
+ ])
+
+
+def _decode(column):
+ return [
+ None if value is None
+ else GenericVariant.from_arrow_struct(value).to_python()
+ for value in column.to_pylist()
+ ]
+
+
+def _sensor_rows(count, offset=0):
+ return [
+ {
+ 'angular_velocity': {
+ 'y': float(index + offset),
+ 'z': float(index + offset) + 0.5,
+ },
+ 'linear_acceleration': {
+ 'y': -float(index + offset),
+ 'z': -float(index + offset) - 0.5,
+ },
+ }
+ for index in range(count)
+ ]
+
+
+_SENSOR_PATHS = (
+ '$.angular_velocity.y',
+ '$.angular_velocity.z',
+ '$.linear_acceleration.y',
+ '$.linear_acceleration.z',
+)
+
+
+class TestVariantSetReplace(unittest.TestCase):
+
+ def test_existing_paths_match_variant_replace(self):
+ column = _variants(_sensor_rows(100) + [None])
+ current = variant_get(
+ column, {path: pa.float64() for path in _SENSOR_PATHS})
+ updates = {
+ path: pc.negate(values) for path, values in current.items()
+ }
+
+ self.assertTrue(
+ variant_set(column, updates).equals(
+ variant_replace(column, updates)))
+
+ def test_negates_four_double_paths(self):
+ rows = _sensor_rows(50)
+ column = _variants(rows)
+ current = variant_get(
+ column, {path: pa.float64() for path in _SENSOR_PATHS})
+
+ result = variant_set(column, {
+ path: pc.negate(values) for path, values in current.items()
+ })
+
+ for row, decoded in zip(rows, _decode(result)):
+ self.assertEqual(decoded, {
+ 'angular_velocity': {
+ 'y': -row['angular_velocity']['y'],
+ 'z': -row['angular_velocity']['z'],
+ },
+ 'linear_acceleration': {
+ 'y': -row['linear_acceleration']['y'],
+ 'z': -row['linear_acceleration']['z'],
+ },
+ })
+
+ def test_replaces_root_path(self):
+ column = _variants([1.5, -2.5])
+
+ result = variant_set(column, '$', pa.scalar(3.5, type=pa.float64()))
+
+ self.assertEqual(
+ variant_get(result, '$', pa.float64()).to_pylist(), [3.5, 3.5])
+
+
+class TestVariantSetInsert(unittest.TestCase):
+
+ def test_inserts_bool_and_string_marks(self):
+ column = _variants([{'value': 1.0}, {'value': 2.0}])
+
+ flagged = variant_set(column, '$.processed', pa.scalar(True))
+ tagged = variant_set(column, '$.tag', pa.scalar('done'))
+
+ self.assertEqual(_decode(flagged), [
+ {'value': 1.0, 'processed': True},
+ {'value': 2.0, 'processed': True},
+ ])
+ self.assertEqual(_decode(tagged), [
+ {'value': 1.0, 'tag': 'done'},
+ {'value': 2.0, 'tag': 'done'},
+ ])
+
+ def test_insert_extends_metadata_dictionary(self):
+ column = _variants([{'value': 1.0}])
+
+ result = variant_set(column, '$.processed', pa.scalar(True))
+
+ metadata = result.to_pylist()[0]['metadata']
+ self.assertEqual(
+ _metadata_key_ids(metadata), {'value': 0, 'processed': 1})
+
+ def test_insert_reuses_metadata_key_and_buffer(self):
+ metadata = GenericVariant.from_python(
+ {'value': 0, 'flag': 0}).metadata()
+ key_ids = _metadata_key_ids(metadata)
+ value = _build_object_value([
+ (key_ids['value'],
+ _encode_scalar_to_value_bytes(1.5, pa.float64())),
+ ])
+ column = GenericVariant.to_arrow_array(
+ [GenericVariant(value, metadata)] * 100)
+
+ result = variant_set(column, '$.flag', pa.scalar(False))
+
+ self.assertEqual(
+ _decode(result), [{'value': 1.5, 'flag': False}] * 100)
+ self.assertEqual(
+ result.to_pylist()[0]['metadata'], metadata)
+ self.assertEqual(
+ column.field('metadata').buffers()[2].address,
+ result.field('metadata').buffers()[2].address,
+ )
+
+ def test_insert_into_empty_object(self):
+ column = _variants([{}])
+
+ result = variant_set(column, '$.first', pa.scalar(7, pa.int64()))
+
+ self.assertEqual(_decode(result), [{'first': 7}])
+
+ def test_inserts_same_key_into_two_objects(self):
+ column = _variants([{'left': {}, 'right': {}}])
+
+ result = variant_set(column, {
+ '$.left.flag': pa.scalar(True),
+ '$.right.flag': pa.scalar(True),
+ })
+
+ metadata = result.to_pylist()[0]['metadata']
+ self.assertEqual(
+ _metadata_key_ids(metadata), {'left': 0, 'right': 1, 'flag': 2})
+ self.assertEqual(_decode(result), [
+ {'left': {'flag': True}, 'right': {'flag': True}}])
+ self.assertEqual(
+ variant_get(result, '$.left.flag', pa.bool_()).to_pylist(),
+ [True])
+
+ def test_inserted_fields_stay_sorted_for_java_binary_search(self):
+ payload = {'k%02d' % index: float(index) for index in range(40)}
+ column = _variants([payload])
+
+ result = variant_set(column, '$.a_mark', pa.scalar('inserted'))
+
+ decoded = _decode(result)[0]
+ expected = dict(payload)
+ expected['a_mark'] = 'inserted'
+ self.assertEqual(decoded, expected)
+ self.assertEqual(list(decoded), sorted(decoded))
+
+ def test_inserted_fields_stay_sorted_by_utf8(self):
+ key_sup = chr(0x10000)
+ key_bmp = chr(0xE000)
+ payload = {key_bmp: 1.0, key_sup: 2.0}
+ payload.update({'k%02d' % i: float(i) for i in range(40)})
+
+ result = variant_set(_variants([payload]), '$.aaa', pa.scalar(3.0))
+
+ decoded = _decode(result)[0]
+ self.assertEqual(decoded[key_sup], 2.0)
+ self.assertEqual(decoded[key_bmp], 1.0)
+ self.assertEqual(list(decoded), sorted(
+ list(payload.keys()) + ['aaa'],
+ key=lambda name: name.encode('utf-8')))
+
+ def test_mixed_rows_in_one_chunk(self):
+ metadata = GenericVariant.from_python({'a': 0, 'b': 0}).metadata()
+ key_ids = _metadata_key_ids(metadata)
+ reversed_fields = _build_object_value_ordered([
+ (key_ids['b'], _encode_scalar_to_value_bytes(2.0, pa.float64())),
+ (key_ids['a'], _encode_scalar_to_value_bytes(1.0, pa.float64())),
+ ])
+ column = pa.concat_arrays([
+ _variants([{'a': 1.0}, {'a': 1.0, 'mark': 'old'}]),
+ GenericVariant.to_arrow_array(
+ [GenericVariant(reversed_fields, metadata)]),
+ ])
+
+ result = variant_set(column, '$.mark', pa.scalar('new'))
+
+ self.assertEqual(_decode(result), [
+ {'a': 1.0, 'mark': 'new'},
+ {'a': 1.0, 'mark': 'new'},
+ {'b': 2.0, 'a': 1.0, 'mark': 'new'},
+ ])
+
+ def test_replace_and_insert_multiple_paths(self):
+ column = _variants(_sensor_rows(10))
+ current = variant_get(
+ column, {path: pa.float64() for path in _SENSOR_PATHS})
+ updates = {
+ path: pc.negate(values) for path, values in current.items()
+ }
+ updates['$.processed'] = pa.scalar(True, type=pa.bool_())
+
+ result = variant_set(column, updates)
+
+ decoded = _decode(result)
+ self.assertTrue(all(row['processed'] is True for row in decoded))
+ self.assertEqual(
+ [row['angular_velocity']['y'] for row in decoded],
+ [-float(index) for index in range(10)],
+ )
+
+ def test_scalar_array_and_chunked_replacements(self):
+ column = pa.chunked_array([
+ _variants([{'value': 1.0}, {'value': 2.0}]),
+ _variants([{'value': 3.0}]),
+ ])
+
+ result = variant_set(column, {
+ '$.value': pa.chunked_array(
+ [[10.0, 20.0], [30.0]], type=pa.float64()),
+ '$.rank': pa.array([1, 2, 3], type=pa.int64()),
+ '$.processed': pa.scalar(True),
+ })
+
+ self.assertIsInstance(result, pa.ChunkedArray)
+ self.assertEqual(result.num_chunks, 2)
+ self.assertEqual(_decode(result), [
+ {'value': 10.0, 'rank': 1, 'processed': True},
+ {'value': 20.0, 'rank': 2, 'processed': True},
+ {'value': 30.0, 'rank': 3, 'processed': True},
+ ])
+
+
+class TestVariantSetNullSemantics(unittest.TestCase):
+
+ def test_sql_null_rows_stay_null(self):
+ column = _variants([None, {'value': 1.0}])
+
+ result = variant_set(column, '$.processed', pa.scalar(True))
+
+ self.assertEqual(_decode(result), [
+ None, {'value': 1.0, 'processed': True},
+ ])
+ self.assertTrue(result.is_null()[0].as_py())
+
+ def test_arrow_null_becomes_variant_null(self):
+ column = _variants([{'value': 1.0}, {'value': 2.0}])
+
+ result = variant_set(column, {
+ '$.value': pa.array([None, -2.0], type=pa.float64()),
+ '$.mark': pa.array([None, 'done'], type=pa.string()),
+ })
+
+ self.assertEqual(_decode(result), [
+ {'value': None, 'mark': None},
+ {'value': -2.0, 'mark': 'done'},
+ ])
+ self.assertEqual(result.null_count, 0)
+
+ def test_untyped_arrow_null_becomes_variant_null(self):
+ replacements = [
+ pa.scalar(None),
+ pa.nulls(2),
+ pa.chunked_array([pa.nulls(1), pa.nulls(1)]),
+ ]
+ for replacement in replacements:
+ with self.subTest(replacement=type(replacement).__name__):
+ result = variant_set(
+ _variants([{'value': 1.0}, {'value': 2.0}]),
+ '$.value',
+ replacement,
+ )
+ self.assertEqual(_decode(result), [
+ {'value': None}, {'value': None},
+ ])
+
+ def test_variant_null_parent_is_not_an_object(self):
+ column = _variants([{'parent': None}])
+
+ with self.assertRaisesRegex(ValueError, "is not an object"):
+ variant_set(column, '$.parent.child', pa.scalar(1.0))
+
+ def test_missing_intermediate_parent_fails(self):
+ column = _variants([{'other': 1.0}] * 100)
+
+ with self.assertRaisesRegex(ValueError, "parent path does not"):
+ variant_set(column, '$.missing.child', pa.scalar(1.0))
+
+ def test_non_object_parent_fails(self):
+ column = _variants([{'value': 1.0}] * 100)
+
+ with self.assertRaisesRegex(ValueError, "is not an object"):
+ variant_set(column, '$.value.child', pa.scalar(1.0))
+
+ def test_replaces_array_element_of_a_different_size(self):
+ column = _variants([{'items': ['aa', 'bb'], 'n': 1.0}])
+
+ result = variant_set(column, '$.items[0]', pa.scalar('cccc'))
+
+ self.assertEqual(
+ _decode(result), [{'items': ['cccc', 'bb'], 'n': 1.0}])
+
+ def test_array_insertion_is_not_supported(self):
+ column = _variants([{'items': [1.0]}])
+
+ with self.assertRaisesRegex(ValueError, "not supported"):
+ variant_set(column, '$.items[3]', pa.scalar(1.0))
+ result = variant_set(column, '$.items[0]', pa.scalar(-1.0))
+ self.assertEqual(_decode(result), [{'items': [-1.0]}])
+
+
+class TestVariantSetLayouts(unittest.TestCase):
+
+ def test_sliced_input(self):
+ base = _variants([
+ {'value': float(index), 'padding': 'x' * 100}
+ for index in range(100)
+ ])
+ column = base.slice(50, 3)
+
+ result = variant_set(column, {
+ '$.value': pa.scalar(-1.0),
+ '$.processed': pa.scalar(True),
+ })
+
+ self.assertEqual(
+ [(row['value'], row['processed']) for row in _decode(result)],
+ [(-1.0, True)] * 3,
+ )
+
+ def test_large_binary_input(self):
+ column = _variants([{'value': 1.0}])
+ large = pa.StructArray.from_arrays(
+ [
+ column.field('value').cast(pa.large_binary()),
+ column.field('metadata').cast(pa.large_binary()),
+ ],
+ names=['value', 'metadata'],
+ )
+
+ result = variant_set(large, '$.processed', pa.scalar(True))
+
+ self.assertTrue(pa.types.is_large_binary(result.type[0].type))
+ self.assertTrue(pa.types.is_large_binary(result.type[1].type))
+ self.assertEqual(
+ _decode(result), [{'value': 1.0, 'processed': True}])
+
+ def test_preserves_chunk_boundaries_without_combine(self):
+ column = pa.chunked_array([
+ _variants([{'value': 1.0}]),
+ _variants([{'value': 2.0}, {'value': 3.0}]),
+ ])
+
+ with patch(
+ 'pypaimon.data.variant_path._rebuilt_offsets',
+ wraps=_rebuilt_offsets,
+ ) as rebuilt_offsets:
+ result = variant_set(column, '$.processed', pa.scalar(True))
+
+ self.assertEqual(
+ [len(chunk) for chunk in result.chunks],
+ [len(chunk) for chunk in column.chunks],
+ )
+ # Offsets are rebuilt per chunk, never for the combined column.
+ self.assertTrue(rebuilt_offsets.called)
+ self.assertEqual(
+ max(len(call[0][0])
+ for call in rebuilt_offsets.call_args_list),
+ 2,
+ )
+ self.assertEqual(_decode(result), [
+ {'value': 1.0, 'processed': True},
+ {'value': 2.0, 'processed': True},
+ {'value': 3.0, 'processed': True},
+ ])
+
+ def test_offset_overflow_guard_is_low_memory(self):
+ lengths = np.array([(1 << 31) - 8, 16], dtype=np.int64)
+
+ with self.assertRaisesRegex(ValueError, "use LargeBinary"):
+ _rebuilt_offsets(lengths, '<i')
+ self.assertEqual(
+ _rebuilt_offsets(lengths, '<q')[-1], (1 << 31) + 8)
+
+ def test_input_is_not_modified(self):
+ column = _variants([{'value': 1.0}, None, {'value': 2.0}])
+ original_rows = column.to_pylist()
+ original_value = column.field('value').buffers()[2].to_pybytes()
+ original_metadata = (
+ column.field('metadata').buffers()[2].to_pybytes())
+
+ variant_set(column, {
+ '$.value': pa.scalar(-1.0),
+ '$.processed': pa.scalar(True),
+ })
+
+ self.assertEqual(column.to_pylist(), original_rows)
+ self.assertEqual(
+ column.field('value').buffers()[2].to_pybytes(),
+ original_value,
+ )
+ self.assertEqual(
+ column.field('metadata').buffers()[2].to_pybytes(),
+ original_metadata,
+ )
+
+
+class TestVariantSetFastPaths(unittest.TestCase):
+
+ def test_replace_avoids_full_decode(self):
+ column = _variants([{'value': float(index)} for index in range(100)])
+
+ with patch.object(
+ GenericVariant, 'to_python',
+ side_effect=AssertionError("full decode is not allowed")), \
+ patch.object(
+ GenericVariant, 'from_python',
+ side_effect=AssertionError(
+ "full encode is not allowed")):
+ result = variant_set(column, '$.value', pa.scalar(-1.0))
+
+ self.assertEqual(
+ variant_get(result, '$.value', pa.float64()).to_pylist(),
+ [-1.0] * 100,
+ )
+
+ def test_insert_avoids_full_decode(self):
+ column = _variants([{'value': float(index)} for index in range(100)])
+
+ with patch.object(
+ GenericVariant, 'to_python',
+ side_effect=AssertionError("full decode is not allowed")), \
+ patch.object(
+ GenericVariant, 'from_python',
+ side_effect=AssertionError(
+ "full encode is not allowed")):
+ result = variant_set(column, '$.processed', pa.scalar(True))
+
+ self.assertEqual(
+ variant_get(result, '$.processed', pa.bool_()).to_pylist(),
+ [True] * 100,
+ )
+
+ def test_replace_fast_path_stays_vectorized(self):
+ column = _variants(
+ [{'value': float(index)} for index in range(4096)])
+
+ with patch(
+ 'pypaimon.data.variant_path._path_positions',
+ wraps=_path_positions,
+ ) as slow_path, patch(
+ 'pypaimon.data.variant_path._apply_edits',
+ wraps=_apply_edits,
+ ) as rebuild:
+ result = variant_set(column, '$.value', pa.scalar(-1.0))
+
+ slow_path.assert_not_called()
+ rebuild.assert_not_called()
+ self.assertEqual(
+ variant_get(result, '$.value', pa.float64()).to_pylist(),
+ [-1.0] * 4096,
+ )
+
+ def test_insert_avoids_per_row_planning(self):
+ column = _variants(
+ [{'value': float(index)} for index in range(4096)])
+
+ with patch(
+ 'pypaimon.data.variant_path._path_positions',
+ wraps=_path_positions,
+ ) as slow_path, patch(
+ 'pypaimon.data.variant_path._metadata_key_ids',
+ wraps=_metadata_key_ids,
+ ) as metadata_parse:
+ result = variant_set(column, '$.processed', pa.scalar(True))
+
+ slow_path.assert_not_called()
+ self.assertLessEqual(metadata_parse.call_count, 2)
+ self.assertEqual(
+ variant_get(result, '$.processed', pa.bool_()).to_pylist(),
+ [True] * 4096,
+ )
+
+ def test_insert_fuses_root_validation_with_rebuild(self):
+ column = _variants([
+ {'nested': {'value': float(index)}, 'other': float(index)}
+ for index in range(100)
+ ])
+
+ with patch(
+ 'pypaimon.data.variant_path._validate_value_field_ids',
+ wraps=_validate_value_field_ids,
+ ) as subtree_validation:
+ result = variant_set(column, '$.processed', pa.scalar(True))
+
+ self.assertFalse(any(
+ args[1] == 0
+ for args, _ in subtree_validation.call_args_list
+ ))
+ self.assertEqual(
+ variant_get(result, '$.processed', pa.bool_()).to_pylist(),
+ [True] * 100,
+ )
+
+ def test_insert_validates_deep_unmodified_sibling_iteratively(self):
+ metadata = GenericVariant.from_python(
+ {'sibling': [], 'target': {}}).metadata()
+ key_ids = _metadata_key_ids(metadata)
+ sibling = _encode_scalar_to_value_bytes(1.0, pa.float64())
+ for _ in range(1020):
+ sibling = _build_array_value([sibling])
+ root = _build_object_value([
+ (key_ids['sibling'], sibling),
+ (key_ids['target'], _build_object_value([])),
+ ])
+ column = GenericVariant.to_arrow_array([
+ GenericVariant(root, metadata),
+ ])
+
+ result = variant_set(column, '$.target.new', pa.scalar(True))
+
+ self.assertEqual(
+ variant_get(result, '$.target.new', pa.bool_()).to_pylist(),
+ [True],
+ )
+
+ def test_insert_rebuilds_deep_modified_path_iteratively(self):
+ metadata = GenericVariant.from_python({'target': {}}).metadata()
+ key_ids = _metadata_key_ids(metadata)
+ target = _build_object_value([])
+ for _ in range(1020):
+ target = _build_array_value([target])
+ root = _build_object_value([
+ (key_ids['target'], target),
+ ])
+ column = GenericVariant.to_arrow_array([
+ GenericVariant(root, metadata),
+ ])
+ path = '$.target' + '[0]' * 1020 + '.new'
+
+ with patch(
+ 'pypaimon.data.variant_path._materialize_value',
+ wraps=_materialize_value,
+ ) as materialize:
+ result = variant_set(column, path, pa.scalar(True))
+
+ self.assertEqual(materialize.call_count, 1)
+ self.assertEqual(
+ variant_get(result, path, pa.bool_()).to_pylist(),
+ [True],
+ )
+
+ def test_insert_offset_width_boundary_mixed_rows(self):
+ # Rows crossing the 1-byte offset limit after the insert must be
+ # rebuilt with a wider offset table inside the same plan group.
+ rows = []
+ for index in range(100):
+ padding = 'x' * (240 if index % 3 == 0 else 10)
+ rows.append({'value': float(index), 'padding': padding})
+ column = _variants(rows)
+ mark = 'm' * 30
+
+ result = variant_set(column, '$.mark', pa.scalar(mark))
+
+ for index, decoded in enumerate(_decode(result)):
+ self.assertEqual(decoded, {
+ 'value': float(index),
+ 'padding': rows[index]['padding'],
+ 'mark': mark,
+ })
+
+
+class TestVariantSetErrors(unittest.TestCase):
+
+ def test_variant_size_limit_boundary(self):
+ with patch('pypaimon.data.generic_variant._SIZE_LIMIT', 64):
+ _check_variant_sizes(64, 64)
+ with self.assertRaisesRegex(
+ ValueError, 'VARIANT_CONSTRUCTOR_SIZE_LIMIT'):
+ _check_variant_sizes(65, 64)
+ with self.assertRaisesRegex(
+ ValueError, 'VARIANT_CONSTRUCTOR_SIZE_LIMIT'):
+ _check_variant_sizes(64, 65)
+
+ def test_rejects_oversized_value_and_metadata(self):
+ column = _variants([{'value': 'a'}])
+ with patch('pypaimon.data.generic_variant._SIZE_LIMIT', 64):
+ with self.assertRaisesRegex(
+ ValueError, 'VARIANT_CONSTRUCTOR_SIZE_LIMIT'):
+ variant_set(column, '$.value', pa.scalar('x' * 128))
+ with self.assertRaisesRegex(
+ ValueError, 'VARIANT_CONSTRUCTOR_SIZE_LIMIT'):
+ variant_set(column, '$.' + 'k' * 128, pa.scalar(True))
+
+ def test_rejects_type_and_length_mismatches(self):
+ column = _variants([{'value': 1.0}, {'value': 2.0}])
+ cases = [
+ ('$.value', pa.scalar('text'), TypeError, "does not match"),
+ ('$.value', pa.scalar(1.0, type=pa.float32()),
+ TypeError, "does not match"),
+ ('$.value', pa.array([1.0]), ValueError, "length must match"),
+ ('$.value', 1.0, TypeError, "Arrow Scalar or Array"),
+ ('value', pa.scalar(1.0), ValueError, "Invalid VARIANT path"),
+ ]
+ for path, replacement, error_type, message in cases:
+ with self.subTest(path=path):
+ with self.assertRaisesRegex(error_type, message):
+ variant_set(column, path, replacement)
+
+ with self.assertRaisesRegex(TypeError, "must be omitted"):
+ variant_set(
+ column, {'$.value': pa.scalar(1.0)}, pa.scalar(2.0))
+
+ def test_rejects_duplicate_and_overlapping_paths(self):
+ column = _variants([{'x': {'y': 1.0}}])
+
+ with self.assertRaisesRegex(ValueError, "must not overlap"):
+ variant_set(column, {
+ '$.x': pa.scalar(1.0),
+ "$['x']": pa.scalar(2.0),
+ })
+ with self.assertRaisesRegex(ValueError, "must not overlap"):
+ variant_set(column, {
+ '$.x': pa.scalar(1.0),
+ '$.x.y': pa.scalar(2.0),
+ })
+
+ def test_rejects_malformed_metadata(self):
+ valid = GenericVariant.from_python({'value': 1.0})
+ column = pa.StructArray.from_arrays(
+ [
+ pa.array([valid.value()]),
+ pa.array([valid.metadata()[:-2]]),
+ ],
+ names=['value', 'metadata'],
+ )
+
+ with self.assertRaisesRegex(ValueError, "MALFORMED_VARIANT"):
+ variant_set(column, '$.processed', pa.scalar(True))
+
+ def test_rejects_unknown_field_id_on_insert(self):
+ metadata = GenericVariant.from_python({'value': 0}).metadata()
+ orphan = _build_object_value([
+ (7, _encode_scalar_to_value_bytes(1.0, pa.float64())),
+ ])
+ column = GenericVariant.to_arrow_array(
+ [GenericVariant(orphan, metadata)])
+
+ with self.assertRaisesRegex(ValueError, "MALFORMED_VARIANT"):
+ variant_set(column, '$.processed', pa.scalar(True))
+
+ def test_rejects_field_id_colliding_with_inserted_key(self):
+ # 'processed' will be assigned id 1; a corrupt source already using
+ # id 1 must be rejected rather than silently producing a duplicate.
+ metadata = GenericVariant.from_python({'value': 0}).metadata()
+ corrupt = _build_object_value([
+ (1, _encode_scalar_to_value_bytes(2.0, pa.float64())),
+ ])
+ column = GenericVariant.to_arrow_array(
+ [GenericVariant(corrupt, metadata)])
+
+ with self.assertRaisesRegex(ValueError, "MALFORMED_VARIANT"):
+ variant_set(column, '$.processed', pa.scalar(True))
+
+ def test_rejects_nested_insert_exposing_invalid_sibling_field_id(self):
+ metadata = GenericVariant.from_python(
+ {'a': 0, 'b': 0, 'child': {}, 'sibling': {}}).metadata()
+ key_ids = _metadata_key_ids(metadata)
+ corrupt_sibling = _build_object_value([
+ (
+ len(key_ids),
+ _encode_scalar_to_value_bytes(2.0, pa.float64()),
+ ),
+ ])
+ corrupt_root = _build_object_value([
+ (key_ids['child'], _build_object_value([])),
+ (key_ids['sibling'], corrupt_sibling),
+ ])
+ column = GenericVariant.to_arrow_array([
+ GenericVariant(corrupt_root, metadata),
+ ])
+
+ with self.assertRaisesRegex(ValueError, "MALFORMED_VARIANT"):
+ variant_set(column, '$.child.new', pa.scalar(True))
+
+ def test_rejects_duplicate_source_field_id(self):
+ metadata = GenericVariant.from_python({'value': 0}).metadata()
+ corrupt = _build_object_value([
+ (0, _encode_scalar_to_value_bytes(1.0, pa.float64())),
+ (0, _encode_scalar_to_value_bytes(2.0, pa.float64())),
+ ])
+ column = GenericVariant.to_arrow_array(
+ [GenericVariant(corrupt, metadata)])
+
+ with self.assertRaisesRegex(ValueError, "MALFORMED_VARIANT"):
+ variant_set(column, '$.value', pa.scalar(9.0))
+
+ def test_rejects_duplicate_source_field_id_in_peer_row(self):
+ metadata = GenericVariant.from_python({'a': 0, 'b': 0}).metadata()
+ duplicate = _build_object_value([
+ (0, _encode_scalar_to_value_bytes(1.0, pa.float64())),
+ (0, _encode_scalar_to_value_bytes(2.0, pa.float64())),
+ ])
+ column = GenericVariant.to_arrow_array([
+ GenericVariant.from_python({'a': 1.0, 'b': 2.0}),
+ GenericVariant(duplicate, metadata),
+ ])
+
+ for updater in (variant_replace, variant_set):
+ with self.subTest(updater=updater.__name__):
+ with self.assertRaisesRegex(
+ ValueError, "MALFORMED_VARIANT"):
+ updater(column, '$.a', pa.scalar(9.0))
+
+ def test_rejects_truncated_child_offsets(self):
+ valid = GenericVariant.from_python({'a': 1.0, 'b': 2.0})
+ truncated = _build_object_value([
+ (0, _encode_scalar_to_value_bytes(1.0, pa.float64())[:-2]),
+ (1, _encode_scalar_to_value_bytes(2.0, pa.float64())),
+ ])
+ column = GenericVariant.to_arrow_array(
+ [GenericVariant(truncated, valid.metadata())])
+ original = column.to_pylist()
+
+ with self.assertRaisesRegex(ValueError, "MALFORMED_VARIANT"):
+ variant_set(column, '$.c', pa.scalar(True))
+ self.assertEqual(column.to_pylist(), original)
+
+
+class TestVariantSetJavaInterop(unittest.TestCase):
+
+ def test_from_python_orders_object_fields_by_utf8(self):
+ key_sup = chr(0x10000)
+ key_bmp = chr(0xE000)
+ payload = {'k%02d' % i: float(i) for i in range(40)}
+ payload[key_sup] = 1.0
+ payload[key_bmp] = 2.0
+
+ variant = GenericVariant.from_python(payload)
+ order = list(variant.to_python().keys())
+ expected = sorted(
+ list(payload.keys()),
+ key=lambda name: name.encode('utf-8'))
+ self.assertEqual(order, expected)
+
+ def test_reads_java_generated_variant(self):
+ column = GenericVariant.to_arrow_array(
+ [GenericVariant(_JAVA_VALUE, _JAVA_METADATA)])
+
+ self.assertEqual(_decode(column), [_JAVA_PYTHON_VALUE])
+ self.assertEqual(
+ variant_get(
+ column, '$.angular_velocity.y', pa.float64()).to_pylist(),
+ [1.5],
+ )
+
+ def test_updates_java_generated_variant(self):
+ column = GenericVariant.to_arrow_array(
+ [GenericVariant(_JAVA_VALUE, _JAVA_METADATA)])
+
+ result = variant_set(column, {
+ '$.angular_velocity.y': pa.scalar(-1.5, type=pa.float64()),
+ '$.processed': pa.scalar(False),
+ '$.mark': pa.scalar('py'),
+ })
+
+ expected = {
+ 'angular_velocity': {'y': -1.5, 'z': -2.5},
+ 'linear_acceleration': {'y': 0.25, 'z': 4.0},
+ 'processed': False,
+ 'seq': 7,
+ 'mark': 'py',
+ }
+ decoded = _decode(result)[0]
+ self.assertEqual(decoded, expected)
+ self.assertEqual(list(decoded), sorted(decoded))
+
+ def test_produces_java_equivalent_encoding(self):
+ # This update was verified to round-trip through the Java
+ # GenericVariant reader (toJson/getFieldByKey, incl. binary search).
+ column = _variants([{
+ 'angular_velocity': {'y': -1.5, 'z': 2.5},
+ 'linear_acceleration': {'y': -0.25, 'z': -4.0},
+ 'seq': 7,
+ }])
+
+ result = variant_set(column, {
+ '$.angular_velocity.y': pa.scalar(1.5, type=pa.float64()),
+ '$.angular_velocity.z': pa.scalar(-2.5, type=pa.float64()),
+ '$.linear_acceleration.y': pa.scalar(0.25, type=pa.float64()),
+ '$.linear_acceleration.z': pa.scalar(4.0, type=pa.float64()),
+ '$.processed': pa.scalar(True, type=pa.bool_()),
+ })
+
+ decoded = _decode(result)[0]
+ self.assertEqual(decoded, _JAVA_PYTHON_VALUE)
+ java_decoded = GenericVariant(
+ _JAVA_VALUE, _JAVA_METADATA).to_python()
+ self.assertEqual(decoded, java_decoded)
+ self.assertEqual(list(decoded), sorted(decoded))
+
+
+class TestMetadataWithKeys(unittest.TestCase):
+
+ def test_reuses_existing_keys(self):
+ metadata = GenericVariant.from_python({'a': 0, 'b': 0}).metadata()
+
+ new_metadata, key_ids, names_by_id = _metadata_with_keys(
+ metadata, ('b',))
+
+ self.assertIsNone(new_metadata)
+ self.assertEqual(key_ids, {'a': 0, 'b': 1})
+ self.assertEqual(names_by_id, {0: 'a', 1: 'b'})
+
+ def test_appends_missing_keys(self):
+ metadata = GenericVariant.from_python({'a': 0}).metadata()
+
+ new_metadata, key_ids, names_by_id = _metadata_with_keys(
+ metadata, ('b', 'c'))
+
+ self.assertEqual(key_ids, {'a': 0, 'b': 1, 'c': 2})
+ self.assertEqual(names_by_id, {0: 'a', 1: 'b', 2: 'c'})
+ self.assertEqual(
+ _metadata_key_ids(new_metadata), {'a': 0, 'b': 1, 'c': 2})
+
+
+if __name__ == '__main__':
+ unittest.main()