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 957cde3b3c [common][python] Support SUBSTRING and TRIM in column
masking (#9212)
957cde3b3c is described below
commit 957cde3b3c3560e957b6e25bbb108f382082996c
Author: Jiajia Li <[email protected]>
AuthorDate: Wed Aug 26 11:47:47 2026 +0800
[common][python] Support SUBSTRING and TRIM in column masking (#9212)
---
.../apache/paimon/predicate/StringTransform.java | 41 +-
.../paimon/predicate/SubstringTransform.java | 53 +-
.../org/apache/paimon/predicate/Transform.java | 2 +
.../org/apache/paimon/predicate/TrimTransform.java | 63 ++-
.../paimon/predicate/TransformJsonSerdeTest.java | 242 +++++++++-
.../pypaimon/common/predicate_json_parser.py | 235 +++++++++
.../pypaimon/tests/auth_masking_reader_test.py | 534 +++++++++++++++++++++
7 files changed, 1154 insertions(+), 16 deletions(-)
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 1dd2405ca6..66291572c6 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 abstract class StringTransform implements Transform {
}
}
+ return otherInput(node, context);
+ }
+
+ /** Inputs a subclass accepts beyond strings and {@link FieldRef}s. */
+ protected Object otherInput(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 abstract class StringTransform implements Transform
{
@JsonGetter(FIELD_INPUTS)
public final List<Object> inputsForJson() {
- List<Object> 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 abstract class StringTransform implements Transform
{
@Override
public String toString() {
- List<String> inputs =
-
this.inputs.stream().map(String::valueOf).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<Object> inputsForJson(List<Object> inputs) {
+ List<Object> 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<Object> 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 7be7e66e8e..99f46c11f9 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.data.InternalRow;
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,42 @@ public class SubstringTransform implements Transform {
private final List<Object> inputs;
- public SubstringTransform(List<Object> inputs) {
+ @JsonCreator
+ public SubstringTransform(
+ @JsonProperty(StringTransform.FIELD_INPUTS)
+ @JsonDeserialize(contentUsing = InputDeserializer.class)
+ List<Object> inputs) {
checkArgument(inputs.size() == 2 || inputs.size() == 3);
+ Object source = inputs.get(0);
+ // transform() casts this slot to BinaryString
+ checkArgument(
+ source == null || source instanceof FieldRef || source
instanceof BinaryString,
+ "SUBSTRING source must be a string or a field reference");
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 otherInput(JsonNode node, DeserializationContext
context)
+ throws IOException {
+ if (node.isNumber()) {
+ // canConvertToInt checks the range, 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.otherInput(node, context);
+ }
+ }
+
@Override
public String name() {
return NAME;
@@ -119,10 +159,16 @@ public class SubstringTransform implements Transform {
}
@Override
+ @JsonIgnore
public final List<Object> inputs() {
return inputs;
}
+ @JsonGetter(StringTransform.FIELD_INPUTS)
+ public final List<Object> inputsForJson() {
+ return StringTransform.inputsForJson(inputs);
+ }
+
@Override
public boolean equals(Object o) {
if (o == null || getClass() != o.getClass()) {
@@ -141,4 +187,9 @@ public class SubstringTransform implements Transform {
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 826cfd8002..ad01afcfb7 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 @@ import java.util.List;
@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 cd97456019..d4a882de6f 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
@@ -20,9 +20,21 @@ package org.apache.paimon.predicate;
import org.apache.paimon.data.BinaryString;
+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.core.JsonParser;
+import
org.apache.paimon.shade.jackson2.com.fasterxml.jackson.databind.DeserializationContext;
+import
org.apache.paimon.shade.jackson2.com.fasterxml.jackson.databind.JsonDeserializer;
+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;
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,12 +46,40 @@ public class TrimTransform extends StringTransform {
/** The one-input form trims spaces only, not every whitespace character.
*/
private static final BinaryString SPACE = BinaryString.fromString(" ");
+ public static final String FIELD_TRIM_FLAG = "trimFlag";
+
private final Flag trimFlag;
- public TrimTransform(List<Object> inputs, Flag trimFlag) {
+ @JsonCreator
+ public TrimTransform(
+ @JsonProperty(StringTransform.FIELD_INPUTS)
+ @JsonDeserialize(contentUsing =
StringTransform.InputDeserializer.class)
+ List<Object> inputs,
+ @JsonProperty(FIELD_TRIM_FLAG) @JsonDeserialize(using =
FlagDeserializer.class)
+ Flag trimFlag) {
super(inputs);
- this.trimFlag = trimFlag;
checkArgument(inputs.size() == 1 || inputs.size() == 2);
+ this.trimFlag = checkNotNull(trimFlag, "trimFlag must not be null");
+ }
+
+ /** Deserializer for {@link Flag}: Jackson would also accept an ordinal or
its text. */
+ public static class FlagDeserializer extends JsonDeserializer<Flag> {
+
+ @Override
+ public Flag deserialize(JsonParser parser, DeserializationContext
context)
+ throws IOException {
+ JsonNode node = parser.readValueAsTree();
+ if (node.isTextual()) {
+ for (Flag flag : Flag.values()) {
+ if (flag.name().equals(node.asText())) {
+ return flag;
+ }
+ }
+ }
+ context.reportInputMismatch(
+ Flag.class, "TRIM trimFlag must be one of LEADING,
TRAILING, BOTH: %s", node);
+ return null;
+ }
}
@Override
@@ -47,6 +87,11 @@ public class TrimTransform extends StringTransform {
return NAME;
}
+ @JsonGetter(FIELD_TRIM_FLAG)
+ public Flag trimFlag() {
+ return trimFlag;
+ }
+
@Override
public BinaryString transform(List<BinaryString> inputs) {
if (inputs.get(0) == null) {
@@ -77,6 +122,20 @@ public class TrimTransform extends StringTransform {
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/TransformJsonSerdeTest.java
b/paimon-common/src/test/java/org/apache/paimon/predicate/TransformJsonSerdeTest.java
index 455add1181..3e6b97a2a2 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 @@ class TransformJsonSerdeTest {
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,69 @@ class TransformJsonSerdeTest {
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"),
+
TestSpec.forJson("{\"name\":\"SUBSTRING\",\"inputs\":[123,1,1]}")
+ .expectErrorMessage(
+ "SUBSTRING source must be a string or a field
reference"));
}
@ParameterizedTest(name = "{index}: {0}")
@@ -136,6 +205,14 @@ class TransformJsonSerdeTest {
}
}
+ @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 +222,169 @@ class TransformJsonSerdeTest {
}
}
+ @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 testPositionsOutsideTheIntegerRangeAreNotTruncated() {
+ // canConvertToInt is false beyond an int, so the value stays wide and
the
+ // per-row parse rejects it rather than silently truncating
+ Transform parsed =
parse("{\"name\":\"SUBSTRING\",\"inputs\":[\"abcdef\",3000000000]}");
+ assertThatThrownBy(() -> parsed.transform(GenericRow.of()))
+ .isInstanceOf(NumberFormatException.class);
+ }
+
+ @Test
+ void testTextualPositionIsParsedPerRow() {
+ // Jackson keeps a textual position as a string; Java parses it when a
row
+ // reaches it, which is the contract the Python client mirrors
+ Transform parsed =
+
parse("{\"name\":\"SUBSTRING\",\"inputs\":[\"123-45-6789\",\"8\",\"4\"]}");
+
assertThat(parsed.transform(GenericRow.of())).isEqualTo(BinaryString.fromString("6789"));
+
+ Transform bad =
parse("{\"name\":\"SUBSTRING\",\"inputs\":[\"abcdef\",\"1_0\"]}");
+ assertThatThrownBy(() -> bad.transform(GenericRow.of()))
+ .isInstanceOf(NumberFormatException.class);
+ }
+
+ @Test
+ void testTrimSourceTypeIsCheckedWhenTheRuleIsRead() {
+ for (String type : new String[] {"INT", "VARCHAR(0)", "STRING ARRAY"})
{
+ assertThatThrownBy(
+ () ->
+ parse(
+
"{\"name\":\"TRIM\",\"inputs\":[{\"index\":0,\"name\":\"s\",\"type\":\""
+ + type
+ +
"\"}],\"trimFlag\":\"BOTH\"}"))
+ .isInstanceOf(RuntimeException.class);
+ }
+ for (String type :
+ new String[] {"STRING", "STRING NULL", "CHAR(3) NOT NULL",
"VARCHAR(10)NULL"}) {
+ assertThat(
+ parse(
+
"{\"name\":\"TRIM\",\"inputs\":[{\"index\":0,\"name\":\"s\",\"type\":\""
+ + type
+ + "\"}],\"trimFlag\":\"BOTH\"}"))
+ .isNotNull();
+ }
+ }
+
+ @Test
+ void testWrongArityIsRejected() {
+ for (String inputs : new String[] {"[\"abc\"]", "[\"abc\",1,2,3]"}) {
+ assertThatThrownBy(() ->
parse("{\"name\":\"SUBSTRING\",\"inputs\":" + inputs + "}"))
+ .isInstanceOf(RuntimeException.class);
+ }
+ for (String inputs : new String[] {"[]", "[\"a\",\"b\",\"c\"]"}) {
+ assertThatThrownBy(
+ () ->
+ parse(
+ "{\"name\":\"TRIM\",\"inputs\":"
+ + inputs
+ +
",\"trimFlag\":\"BOTH\"}"))
+ .isInstanceOf(RuntimeException.class);
+ }
+ }
+
+ @Test
+ void testCopyWithNewInputsKeepsTheFlag() {
+ // auth remapping rebuilds every transform through copyWithNewInputs,
so a flag
+ // lost there would silently turn LEADING into BOTH
+ FieldRef f0 = new FieldRef(0, "f0", DataTypes.STRING());
+ GenericRow row = GenericRow.of(BinaryString.fromString(" x "));
+ for (TrimTransform.Flag flag : TrimTransform.Flag.values()) {
+ Transform copied =
+ new
TrimTransform(Collections.singletonList(BinaryString.fromString("")), flag)
+ .copyWithNewInputs(Collections.singletonList(f0));
+ assertThat(copied.transform(row))
+ .isEqualTo(
+ new TrimTransform(Collections.singletonList(f0),
flag).transform(row));
+ }
+ assertThat(
+ new TrimTransform(Collections.singletonList(f0),
TrimTransform.Flag.LEADING)
+
.copyWithNewInputs(Collections.singletonList(f0))
+ .transform(row))
+ .isEqualTo(BinaryString.fromString("x "));
+ }
+
+ @Test
+ void testTrimFlagMustBeItsName() {
+ for (String flag :
+ new String[] {"0", "\"0\"", "2", "\"LTRIM\"", "null",
"\"both\"", "\"Both\""}) {
+ assertThatThrownBy(
+ () ->
+ parse(
+ "{\"name\":\"TRIM\",\"inputs\":[\"
x \"],\"trimFlag\":"
+ + flag
+ + "}"))
+ .isInstanceOf(RuntimeException.class);
+ }
+
+ assertThat(parse("{\"name\":\"TRIM\",\"inputs\":[\" x
\"],\"trimFlag\":\"LEADING\"}"))
+ .isEqualTo(
+ new TrimTransform(
+
Collections.singletonList(BinaryString.fromString(" x ")),
+ TrimTransform.Flag.LEADING));
+ }
+
+ @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-python/pypaimon/common/predicate_json_parser.py
b/paimon-python/pypaimon/common/predicate_json_parser.py
index c89da9fce0..0f2e5b4e3e 100644
--- a/paimon-python/pypaimon/common/predicate_json_parser.py
+++ b/paimon-python/pypaimon/common/predicate_json_parser.py
@@ -23,6 +23,26 @@ from typing import Callable
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: an optional sign and Unicode decimal digits, which
+# Character.digit accepts, but no whitespace or underscore, which int() would.
+# Java reads UTF-16 chars, so a supplementary-plane digit fails there.
+_JAVA_INT = re.compile(r"[+-]?\d+\Z")
+
+# an omitted third input, as opposed to one that is explicitly null
+_ABSENT = object()
+
+# 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 +123,227 @@ 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 not isinstance(inputs, list):
+ raise ValueError(f"SUBSTRING inputs must be a list, got {inputs!r}")
+ 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
+
+ _check_string_input("SUBSTRING source", inputs[0], batch)
+
+ # Jackson refuses a non-integral number, a boolean or an array when the
rule is
+ # read; a null is not malformed and propagates to a null result, as in SQL
+ for position in (begin,) + ((length,) if length is not _ABSENT else ()):
+ if position is None or isinstance(position, dict):
+ continue
+ if isinstance(position, bool) or not isinstance(position, (int, str)):
+ raise ValueError(f"SUBSTRING position must be an integer:
{position!r}")
+
+ # a malformed literal is left to the per-row path, which raises where Java
does
+ begin_literal = _literal_position(begin)
+ length_literal = _literal_position(length) if length is not _ABSENT else
None
+
+ # the kernel only matches the SQL semantics for a positive begin and length
+ 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_0" or " 2 ", 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) or any(ord(c) > 0xFFFF for c in 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 _field_column(inp, batch: pa.RecordBatch):
+ """The column a field reference names. Java rebuilds the reference from
the schema
+ before a rule runs, so the index it carries and, for SUBSTRING, the type
it carries
+ never decide anything; only the name is used to find the column."""
+ name = inp.get("name")
+ if not isinstance(name, str) or name not in batch.schema.names:
+ raise ValueError(f"Column masking refers to a field that is not
present: {inp!r}")
+ return batch.column(name)
+
+
+# StringTransform validates a stored FieldRef type when the rule is read,
before
+# TableQueryAuthResult replaces it from the schema. The whole spelling has to
match:
+# Java parses "STRING ARRAY" as an array type, accepts an explicit NULL as
well as
+# NOT NULL, needs no space after a ")" but does after a bare keyword, and
rejects a
+# length outside CharType and VarCharType's [1, MAX_VALUE].
+_CHARACTER_TYPE = re.compile(
+ r"\s*(?:STRING(?=\s|\Z)|(?:CHAR|VARCHAR)\s*(?:\(\s*(\d+)\s*\)|(?=\s|\Z)))"
+ r"(?:\s*(?:NOT\s+)?NULL)?\s*\Z",
+ re.IGNORECASE,
+)
+
+
+def _check_stored_character_type(slot: str, inp) -> None:
+ stored = inp.get("type")
+ matched = _CHARACTER_TYPE.match(stored) if isinstance(stored, str) else
None
+ length = matched.group(1) if matched else None
+ if matched is None or (length is not None and not 1 <= int(length) <=
_INT_MAX):
+ raise ValueError(f"{slot} field must be a string: {stored!r}")
+
+
+def _check_string_input(slot: str, value, batch: pa.RecordBatch, stored: bool
= False) -> None:
+ """A string input is a literal or a field reference to a string column."""
+ if value is None or isinstance(value, str):
+ return
+ if not isinstance(value, dict):
+ raise ValueError(f"{slot} must be a string or a field: {value!r}")
+ if stored:
+ _check_stored_character_type(slot, value)
+ column = _field_column(value, batch)
+ if not pa.types.is_string(column.type) and not
pa.types.is_large_string(column.type):
+ raise ValueError(f"{slot} field must be a string: {column.type}")
+
+
+class _Positions:
+ """A position slot resolved lazily, as Java reads one only when a row
reaches it."""
+
+ def __init__(self, inp, batch: pa.RecordBatch):
+ self._inp = inp
+ self._batch = batch
+ self._values = None
+
+ def value(self, index: int):
+ if self._values is None:
+ self._values = self._resolve()
+ return self._values[index]
+
+ def _resolve(self) -> list:
+ inp = self._inp
+ if not isinstance(inp, dict):
+ return [inp] * len(self._batch)
+ column = _field_column(inp, self._batch)
+ if not pa.types.is_integer(column.type):
+ raise ValueError(
+ f"SUBSTRING position field must be an integer type:
{column.type}")
+ return column.to_pylist()
+
+
+def _substring_sql(value: str, pos: int, length: int) -> str:
+ """BinaryString.substringSQL: one-based, zero means one, negative counts
from the end."""
+ chars = len(value)
+ start = pos - 1 if pos > 0 else (chars + pos if pos < 0 else 0)
+ end = start + length
+ # Java computes the end in long arithmetic and saturates it into an int
+ end = min(max(end, _INT_MIN), _INT_MAX)
+ if end <= start or start >= chars:
+ return ""
+ return value[max(start, 0):min(max(end, 0), chars)]
+
+
+def _substring_per_row(source: pa.Array, begin, length, batch: pa.RecordBatch)
-> pa.Array:
+ # mirrors SubstringTransform.transform, including the order of its checks
+ begins = _Positions(begin, batch)
+ has_length = length is not _ABSENT
+ lengths = _Positions(length, batch) if has_length else None
+ result = []
+ for i, value in enumerate(source.to_pylist()):
+ if value is None:
+ result.append(None)
+ continue
+ raw_begin = begins.value(i)
+ # SQL null propagation: every position is checked before any is parsed
+ if raw_begin is None or (has_length and lengths.value(i) is None):
+ result.append(None)
+ continue
+ pos = _int_position(raw_begin)
+ length_value = _int_position(lengths.value(i)) if has_length else
_INT_MAX
+ result.append(_substring_sql(value, pos, length_value))
+ return pa.array(result, type=source.type)
+
+
+def _trim(inputs, flag: str, batch: pa.RecordBatch) -> pa.Array:
+ if not isinstance(inputs, list):
+ raise ValueError(f"TRIM inputs must be a list, got {inputs!r}")
+ if len(inputs) not in (1, 2):
+ raise ValueError(f"TRIM takes 1 or 2 inputs, got {len(inputs)}")
+ _check_string_input("TRIM source", inputs[0], batch, stored=True)
+ if len(inputs) == 2:
+ _check_string_input("TRIM characters", inputs[1], batch, stored=True)
+ 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
+ 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=source.type)
+
+ 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=source.type)
+
+
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 7745f94a5a..66b837cb97 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.auth_masking_reader import (
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,537 @@ class TestAuthMaskingReaderTransforms(unittest.TestCase):
)
+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 _mask_with_bigint_position(self, begin):
+ batch = pa.RecordBatch.from_arrays(
+ [pa.array(["abcdef"], type=pa.string()), pa.array([begin],
type=pa.int64())],
+ names=["ssn", "begin"])
+ reader = AuthMaskingReader(
+ _FakeBatchReader([batch]),
+ {"ssn": json.dumps({"name": "SUBSTRING", "inputs": [
+ self._ssn_ref(), {"index": 1, "name": "begin", "type":
"BIGINT"}]})},
+ [_FakeField("ssn"), _FakeField("begin")])
+ return reader.read_arrow_batch().column("ssn").to_pylist()
+
+ def _mask_typed(self, inputs, cols, names):
+ batch = pa.RecordBatch.from_arrays(cols, names=names)
+ reader = AuthMaskingReader(
+ _FakeBatchReader([batch]),
+ {names[0]: json.dumps({"name": "SUBSTRING", "inputs": inputs})},
+ [_FakeField(n) for n in names])
+ return reader.read_arrow_batch().column(names[0]).to_pylist()
+
+ def test_non_integer_position_field_rejected_like_java(self):
+ with self.assertRaisesRegex(ValueError, "must be an integer type"):
+ self._mask_typed(
+ [self._ssn_ref(), 99, {"index": 1, "name": "n", "type":
"STRING"}],
+ [pa.array(["abc"]), pa.array(["2"])], ["ssn", "n"])
+ with self.assertRaisesRegex(ValueError, "must be an integer type"):
+ self._mask_typed(
+ [self._ssn_ref(), {"index": 1, "name": "n", "type": "STRING"}],
+ [pa.array(["abc"]), pa.array([None], type=pa.string())],
["ssn", "n"])
+
+ def test_malformed_position_rejected_even_when_no_row_reads_it(self):
+ # Jackson refuses these when the rule is read, so a null source must
not hide them
+ for inputs in ([None, 1.5], [None, 2, 1.5], [None, True]):
+ with self.assertRaisesRegex(ValueError, "position must be an
integer"):
+ self._mask({"name": "SUBSTRING", "inputs": inputs})
+
+ def test_negative_begin_beyond_the_length_is_clamped(self):
+ # start and end both go negative here; without the clamps the slice
would
+ # silently return most of the value instead of nothing
+ for begin, length, expected in ((-9, 2, ""), (-9, 5, "ab"), (-7, 2,
"a")):
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": ["abcdef", begin,
length]}),
+ [expected] * 3)
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": ["abcdef", -9]}),
+ ["abcdef"] * 3)
+
+ def test_per_row_result_keeps_the_source_column_type(self):
+ batch = pa.RecordBatch.from_arrays(
+ [pa.array(["abcdef"], type=pa.large_string()), pa.array([2],
type=pa.int32())],
+ names=["ssn", "p"])
+ reader = AuthMaskingReader(
+ _FakeBatchReader([batch]),
+ {"ssn": json.dumps({"name": "SUBSTRING", "inputs": [
+ self._ssn_ref(), {"index": 1, "name": "p", "type": "INT"}]})},
+ [_FakeField("ssn"), _FakeField("p")])
+ self.assertEqual(reader.read_arrow_batch().column("ssn").type,
pa.large_string())
+
+ def test_substring_ignores_the_stored_source_type(self):
+ # SubstringTransform does not extend StringTransform, so its
constructor never
+ # looks at the stored type; only the remapped column decides
+ self.assertEqual(
+ self._mask_typed([{"index": 0, "name": "ssn", "type": "INT"}, 2],
+ [pa.array(["abcdef"], type=pa.string())],
["ssn"]),
+ ["bcdef"])
+
+ def test_source_field_must_be_a_string_column(self):
+ # the stored type is discarded by the remap, so the actual column
decides
+ with self.assertRaisesRegex(ValueError, "source field must be a
string"):
+ self._mask_typed([{"index": 0, "name": "n", "type": "STRING"}, 1],
+ [pa.array([1], type=pa.int32())], ["n"])
+
+ def test_position_field_untouched_when_the_source_is_null(self):
+ self.assertEqual(
+ self._mask_typed([None, {"index": 1, "name": "p"}],
+ [pa.array(["ignored"]), pa.array(["notint"],
type=pa.string())],
+ ["ssn", "p"]),
+ [None])
+
+ 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_yields_empty_for_field_positions(self):
+ self.assertEqual(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_number_in_the_source_slot_rejected(self):
+ with self.assertRaisesRegex(ValueError, "source must be a string or a
field"):
+ self._mask({"name": "SUBSTRING", "inputs": [123, 1, 1]})
+
+ def test_supplementary_plane_digit_rejected(self):
+ with self.assertRaisesRegex(ValueError, "must be an integer"):
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(),
"\U0001D7DA", 4]})
+
+ def
test_malformed_position_literal_fails_even_when_the_row_short_circuits(self):
+ for length in [1.5, True]:
+ with self.assertRaisesRegex(ValueError, "must be an integer"):
+ self._mask({"name": "SUBSTRING", "inputs": ["abcdef", 99,
length]})
+
+ def test_unicode_digit_position_accepted_like_character_digit(self):
+ for begin in ["8", "\u0668", "\uff18"]:
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(),
begin, 4]}),
+ ["6789", "4321", None],
+ begin,
+ )
+
+ def test_explicit_null_position_propagates(self):
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 8,
None]}),
+ [None, None, None],
+ )
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(),
None]}),
+ [None, None, None],
+ )
+
+ def test_null_length_propagates_even_when_begin_is_past_the_end(self):
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 99,
None]}),
+ [None, None, 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_any_integer_width_is_read(self):
+ for declared, arrow_type in [("BIGINT", pa.int64()), ("INT",
pa.int32()),
+ ("SMALLINT", pa.int16()), ("TINYINT",
pa.int8())]:
+ batch = pa.RecordBatch.from_arrays(
+ [pa.array(["abcdef"], type=pa.string()), pa.array([2],
type=arrow_type)],
+ names=["ssn", "begin"],
+ )
+ reader = AuthMaskingReader(
+ _FakeBatchReader([batch]),
+ {"ssn": json.dumps({
+ "name": "SUBSTRING",
+ "inputs": [self._ssn_ref(), {"index": 1, "name": "begin",
"type": declared}],
+ })},
+ [_FakeField("ssn"), _FakeField("begin")],
+ )
+
self.assertEqual(reader.read_arrow_batch().column("ssn").to_pylist(), ["bcdef"])
+
+ self.assertEqual(self._mask_with_bigint_position(None), [None])
+
+ 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, "bad"]},
+ batch=batch, fields=[_FakeField("ssn")]),
+ [None],
+ )
+
+ def test_malformed_length_is_read_even_when_begin_is_past_the_end(self):
+ # substringSQL reads both positions; Java no longer returns early on
begin
+ with self.assertRaisesRegex(ValueError, "position must be an integer"):
+ self._mask({"name": "SUBSTRING", "inputs": ["abc", 99, "bad"]})
+
+ def test_null_length_wins_over_a_malformed_begin(self):
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": ["abc", "bad", None]}),
+ [None, None, None],
+ )
+
+ def test_end_overflowing_the_integer_range_is_clamped(self):
+ # Java saturates the end in long arithmetic instead of wrapping
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 2, 2
** 31 - 1]}),
+ ["23-45-6789", "87-65-4321", None])
+
+ 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_null_length_field_propagates_even_when_begin_is_past_the_end(self):
+ self.assertEqual(self._mask_with_position_fields(99, None), [None])
+
+ def test_supplementary_characters_count_code_points_like_java(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_yields_empty(self):
+ for length in (0, -1):
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(),
1, length]}),
+ ["", "", None])
+
+ def test_begin_zero_means_one(self):
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), 0]}),
+ ["123-45-6789", "987-65-4321", None])
+
+ def test_negative_begin_counts_from_the_end(self):
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), -4]}),
+ ["6789", "4321", None])
+ self.assertEqual(
+ self._mask({"name": "SUBSTRING", "inputs": [self._ssn_ref(), -4,
2]}),
+ ["67", "43", None])
+
+ def test_begin_zero_means_one_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")],
+ )
+ self.assertEqual(
+ reader.read_arrow_batch().column("ssn").to_pylist(),
["123-45-6789"])
+
+ def test_begin_past_end_wins_over_a_zero_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_numeric_trim_flag_rejected(self):
+ for flag in (0, "0", 2):
+ with self.assertRaisesRegex(ValueError, "trimFlag"):
+ self._mask({"name": "TRIM", "inputs": [" x "], "trimFlag":
flag})
+
+ def test_non_string_input_rejected_naming_the_slot(self):
+ with self.assertRaisesRegex(ValueError, "TRIM source must be a
string"):
+ self._mask({"name": "TRIM", "inputs": [123], "trimFlag": "BOTH"})
+ with self.assertRaisesRegex(ValueError, "TRIM characters must be a
string"):
+ self._mask({"name": "TRIM", "inputs": [" x ", 123], "trimFlag":
"BOTH"})
+
+ def test_stored_source_type_is_checked_when_the_rule_is_read(self):
+ for bad in ("INT", "VARCHAR(0)", "STRING ARRAY", "VARCHAR(10)
GARBAGE"):
+ with self.assertRaisesRegex(ValueError, "TRIM source"):
+ self._mask({"name": "TRIM",
+ "inputs": [{"index": 0, "name": "s", "type": bad}],
+ "trimFlag": "BOTH"})
+ with self.assertRaisesRegex(ValueError, "TRIM source"):
+ self._mask({"name": "TRIM", "inputs": [{"index": 0, "name": "s"}],
+ "trimFlag": "BOTH"})
+ for bad in ("STRINGNULL", "STRINGNOT NULL"):
+ with self.assertRaisesRegex(ValueError, "TRIM source"):
+ self._mask({"name": "TRIM",
+ "inputs": [{"index": 0, "name": "s", "type": bad}],
+ "trimFlag": "BOTH"})
+ # Java's tokenizer needs no space after a ")" but does after a bare
keyword
+ for good in ("STRING NOT NULL", "STRING NULL", "CHAR(3)",
"varchar(10)",
+ "VARCHAR(10)NULL", "CHAR(3)NOT NULL"):
+ self.assertEqual(
+ self._mask({"name": "TRIM",
+ "inputs": [{"index": 0, "name": "s", "type":
good}],
+ "trimFlag": "BOTH"}),
+ ["x", "\ty\t", None])
+
+ self.assertEqual(
+ self._mask({"name": "TRIM",
+ "inputs": [{"index": 0, "name": "s", "type":
"VARCHAR(10)"}],
+ "trimFlag": "BOTH"}),
+ ["x", "\ty\t", None])
+
+ def test_inputs_must_be_a_list(self):
+ with self.assertRaisesRegex(ValueError, "must be a list"):
+ self._mask({"name": "TRIM", "inputs": "x", "trimFlag": "BOTH"})
+
+ with self.assertRaisesRegex(RuntimeError, "not present"):
+ self._mask({"name": "TRIM",
+ "inputs": [{"index": 0, "name": "nope", "type":
"STRING"}],
+ "trimFlag": "BOTH"})
+
+ def test_structured_position_rejected_before_any_shortcut(self):
+ with self.assertRaisesRegex(ValueError, "position must be an integer"):
+ self._mask({"name": "SUBSTRING", "inputs": ["abc", 99, []]})
+ with self.assertRaisesRegex(ValueError, "position must be an integer"):
+ self._mask({"name": "SUBSTRING", "inputs": ["abc", []]})
+
+ 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_chars_from_a_field_are_a_set_not_an_affix(self):
+ batch = pa.RecordBatch.from_pydict({"s": ["zxayxz"], "chars": ["xz"]})
+ 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(),
["ay"])
+
+ 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):