diff --git a/paimon-common/src/main/java/org/apache/paimon/predicate/StringTransform.java b/paimon-common/src/main/java/org/apache/paimon/predicate/StringTransform.java index 7b87a67c2c5b..2b68d271c681 100644 --- a/paimon-common/src/main/java/org/apache/paimon/predicate/StringTransform.java +++ b/paimon-common/src/main/java/org/apache/paimon/predicate/StringTransform.java @@ -94,6 +94,12 @@ public Object deserialize(JsonParser parser, DeserializationContext context) } } + return unsupported(node, context); + } + + /** Inputs a subclass accepts beyond strings and {@link FieldRef}s. */ + protected Object unsupported(JsonNode node, DeserializationContext context) + throws java.io.IOException { context.reportInputMismatch( Object.class, "Unsupported StringTransform input JSON: %s", node.toString()); return null; @@ -108,15 +114,7 @@ public final List inputs() { @JsonGetter(FIELD_INPUTS) public final List inputsForJson() { - List serialized = new ArrayList<>(inputs.size()); - for (Object input : inputs) { - if (input instanceof BinaryString) { - serialized.add(input.toString()); - } else { - serialized.add(input); - } - } - return serialized; + return inputsForJson(inputs); } @Override @@ -157,8 +155,27 @@ public int hashCode() { @Override public String toString() { - List inputs = - this.inputs.stream().map(Object::toString).collect(Collectors.toList()); - return name() + "(" + String.join(", ", inputs) + ')'; + return formatCall(name(), inputs); + } + + /** Inputs as written to JSON: {@link BinaryString} literals become JSON strings. */ + static List inputsForJson(List inputs) { + List serialized = new ArrayList<>(inputs.size()); + for (Object input : inputs) { + if (input instanceof BinaryString) { + serialized.add(input.toString()); + } else { + serialized.add(input); + } + } + return serialized; + } + + /** Renders a transform as {@code NAME(input, input)}. */ + static String formatCall(String name, List inputs) { + return name + + "(" + + inputs.stream().map(String::valueOf).collect(Collectors.joining(", ")) + + ')'; } } diff --git a/paimon-common/src/main/java/org/apache/paimon/predicate/SubstringTransform.java b/paimon-common/src/main/java/org/apache/paimon/predicate/SubstringTransform.java index 054422a20125..c54b5d972d36 100644 --- a/paimon-common/src/main/java/org/apache/paimon/predicate/SubstringTransform.java +++ b/paimon-common/src/main/java/org/apache/paimon/predicate/SubstringTransform.java @@ -23,6 +23,15 @@ import org.apache.paimon.types.DataType; import org.apache.paimon.types.DataTypes; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.annotation.JsonCreator; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.annotation.JsonGetter; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.annotation.JsonIgnore; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.annotation.JsonProperty; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.databind.DeserializationContext; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.databind.JsonNode; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.databind.annotation.JsonDeserialize; + +import java.io.IOException; import java.util.List; import java.util.Objects; @@ -39,11 +48,37 @@ public class SubstringTransform implements Transform { private final List inputs; - public SubstringTransform(List inputs) { + @JsonCreator + public SubstringTransform( + @JsonProperty(StringTransform.FIELD_INPUTS) + @JsonDeserialize(contentUsing = InputDeserializer.class) + List inputs) { checkArgument(inputs.size() == 2 || inputs.size() == 3); this.inputs = inputs; } + /** Deserializer for {@link SubstringTransform} inputs, which may also be integers. */ + public static class InputDeserializer extends StringTransform.InputDeserializer { + + private static final long serialVersionUID = 1L; + + @Override + protected Object unsupported(JsonNode node, DeserializationContext context) + throws IOException { + if (node.isNumber()) { + // canConvertToInt checks the range but not integrality + if (!node.isIntegralNumber()) { + context.reportInputMismatch( + Object.class, + "SubstringTransform position must be an integer: %s", + node.toString()); + } + return node.canConvertToInt() ? node.intValue() : node.numberValue(); + } + return super.unsupported(node, context); + } + } + @Override public String name() { return NAME; @@ -71,6 +106,10 @@ public final Object transform(InternalRow row) { if (begin instanceof FieldRef) { FieldRef beginRef = (FieldRef) begin; checkArgument(beginRef.type().is(INTEGER_NUMERIC)); + // getInt on a null reads an undefined value on columnar rows + if (row.isNullAt(beginRef.index())) { + return null; + } beginIndex = row.getInt(beginRef.index()); } else { beginIndex = Integer.parseInt(inputs.get(1).toString()); @@ -85,6 +124,9 @@ public final Object transform(InternalRow row) { if (end instanceof FieldRef) { FieldRef endRef = (FieldRef) inputs.get(2); checkArgument(endRef.type().is(INTEGER_NUMERIC)); + if (row.isNullAt(endRef.index())) { + return null; + } endIndex = beginIndex + row.getInt(endRef.index()) - 1; } else { endIndex = beginIndex + Integer.parseInt(inputs.get(2).toString()) - 1; @@ -103,10 +145,16 @@ public Transform copyWithNewInputs(List inputs) { } @Override + @JsonIgnore public final List inputs() { return inputs; } + @JsonGetter(StringTransform.FIELD_INPUTS) + public final List inputsForJson() { + return StringTransform.inputsForJson(inputs); + } + @Override public boolean equals(Object o) { if (o == null || getClass() != o.getClass()) { @@ -125,4 +173,9 @@ public DataType outputType() { public int hashCode() { return Objects.hashCode(inputs); } + + @Override + public String toString() { + return StringTransform.formatCall(name(), inputs); + } } diff --git a/paimon-common/src/main/java/org/apache/paimon/predicate/Transform.java b/paimon-common/src/main/java/org/apache/paimon/predicate/Transform.java index 826cfd800255..ad01afcfb7be 100644 --- a/paimon-common/src/main/java/org/apache/paimon/predicate/Transform.java +++ b/paimon-common/src/main/java/org/apache/paimon/predicate/Transform.java @@ -39,6 +39,8 @@ @JsonSubTypes.Type(value = ConcatWsTransform.class, name = ConcatWsTransform.NAME), @JsonSubTypes.Type(value = UpperTransform.class, name = UpperTransform.NAME), @JsonSubTypes.Type(value = LowerTransform.class, name = LowerTransform.NAME), + @JsonSubTypes.Type(value = SubstringTransform.class, name = SubstringTransform.NAME), + @JsonSubTypes.Type(value = TrimTransform.class, name = TrimTransform.NAME), @JsonSubTypes.Type(value = NullTransform.class, name = NullTransform.NAME) }) public interface Transform extends Serializable { diff --git a/paimon-common/src/main/java/org/apache/paimon/predicate/TrimTransform.java b/paimon-common/src/main/java/org/apache/paimon/predicate/TrimTransform.java index 6182335bb221..4441fa52ed60 100644 --- a/paimon-common/src/main/java/org/apache/paimon/predicate/TrimTransform.java +++ b/paimon-common/src/main/java/org/apache/paimon/predicate/TrimTransform.java @@ -21,9 +21,16 @@ import org.apache.paimon.data.BinaryString; import org.apache.paimon.utils.StringUtils; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.annotation.JsonCreator; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.annotation.JsonGetter; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.annotation.JsonProperty; +import org.apache.paimon.shade.jackson2.com.fasterxml.jackson.databind.annotation.JsonDeserialize; + import java.util.List; +import java.util.Objects; import static org.apache.paimon.utils.Preconditions.checkArgument; +import static org.apache.paimon.utils.Preconditions.checkNotNull; /** TRIM/LTRIM/RTRIM {@link Transform}. */ public class TrimTransform extends StringTransform { @@ -34,10 +41,17 @@ public class TrimTransform extends StringTransform { private final Flag trimFlag; - public TrimTransform(List inputs, Flag trimFlag) { + public static final String FIELD_TRIM_FLAG = "trimFlag"; + + @JsonCreator + public TrimTransform( + @JsonProperty(StringTransform.FIELD_INPUTS) + @JsonDeserialize(contentUsing = StringTransform.InputDeserializer.class) + List inputs, + @JsonProperty(FIELD_TRIM_FLAG) Flag trimFlag) { super(inputs); - this.trimFlag = trimFlag; checkArgument(inputs.size() == 1 || inputs.size() == 2); + this.trimFlag = checkNotNull(trimFlag, "trimFlag must not be null"); } @Override @@ -45,13 +59,25 @@ public String name() { return NAME; } + @JsonGetter(FIELD_TRIM_FLAG) + public Flag trimFlag() { + return trimFlag; + } + @Override public BinaryString transform(List inputs) { if (inputs.get(0) == null) { return null; } String sourceString = inputs.get(0).toString(); - String charsToTrim = inputs.size() == 1 ? " " : inputs.get(1).toString(); + String charsToTrim = " "; + if (inputs.size() == 2) { + if (inputs.get(1) == null) { + // StringUtils.ltrim/rtrim treat a null charsToTrim as a null result + return null; + } + charsToTrim = inputs.get(1).toString(); + } switch (trimFlag) { case BOTH: return BinaryString.fromString(StringUtils.trim(sourceString, charsToTrim)); @@ -69,6 +95,20 @@ public Transform copyWithNewInputs(List inputs) { return new TrimTransform(inputs, this.trimFlag); } + @Override + public boolean equals(Object o) { + if (!super.equals(o)) { + return false; + } + TrimTransform that = (TrimTransform) o; + return trimFlag == that.trimFlag; + } + + @Override + public int hashCode() { + return Objects.hash(super.hashCode(), trimFlag); + } + /** Enum of trim functions. */ public enum Flag { LEADING, diff --git a/paimon-common/src/test/java/org/apache/paimon/predicate/ConcatTransformTest.java b/paimon-common/src/test/java/org/apache/paimon/predicate/ConcatTransformTest.java index e776040f89f0..a9cc9e7a5524 100644 --- a/paimon-common/src/test/java/org/apache/paimon/predicate/ConcatTransformTest.java +++ b/paimon-common/src/test/java/org/apache/paimon/predicate/ConcatTransformTest.java @@ -72,4 +72,13 @@ public void testConcatHybridInputs() { BinaryString.fromString("-he"))); assertThat(result).isEqualTo(BinaryString.fromString("ha-he")); } + + @Test + public void testToStringWithNullInput() { + List inputs = new ArrayList<>(); + inputs.add(BinaryString.fromString("a")); + inputs.add(null); + + assertThat(new ConcatTransform(inputs).toString()).isEqualTo("CONCAT(a, null)"); + } } diff --git a/paimon-common/src/test/java/org/apache/paimon/predicate/SubstringTransformTest.java b/paimon-common/src/test/java/org/apache/paimon/predicate/SubstringTransformTest.java index b4d998bea9ec..ee9bfe0c874c 100644 --- a/paimon-common/src/test/java/org/apache/paimon/predicate/SubstringTransformTest.java +++ b/paimon-common/src/test/java/org/apache/paimon/predicate/SubstringTransformTest.java @@ -100,6 +100,41 @@ public void testSubstringRefInputs() { assertThat(result).isEqualTo(BinaryString.fromString("ell")); } + @Test + public void testNullPositionFieldYieldsNull() { + List inputs = new ArrayList<>(); + inputs.add(new FieldRef(0, "f0", DataTypes.STRING())); + inputs.add(new FieldRef(1, "f1", DataTypes.INT())); + assertThat( + new SubstringTransform(inputs) + .transform( + GenericRow.of( + BinaryString.fromString("123-45-6789"), null))) + .isNull(); + + inputs.add(new FieldRef(2, "f2", DataTypes.INT())); + assertThat( + new SubstringTransform(inputs) + .transform( + GenericRow.of( + BinaryString.fromString("123-45-6789"), 8, null))) + .isNull(); + } + + @Test + public void testPositionsCountUtf16CodeUnits() { + List inputs = new ArrayList<>(); + inputs.add(new FieldRef(0, "f0", DataTypes.STRING())); + inputs.add(2); + inputs.add(2); + + Object result = + new SubstringTransform(inputs) + .transform(GenericRow.of(BinaryString.fromString("๐Ÿ˜€abc"))); + + assertThat(result).isEqualTo(BinaryString.fromString("?a")); + } + @Test public void testSubstringRefInputUsesSourceFieldNullability() { List inputs = new ArrayList<>(); diff --git a/paimon-common/src/test/java/org/apache/paimon/predicate/TransformJsonSerdeTest.java b/paimon-common/src/test/java/org/apache/paimon/predicate/TransformJsonSerdeTest.java index 455add1181e6..bb4414fc09bd 100644 --- a/paimon-common/src/test/java/org/apache/paimon/predicate/TransformJsonSerdeTest.java +++ b/paimon-common/src/test/java/org/apache/paimon/predicate/TransformJsonSerdeTest.java @@ -19,10 +19,13 @@ package org.apache.paimon.predicate; import org.apache.paimon.data.BinaryString; +import org.apache.paimon.data.GenericRow; +import org.apache.paimon.data.InternalRow; import org.apache.paimon.types.DataTypes; import org.apache.paimon.types.IntType; import org.apache.paimon.utils.JsonSerdeUtil; +import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.MethodSource; @@ -70,6 +73,13 @@ private static Stream testData() { new FieldRef(1, "f1", DataTypes.STRING())))) .expectJson( "{\"name\":\"UPPER\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"}]}"), + TestSpec.forTransform( + new LowerTransform( + Collections.singletonList( + new FieldRef(1, "f1", DataTypes.STRING())))) + .expectJson( + "{\"name\":\"LOWER\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"}]}"), + TestSpec.forTransform(NullTransform.INSTANCE).expectJson("{\"name\":\"NULL\"}"), // ConcatTransform - two fields TestSpec.forTransform( @@ -112,10 +122,66 @@ private static Stream testData() { new FieldRef(2, "f2", DataTypes.STRING())))) .expectJson( "{\"name\":\"CONCAT_WS\",\"inputs\":[\"|\",{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"},\"X\",null,{\"index\":2,\"name\":\"f2\",\"type\":\"STRING\"}]}"), + TestSpec.forTransform( + new SubstringTransform( + Arrays.asList( + new FieldRef(1, "f1", DataTypes.STRING()), 8, 4))) + .expectJson( + "{\"name\":\"SUBSTRING\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"},8,4]}"), + TestSpec.forTransform( + new SubstringTransform( + Arrays.asList( + new FieldRef(1, "f1", DataTypes.STRING()), 8))) + .expectJson( + "{\"name\":\"SUBSTRING\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"},8]}"), + TestSpec.forTransform( + new SubstringTransform( + Arrays.asList( + new FieldRef(1, "f1", DataTypes.STRING()), + new FieldRef(3, "f3", DataTypes.INT()), + new FieldRef(4, "f4", DataTypes.INT())))) + .expectJson( + "{\"name\":\"SUBSTRING\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"},{\"index\":3,\"name\":\"f3\",\"type\":\"INT\"},{\"index\":4,\"name\":\"f4\",\"type\":\"INT\"}]}"), + TestSpec.forTransform( + new SubstringTransform( + Arrays.asList(BinaryString.fromString("hello"), 2, 3))) + .expectJson("{\"name\":\"SUBSTRING\",\"inputs\":[\"hello\",2,3]}"), + TestSpec.forTransform(new SubstringTransform(Arrays.asList(null, 1))) + .expectJson("{\"name\":\"SUBSTRING\",\"inputs\":[null,1]}"), + TestSpec.forTransform( + new TrimTransform( + Collections.singletonList( + new FieldRef(1, "f1", DataTypes.STRING())), + TrimTransform.Flag.BOTH)) + .expectJson( + "{\"name\":\"TRIM\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"}],\"trimFlag\":\"BOTH\"}"), + TestSpec.forTransform( + new TrimTransform( + Collections.singletonList( + new FieldRef(1, "f1", DataTypes.STRING())), + TrimTransform.Flag.LEADING)) + .expectJson( + "{\"name\":\"TRIM\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"}],\"trimFlag\":\"LEADING\"}"), + TestSpec.forTransform( + new TrimTransform( + Arrays.asList( + new FieldRef(1, "f1", DataTypes.STRING()), + BinaryString.fromString("x")), + TrimTransform.Flag.TRAILING)) + .expectJson( + "{\"name\":\"TRIM\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"},\"x\"],\"trimFlag\":\"TRAILING\"}"), // error message testing TestSpec.forJson("{\"name\":\"invalid\"}") - .expectErrorMessage("Could not resolve type id 'invalid'")); + .expectErrorMessage("Could not resolve type id 'invalid'"), + TestSpec.forJson( + "{\"name\":\"TRIM\",\"inputs\":[{\"index\":1,\"name\":\"f1\",\"type\":\"STRING\"}]}") + .expectErrorMessage("trimFlag must not be null"), + TestSpec.forJson("{\"name\":\"SUBSTRING\",\"inputs\":[true,1]}") + .expectErrorMessage("Unsupported StringTransform input JSON"), + TestSpec.forJson( + "{\"name\":\"SUBSTRING\",\"inputs\":[{\"index\":0,\"name\":\"f0\",\"type\":\"STRING\"},1.5]}") + .expectErrorMessage("position must be an integer")); } @ParameterizedTest(name = "{index}: {0}") @@ -136,6 +202,14 @@ void testJsonParsing(TestSpec testSpec) { } } + @ParameterizedTest(name = "{index}: {0}") + @MethodSource("testData") + void testSerializedText(TestSpec testSpec) { + if (testSpec.expectedJson != null) { + assertThat(toJson(testSpec.transform)).isEqualTo(testSpec.expectedJson); + } + } + @ParameterizedTest(name = "{index}: {0}") @MethodSource("testData") void testErrorMessage(TestSpec testSpec) { @@ -145,6 +219,67 @@ void testErrorMessage(TestSpec testSpec) { } } + @Test + void testSubstringRoundTripKeepsPositions() { + FieldRef ssn = new FieldRef(0, "ssn", DataTypes.VARCHAR(64)); + assertRoundTrip( + new SubstringTransform(Arrays.asList(ssn, 8, 4)), + GenericRow.of(BinaryString.fromString("123-45-6789")), + BinaryString.fromString("6789")); + + FieldRef phone = new FieldRef(0, "phone", DataTypes.VARCHAR(64)); + assertRoundTrip( + new SubstringTransform(Arrays.asList(phone, 1, 3)), + GenericRow.of(BinaryString.fromString("13812348000")), + BinaryString.fromString("138")); + + assertRoundTrip( + new SubstringTransform( + Arrays.asList( + new FieldRef(0, "f0", DataTypes.STRING()), + new FieldRef(1, "f1", DataTypes.INT()), + new FieldRef(2, "f2", DataTypes.INT()))), + GenericRow.of(BinaryString.fromString("123-45-6789"), 8, 4), + BinaryString.fromString("6789")); + + assertRoundTrip( + new SubstringTransform(Arrays.asList(BinaryString.fromString("123-45-6789"), 8)), + GenericRow.of(), + BinaryString.fromString("6789")); + } + + @Test + void testTrimRoundTripKeepsFlag() { + FieldRef f0 = new FieldRef(0, "f0", DataTypes.STRING()); + GenericRow row = GenericRow.of(BinaryString.fromString(" x ")); + + assertRoundTrip( + new TrimTransform(Collections.singletonList(f0), TrimTransform.Flag.BOTH), + row, + BinaryString.fromString("x")); + assertRoundTrip( + new TrimTransform(Collections.singletonList(f0), TrimTransform.Flag.LEADING), + row, + BinaryString.fromString("x ")); + assertRoundTrip( + new TrimTransform(Collections.singletonList(f0), TrimTransform.Flag.TRAILING), + row, + BinaryString.fromString(" x")); + + assertThat(new TrimTransform(Collections.singletonList(f0), TrimTransform.Flag.LEADING)) + .isNotEqualTo( + new TrimTransform(Collections.singletonList(f0), TrimTransform.Flag.BOTH)); + } + + private static void assertRoundTrip(Transform transform, InternalRow row, Object expected) { + assertThat(transform.transform(row)).isEqualTo(expected); + + Transform parsed = parse(toJson(transform)); + assertThat(parsed.transform(row)).isEqualTo(expected); + assertThat(parsed).isEqualTo(transform); + assertThat(toJson(parsed)).isEqualTo(toJson(transform)); + } + private static String toJson(Transform transform) { return JsonSerdeUtil.toFlatJson(transform); } diff --git a/paimon-common/src/test/java/org/apache/paimon/predicate/TrimTransformTest.java b/paimon-common/src/test/java/org/apache/paimon/predicate/TrimTransformTest.java index b24fda78a7d7..597a71d280a7 100644 --- a/paimon-common/src/test/java/org/apache/paimon/predicate/TrimTransformTest.java +++ b/paimon-common/src/test/java/org/apache/paimon/predicate/TrimTransformTest.java @@ -101,6 +101,18 @@ public void testNormalInputs() { assertThat(result).isEqualTo(BinaryString.fromString(" aa")); } + @Test + public void testNullCharsToTrimYieldsNull() { + List inputs = new ArrayList<>(); + inputs.add(new FieldRef(0, "f0", DataTypes.STRING())); + inputs.add(new FieldRef(1, "f1", DataTypes.STRING())); + GenericRow row = GenericRow.of(BinaryString.fromString(" x "), null); + + for (TrimTransform.Flag flag : TrimTransform.Flag.values()) { + assertThat(new TrimTransform(inputs, flag).transform(row)).isNull(); + } + } + @Test public void testSubstringRefInputs() { List inputs = new ArrayList<>(); diff --git a/paimon-python/pypaimon/common/predicate_json_parser.py b/paimon-python/pypaimon/common/predicate_json_parser.py index c89da9fce010..e2f839d9a665 100644 --- a/paimon-python/pypaimon/common/predicate_json_parser.py +++ b/paimon-python/pypaimon/common/predicate_json_parser.py @@ -23,6 +23,27 @@ import pyarrow as pa import pyarrow.compute as pc +# utf8_slice_codeunits needs an explicit integer stop on pyarrow 6 +_MAX_STOP = 2 ** 31 - 1 + +_INT_MIN, _INT_MAX = -2 ** 31, 2 ** 31 - 1 + +# Integer.parseInt syntax +_JAVA_INT = re.compile(r"[+-]?[0-9]+\Z") + +# an omitted third input, as opposed to one that is explicitly null +_ABSENT = object() + +# the DataTypeFamily.INTEGER_NUMERIC members a position field may declare +_INTEGER_TYPES = ("TINYINT", "SMALLINT", "INT", "BIGINT") + +# per trimFlag: the Arrow kernel, and the str method for the per-row form +_TRIM_OPS = { + "BOTH": (pc.utf8_trim, str.strip), + "LEADING": (pc.utf8_ltrim, str.lstrip), + "TRAILING": (pc.utf8_rtrim, str.rstrip), +} + def parse_predicate_to_batch_filter(json_str: str) -> Callable[[pa.RecordBatch], pa.Array]: data = json.loads(json_str) @@ -103,12 +124,172 @@ def _apply_predicate_transform(transform: dict, batch: pa.RecordBatch, return pa.nulls(len(batch), type=pa.string()) return _concat_ws(sep, values) + elif name == "SUBSTRING": + return _substring(transform["inputs"], batch) + + elif name == "TRIM": + flag = transform.get("trimFlag") + if flag is None: + raise ValueError("TRIM rule is missing trimFlag") + return _trim(transform["inputs"], flag, batch) + elif name == "NULL": return pa.nulls(len(batch), type=null_type) raise ValueError(f"Unknown transform type: {name}") +def _substring(inputs, batch: pa.RecordBatch) -> pa.Array: + if len(inputs) not in (2, 3): + raise ValueError(f"SUBSTRING takes 2 or 3 inputs, got {len(inputs)}") + source = _resolve_transform_input(inputs[0], batch) + begin = inputs[1] + length = inputs[2] if len(inputs) == 3 else _ABSENT + + # a malformed literal is not rejected here: Java only reads a position once a row + # reaches it, so the per-row path raises at the point Java would + begin_literal = _literal_position(begin) + length_literal = _literal_position(length) if length is not _ABSENT else None + + # the kernel only matches Java for a positive begin and length; everything else + # goes per row, where Java's order of checks can be followed + if begin_literal is not None and begin_literal >= 1: + if length is _ABSENT: + return pc.utf8_slice_codeunits(source, start=begin_literal - 1, stop=_MAX_STOP) + if ( + length_literal is not None + and length_literal > 0 + and begin_literal + length_literal - 1 <= _INT_MAX + ): + start = begin_literal - 1 + return pc.utf8_slice_codeunits(source, start=start, stop=start + length_literal) + + return _substring_per_row(source, begin, length, batch) + + +def _int_position(value): + """A SUBSTRING begin/length, with Java's tolerance and no more: Integer.parseInt + takes "+2" and "007" but not "1.5", "1_0" or " 2 ", all of which int() accepts.""" + if isinstance(value, bool) or isinstance(value, float): + raise ValueError(f"SUBSTRING position must be an integer: {value!r}") + if isinstance(value, str): + if not _JAVA_INT.match(value): + raise ValueError(f"SUBSTRING position must be an integer: {value!r}") + position = int(value) + elif isinstance(value, int): + position = value + else: + raise ValueError(f"SUBSTRING position must be an integer: {value!r}") + if not _INT_MIN <= position <= _INT_MAX: + raise ValueError(f"SUBSTRING position is out of the integer range: {value!r}") + return position + + +def _literal_position(value): + """The value of a literal position, or None when it is a field or unusable here.""" + if value is None or isinstance(value, dict): + return None + try: + return _int_position(value) + except ValueError: + return None + + +def _position_values(inp, batch: pa.RecordBatch) -> list: + """Raw, unparsed positions; Java parses one only when a row reaches it.""" + if isinstance(inp, dict): + declared = inp.get("type", "").split("(")[0].split()[0].upper() + if declared not in _INTEGER_TYPES: + raise ValueError(f"SUBSTRING position field must be an integer type: {inp.get('type')}") + return batch.column(inp["name"]).to_pylist() + return [inp] * len(batch) + + +def _substring_per_row(source: pa.Array, begin, length, batch: pa.RecordBatch) -> pa.Array: + # mirrors SubstringTransform.transform, including the order of its checks: each + # position is parsed only once a row actually reaches it + begins = _position_values(begin, batch) + begin_is_field = isinstance(begin, dict) + has_length = length is not _ABSENT + lengths = _position_values(length, batch) if has_length else None + length_is_field = has_length and isinstance(length, dict) + result = [] + for i, value in enumerate(source.to_pylist()): + if value is None: + result.append(None) + continue + raw_begin = begins[i] + if raw_begin is None: + # a null field value is as unusable as a null source, while a literal null + # is what Java reads with toString() and fails on + if not begin_is_field: + raise ValueError("SUBSTRING begin must not be null") + result.append(None) + continue + begin_index = _int_position(raw_begin) + if begin_index > len(value): + result.append("") + continue + stop = len(value) + if has_length: + raw_length = lengths[i] + if raw_length is None: + if not length_is_field: + raise ValueError("SUBSTRING length must not be null") + result.append(None) + continue + length_value = _int_position(raw_length) + end = begin_index + length_value - 1 + if end > _INT_MAX: + # Java adds these in 32 bits, where the sum wraps into a failure + raise ValueError( + f"SUBSTRING end overflows the integer range: " + f"begin={begin_index}, length={length_value}" + ) + stop = min(end, len(value)) + start = begin_index - 1 + if start < 0 or start >= stop: + raise ValueError(f"SUBSTRING out of bounds: begin={begin_index}, stop={stop}") + result.append(value[start:stop]) + return pa.array(result, type=pa.string()) + + +def _trim(inputs, flag: str, batch: pa.RecordBatch) -> pa.Array: + if len(inputs) not in (1, 2): + raise ValueError(f"TRIM takes 1 or 2 inputs, got {len(inputs)}") + source = _resolve_transform_input(inputs[0], batch) + # Java's one-input TRIM trims spaces only, not every whitespace character. + chars = " " if len(inputs) == 1 else inputs[1] + + # validated first: Jackson rejects an unknown flag when the rule is read, so it must + # not survive the null shortcut below either + kernel = _trim_ops(flag)[0] + + if isinstance(chars, dict): + return _trim_per_row(source, flag, batch.column(chars["name"]).to_pylist()) + + if chars is None: + # Java masks the whole column to null for a null charsToTrim + return pa.nulls(len(batch), type=pa.string()) + + return kernel(source, characters=chars) + + +def _trim_ops(flag: str): + ops = _TRIM_OPS.get(flag) + if ops is None: + raise ValueError(f"Unknown trimFlag: {flag}") + return ops + + +def _trim_per_row(source: pa.Array, flag: str, chars_per_row: list) -> pa.Array: + strip = _trim_ops(flag)[1] + result = [] + for value, chars in zip(source.to_pylist(), chars_per_row): + result.append(None if value is None or chars is None else strip(value, chars)) + return pa.array(result, type=pa.string()) + + def _resolve_transform_input(inp, batch: pa.RecordBatch) -> pa.Array: if isinstance(inp, dict): return batch.column(inp["name"]) diff --git a/paimon-python/pypaimon/tests/auth_masking_reader_test.py b/paimon-python/pypaimon/tests/auth_masking_reader_test.py index 7745f94a5a59..80e740c6b356 100644 --- a/paimon-python/pypaimon/tests/auth_masking_reader_test.py +++ b/paimon-python/pypaimon/tests/auth_masking_reader_test.py @@ -29,6 +29,9 @@ from pypaimon.read.reader.iface.record_batch_reader import RecordBatchReader +_NO_LENGTH = object() + + class _FakeField: def __init__(self, name): self.name = name @@ -250,6 +253,346 @@ def test_concat_ws_field_ref_separator(self): ) +class TestSubstringTransform(unittest.TestCase): + + def setUp(self): + self.batch = pa.RecordBatch.from_pydict({ + "ssn": ["123-45-6789", "987-65-4321", None], + "begin": [8, 1, 1], + "length": [4, 3, 3], + }) + self.fields = [_FakeField("ssn"), _FakeField("begin"), _FakeField("length")] + + def _mask(self, transform, batch=None, fields=None): + reader = AuthMaskingReader( + _FakeBatchReader([batch if batch is not None else self.batch]), + {"ssn": json.dumps(transform)}, + fields if fields is not None else self.fields, + ) + return reader.read_arrow_batch().column("ssn").to_pylist() + + @staticmethod + def _ssn_ref(): + return {"index": 0, "name": "ssn", "type": "STRING"} + + def test_begin_and_length(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 8, 4]}), + ["6789", "4321", None], + ) + + def test_begin_only_runs_to_end_of_string(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 8]}), + ["6789", "4321", None], + ) + + def test_begin_past_end_yields_empty_string_not_null(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 99, 4]}), + ["", "", None], + ) + + def test_length_longer_than_string_is_clamped(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 8, 100]}), + ["6789", "4321", None], + ) + + def test_positions_read_from_other_fields(self): + self.assertEqual( + self._mask({ + "name": "SUBSTRING", + "inputs": [ + self._ssn_ref(), + {"index": 1, "name": "begin", "type": "INT"}, + {"index": 2, "name": "length", "type": "INT"}, + ], + }), + ["6789", "987", None], + ) + + def _mask_with_position_fields(self, begin, length=_NO_LENGTH, ssn="123-45-6789"): + cols = {"ssn": pa.array([ssn], type=pa.string()), + "begin": pa.array([begin], type=pa.int32())} + inputs = [self._ssn_ref(), {"index": 1, "name": "begin", "type": "INT"}] + if length is not _NO_LENGTH: + cols["length"] = pa.array([length], type=pa.int32()) + inputs.append({"index": 2, "name": "length", "type": "INT"}) + batch = pa.RecordBatch.from_arrays(list(cols.values()), names=list(cols)) + reader = AuthMaskingReader( + _FakeBatchReader([batch]), + {"ssn": json.dumps({"name": "SUBSTRING", "inputs": inputs})}, + [_FakeField(n) for n in cols], + ) + return reader.read_arrow_batch().column("ssn").to_pylist() + + def test_begin_past_end_yields_empty_string_for_field_positions(self): + self.assertEqual(self._mask_with_position_fields(99, 4), [""]) + + def test_non_positive_length_rejected_for_field_positions(self): + with self.assertRaisesRegex(ValueError, "SUBSTRING out of bounds"): + self._mask_with_position_fields(1, 0) + + def test_begin_only_runs_to_end_for_field_positions(self): + self.assertEqual(self._mask_with_position_fields(8), ["6789"]) + + def test_null_position_masks_the_row_to_null(self): + self.assertEqual(self._mask_with_position_fields(None, 4), [None]) + self.assertEqual(self._mask_with_position_fields(8, None), [None]) + + def test_non_ascii_positions_count_characters_on_both_paths(self): + source = "่บซไปฝ่ฏ12345678" + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 8, 4]}, + batch=pa.RecordBatch.from_arrays( + [pa.array([source], type=pa.string())], names=["ssn"]), + fields=[_FakeField("ssn")]), + ["5678"], + ) + self.assertEqual(self._mask_with_position_fields(8, 4, ssn=source), ["5678"]) + + def test_fractional_position_rejected(self): + for begin in [1.5, 2.0, "1.5", True]: + with self.assertRaisesRegex(ValueError, "must be an integer"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), begin, 2]}) + + def test_textual_position_accepted_like_integer_parse_int(self): + for begin in ["8", "+8", "008"]: + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), begin, 4]}), + ["6789", "4321", None], + begin, + ) + + def test_textual_position_outside_parse_int_syntax_rejected(self): + for begin in ["1_0", " 2 ", "2\n", "ูฃ"]: + with self.assertRaisesRegex(ValueError, "must be an integer"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), begin, 4]}) + + def test_explicit_null_length_is_not_an_omitted_length(self): + with self.assertRaisesRegex(ValueError, "length must not be null"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 8, None]}) + + def test_explicit_null_begin_rejected(self): + with self.assertRaisesRegex(ValueError, "begin must not be null"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), None]}) + + def test_wrong_arity_rejected(self): + with self.assertRaisesRegex(ValueError, "SUBSTRING takes 2 or 3 inputs"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 8, 4, 9]}) + with self.assertRaisesRegex(ValueError, "SUBSTRING takes 2 or 3 inputs"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref()]}) + + def test_position_field_of_a_non_integer_type_rejected(self): + batch = pa.RecordBatch.from_arrays( + [pa.array(["abcdef"], type=pa.string()), pa.array(["2"], type=pa.string())], + names=["ssn", "begin"], + ) + reader = AuthMaskingReader( + _FakeBatchReader([batch]), + {"ssn": json.dumps({ + "name": "SUBSTRING", + "inputs": [self._ssn_ref(), {"index": 1, "name": "begin", "type": "STRING"}], + })}, + [_FakeField("ssn"), _FakeField("begin")], + ) + with self.assertRaisesRegex(ValueError, "must be an integer type"): + reader.read_arrow_batch() + + def test_null_source_wins_over_a_bad_begin(self): + batch = pa.RecordBatch.from_arrays( + [pa.array([None], type=pa.string())], names=["ssn"]) + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [None, None]}, + batch=batch, fields=[_FakeField("ssn")]), + [None], + ) + + def test_begin_past_end_wins_over_a_malformed_length(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": ["abc", 99, "bad"]}), + ["", "", ""], + ) + + def test_end_overflowing_the_integer_range_rejected(self): + with self.assertRaisesRegex(ValueError, "overflows the integer range"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 2, 2 ** 31 - 1]}) + + def test_position_outside_the_integer_range_rejected(self): + with self.assertRaisesRegex(ValueError, "out of the integer range"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 2 ** 31, 4]}) + + def test_begin_past_end_wins_over_a_null_length_field(self): + self.assertEqual(self._mask_with_position_fields(99, None), [""]) + + def test_supplementary_characters_count_code_points(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 2, 2]}, + batch=pa.RecordBatch.from_arrays( + [pa.array(["\U0001F600abc"], type=pa.string())], names=["ssn"]), + fields=[_FakeField("ssn")]), + ["ab"], + ) + + def test_literal_source_instead_of_field(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": ["123-45-6789", 8, 4]}), + ["6789", "6789", "6789"], + ) + + def test_non_positive_length_rejected(self): + with self.assertRaisesRegex(ValueError, "SUBSTRING out of bounds"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 1, 0]}) + + def test_begin_below_one_rejected(self): + with self.assertRaisesRegex(ValueError, "SUBSTRING out of bounds"): + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 0]}) + + def test_begin_below_one_rejected_for_field_positions(self): + batch = pa.RecordBatch.from_pydict({"ssn": ["123-45-6789"], "begin": [0]}) + reader = AuthMaskingReader( + _FakeBatchReader([batch]), + {"ssn": json.dumps({ + "name": "SUBSTRING", + "inputs": [self._ssn_ref(), {"index": 1, "name": "begin", "type": "INT"}], + })}, + [_FakeField("ssn"), _FakeField("begin")], + ) + with self.assertRaisesRegex(ValueError, "SUBSTRING out of bounds"): + reader.read_arrow_batch() + + def test_begin_past_end_wins_over_a_bad_length(self): + self.assertEqual( + self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 99, 0]}), + ["", "", None], + ) + + +class TestTrimTransform(unittest.TestCase): + + def setUp(self): + self.batch = pa.RecordBatch.from_pydict({ + "s": [" x ", "\ty\t", None], + "chars": [" ", "\t", "z"], + }) + self.fields = [_FakeField("s"), _FakeField("chars")] + + def _mask(self, transform): + reader = AuthMaskingReader( + _FakeBatchReader([self.batch]), {"s": json.dumps(transform)}, self.fields + ) + return reader.read_arrow_batch().column("s").to_pylist() + + @staticmethod + def _transform(flag, extra_inputs=()): + return { + "name": "TRIM", + "inputs": [{"index": 0, "name": "s", "type": "STRING"}, *extra_inputs], + "trimFlag": flag, + } + + def test_both(self): + self.assertEqual(self._mask(self._transform("BOTH")), ["x", "\ty\t", None]) + + def test_leading(self): + self.assertEqual(self._mask(self._transform("LEADING")), ["x ", "\ty\t", None]) + + def test_trailing(self): + self.assertEqual(self._mask(self._transform("TRAILING")), [" x", "\ty\t", None]) + + def test_wrong_arity_rejected(self): + with self.assertRaisesRegex(ValueError, "TRIM takes 1 or 2 inputs"): + self._mask(self._transform("BOTH", ["x", "y"])) + + def test_unknown_flag_rejected(self): + for flag in ["both", "LTRIM"]: + with self.assertRaisesRegex(ValueError, "Unknown trimFlag"): + self._mask(self._transform(flag)) + + def test_unknown_flag_rejected_with_null_chars(self): + with self.assertRaisesRegex(ValueError, "Unknown trimFlag"): + self._mask({ + "name": "TRIM", + "inputs": [{"index": 0, "name": "s", "type": "STRING"}, None], + "trimFlag": "LTRIM", + }) + + def _trim_by(self, chars, values): + batch = pa.RecordBatch.from_arrays( + [pa.array(values, type=pa.string())], names=["s"]) + reader = AuthMaskingReader( + _FakeBatchReader([batch]), + {"s": json.dumps({ + "name": "TRIM", + "inputs": [{"index": 0, "name": "s", "type": "STRING"}, chars], + "trimFlag": "BOTH", + })}, + [_FakeField("s")], + ) + return reader.read_arrow_batch().column("s").to_pylist() + + def test_multibyte_trim_characters(self): + self.assertEqual(self._trim_by("ใ€‚", ["ใ€‚ใ€‚xใ€‚ใ€‚", " y "]), ["x", " y "]) + + def test_trim_matches_whole_characters_not_bytes(self): + self.assertEqual(self._trim_by("ใ€", ["ใ€‚xใ€‚"]), ["ใ€‚xใ€‚"]) + + def test_custom_chars_are_treated_as_a_set(self): + batch = pa.RecordBatch.from_pydict({"s": ["xyzaxyz", "zyxaxyz", None]}) + reader = AuthMaskingReader( + _FakeBatchReader([batch]), + {"s": json.dumps({ + "name": "TRIM", + "inputs": [{"index": 0, "name": "s", "type": "STRING"}, "xyz"], + "trimFlag": "BOTH", + })}, + [_FakeField("s")], + ) + self.assertEqual( + reader.read_arrow_batch().column("s").to_pylist(), ["a", "a", None] + ) + + def test_chars_read_from_another_field(self): + self.assertEqual( + self._mask( + self._transform("BOTH", [{"index": 1, "name": "chars", "type": "STRING"}]) + ), + ["x", "y", None], + ) + + def test_literal_null_trim_chars_yields_null(self): + self.assertEqual( + self._mask({ + "name": "TRIM", + "inputs": [{"index": 0, "name": "s", "type": "STRING"}, None], + "trimFlag": "BOTH", + }), + [None, None, None], + ) + + def test_null_trim_chars_yields_null(self): + batch = pa.RecordBatch.from_arrays( + [pa.array([" x "], type=pa.string()), pa.array([None], type=pa.string())], + names=["s", "chars"], + ) + reader = AuthMaskingReader( + _FakeBatchReader([batch]), + {"s": json.dumps( + self._transform("BOTH", [{"index": 1, "name": "chars", "type": "STRING"}]) + )}, + [_FakeField("s"), _FakeField("chars")], + ) + self.assertEqual(reader.read_arrow_batch().column("s").to_pylist(), [None]) + + def test_missing_flag_rejected(self): + with self.assertRaisesRegex(ValueError, "trimFlag"): + self._mask({ + "name": "TRIM", + "inputs": [{"index": 0, "name": "s", "type": "STRING"}], + }) + + class TestMaskingOrderIndependence(unittest.TestCase): def test_cross_reference_uses_original_batch(self):