JingsongLi commented on code in PR #8136: URL: https://github.com/apache/paimon/pull/8136#discussion_r3418042647
########## paimon-python/pypaimon/common/predicate_json_parser.py: ########## @@ -0,0 +1,322 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +################################################################################ + +import json +import re +from typing import Callable + +import pyarrow as pa +import pyarrow.compute as pc + + +def parse_predicate_to_batch_filter(json_str: str) -> Callable[[pa.RecordBatch], pa.Array]: + data = json.loads(json_str) + return _build_filter(data) + + +def _build_filter(data: dict) -> Callable[[pa.RecordBatch], pa.Array]: + kind = data["kind"] + if kind == "LEAF": + return _build_leaf_filter(data) + elif kind == "COMPOUND": + return _build_compound_filter(data) + raise ValueError(f"Unknown predicate kind: {kind}") + + +def _build_leaf_filter(data: dict) -> Callable: + transform = data["transform"] + function = data["function"] + literals = data.get("literals", []) + + def filter_fn(batch: pa.RecordBatch) -> pa.Array: + value_array = _apply_predicate_transform(transform, batch) + return _apply_leaf_function(function, value_array, literals, len(batch)) + + return filter_fn + + +def _build_compound_filter(data: dict) -> Callable: + function = data["function"] + child_filters = [_build_filter(child) for child in data["children"]] + + def filter_fn(batch: pa.RecordBatch) -> pa.Array: + if function == "AND": + result = child_filters[0](batch) + for cf in child_filters[1:]: + result = pc.and_(result, cf(batch)) + return result + elif function == "OR": + result = child_filters[0](batch) + for cf in child_filters[1:]: + result = pc.or_(result, cf(batch)) + return result + raise ValueError(f"Unknown compound function: {function}") + + return filter_fn + + +def _apply_predicate_transform(transform: dict, batch: pa.RecordBatch) -> pa.Array: + name = transform["name"] + + if name == "FIELD_REF": + return batch.column(transform["fieldRef"]["name"]) + + elif name == "CAST": + col = batch.column(transform["fieldRef"]["name"]) + target_type = _paimon_type_to_arrow(transform["type"]) + return pc.cast(col, target_type) + + elif name == "UPPER": + input_col = _resolve_transform_input(transform["inputs"][0], batch) + return pc.utf8_upper(input_col) + + elif name == "LOWER": + input_col = _resolve_transform_input(transform["inputs"][0], batch) + return pc.utf8_lower(input_col) + + elif name == "CONCAT": + resolved = [_resolve_transform_input(inp, batch) for inp in transform["inputs"]] + if not resolved: + return pa.nulls(len(batch), type=pa.string()) + return pc.binary_join_element_wise(*resolved, "") + + elif name == "CONCAT_WS": + sep = _resolve_transform_input(transform["inputs"][0], batch) + values = [_resolve_transform_input(inp, batch) for inp in transform["inputs"][1:]] + if not values: + return pa.nulls(len(batch), type=pa.string()) + return pc.binary_join_element_wise(*values, sep, null_handling='skip') + + elif name == "NULL": + return pa.nulls(len(batch), type=pa.bool_()) + + raise ValueError(f"Unknown transform type in predicate: {name}") + + +def _resolve_transform_input(inp, batch: pa.RecordBatch) -> pa.Array: + if isinstance(inp, dict): + return batch.column(inp["name"]) + elif isinstance(inp, str): + return pa.array([inp] * len(batch), type=pa.string()) + elif inp is None: + return pa.nulls(len(batch), type=pa.string()) + return pa.array([str(inp)] * len(batch), type=pa.string()) + + +def _apply_leaf_function(function: str, value_array: pa.Array, literals: list, batch_len: int) -> pa.Array: + converted = [_convert_literal(lit, value_array.type) for lit in literals] + + if function == "EQUAL": + return pc.equal(value_array, converted[0]) + elif function == "NOT_EQUAL": + return pc.not_equal(value_array, converted[0]) + elif function == "LESS_THAN": + return pc.less(value_array, converted[0]) + elif function == "LESS_OR_EQUAL": + return pc.less_equal(value_array, converted[0]) + elif function == "GREATER_THAN": + return pc.greater(value_array, converted[0]) + elif function == "GREATER_OR_EQUAL": + return pc.greater_equal(value_array, converted[0]) + elif function == "IS_NULL": + return pc.is_null(value_array) + elif function == "IS_NOT_NULL": + return pc.is_valid(value_array) + elif function == "IN": + return pc.is_in(value_array, pa.array(converted, type=value_array.type)) + elif function == "NOT_IN": + return pc.invert(pc.is_in(value_array, pa.array(converted, type=value_array.type))) Review Comment: Java Paimon's `NotIn` returns false when the input field is null, but `pc.is_in(null, ...)` is false and inverting it makes null rows pass. For a row-filter auth rule such as `dept NOT_IN ('blocked')`, Python would expose rows where `dept` is null. Please combine this with `pc.is_valid(value_array)` and also preserve Java's null-literal behavior, where any null literal makes `NOT_IN` false. ########## paimon-python/pypaimon/common/predicate_json_parser.py: ########## @@ -0,0 +1,322 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +################################################################################ + +import json +import re +from typing import Callable + +import pyarrow as pa +import pyarrow.compute as pc + + +def parse_predicate_to_batch_filter(json_str: str) -> Callable[[pa.RecordBatch], pa.Array]: + data = json.loads(json_str) + return _build_filter(data) + + +def _build_filter(data: dict) -> Callable[[pa.RecordBatch], pa.Array]: + kind = data["kind"] + if kind == "LEAF": + return _build_leaf_filter(data) + elif kind == "COMPOUND": + return _build_compound_filter(data) + raise ValueError(f"Unknown predicate kind: {kind}") + + +def _build_leaf_filter(data: dict) -> Callable: + transform = data["transform"] + function = data["function"] + literals = data.get("literals", []) + + def filter_fn(batch: pa.RecordBatch) -> pa.Array: + value_array = _apply_predicate_transform(transform, batch) + return _apply_leaf_function(function, value_array, literals, len(batch)) + + return filter_fn + + +def _build_compound_filter(data: dict) -> Callable: + function = data["function"] + child_filters = [_build_filter(child) for child in data["children"]] + + def filter_fn(batch: pa.RecordBatch) -> pa.Array: + if function == "AND": + result = child_filters[0](batch) + for cf in child_filters[1:]: + result = pc.and_(result, cf(batch)) + return result + elif function == "OR": + result = child_filters[0](batch) + for cf in child_filters[1:]: + result = pc.or_(result, cf(batch)) + return result + raise ValueError(f"Unknown compound function: {function}") + + return filter_fn + + +def _apply_predicate_transform(transform: dict, batch: pa.RecordBatch) -> pa.Array: + name = transform["name"] + + if name == "FIELD_REF": + return batch.column(transform["fieldRef"]["name"]) + + elif name == "CAST": + col = batch.column(transform["fieldRef"]["name"]) + target_type = _paimon_type_to_arrow(transform["type"]) + return pc.cast(col, target_type) + + elif name == "UPPER": + input_col = _resolve_transform_input(transform["inputs"][0], batch) + return pc.utf8_upper(input_col) + + elif name == "LOWER": + input_col = _resolve_transform_input(transform["inputs"][0], batch) + return pc.utf8_lower(input_col) + + elif name == "CONCAT": + resolved = [_resolve_transform_input(inp, batch) for inp in transform["inputs"]] + if not resolved: + return pa.nulls(len(batch), type=pa.string()) + return pc.binary_join_element_wise(*resolved, "") + + elif name == "CONCAT_WS": + sep = _resolve_transform_input(transform["inputs"][0], batch) + values = [_resolve_transform_input(inp, batch) for inp in transform["inputs"][1:]] + if not values: + return pa.nulls(len(batch), type=pa.string()) + return pc.binary_join_element_wise(*values, sep, null_handling='skip') + + elif name == "NULL": + return pa.nulls(len(batch), type=pa.bool_()) + + raise ValueError(f"Unknown transform type in predicate: {name}") + + +def _resolve_transform_input(inp, batch: pa.RecordBatch) -> pa.Array: + if isinstance(inp, dict): + return batch.column(inp["name"]) + elif isinstance(inp, str): + return pa.array([inp] * len(batch), type=pa.string()) + elif inp is None: + return pa.nulls(len(batch), type=pa.string()) + return pa.array([str(inp)] * len(batch), type=pa.string()) + + +def _apply_leaf_function(function: str, value_array: pa.Array, literals: list, batch_len: int) -> pa.Array: + converted = [_convert_literal(lit, value_array.type) for lit in literals] + + if function == "EQUAL": + return pc.equal(value_array, converted[0]) + elif function == "NOT_EQUAL": + return pc.not_equal(value_array, converted[0]) + elif function == "LESS_THAN": + return pc.less(value_array, converted[0]) + elif function == "LESS_OR_EQUAL": + return pc.less_equal(value_array, converted[0]) + elif function == "GREATER_THAN": + return pc.greater(value_array, converted[0]) + elif function == "GREATER_OR_EQUAL": + return pc.greater_equal(value_array, converted[0]) + elif function == "IS_NULL": + return pc.is_null(value_array) + elif function == "IS_NOT_NULL": + return pc.is_valid(value_array) + elif function == "IN": + return pc.is_in(value_array, pa.array(converted, type=value_array.type)) + elif function == "NOT_IN": + return pc.invert(pc.is_in(value_array, pa.array(converted, type=value_array.type))) + elif function == "BETWEEN": + return pc.and_(pc.greater_equal(value_array, converted[0]), + pc.less_equal(value_array, converted[1])) + elif function == "NOT_BETWEEN": + return pc.or_(pc.less(value_array, converted[0]), + pc.greater(value_array, converted[1])) + elif function == "STARTS_WITH": + return pc.starts_with(value_array, converted[0]) + elif function == "ENDS_WITH": + return pc.ends_with(value_array, converted[0]) + elif function == "CONTAINS": + return pc.match_substring(value_array, converted[0]) + elif function == "LIKE": + raw = literals[0] + pattern = _sql_to_regex_like(raw) + return pc.match_substring_regex(value_array, f"^{pattern}$") + elif function == "TRUE": + return pa.array([True] * batch_len, type=pa.bool_()) + elif function == "FALSE": + return pa.array([False] * batch_len, type=pa.bool_()) + raise ValueError(f"Unknown leaf function: {function}") Review Comment: `IS_NAN` is a valid Java Paimon predicate (`PredicateBuilder.isNaN` / `IsNaN.NAME`) and can be serialized in REST auth filters. With the current switch it fails every Python read with `Unknown leaf function: IS_NAN`; please add a branch using `pc.is_nan` for float/double arrays. ########## paimon-python/pypaimon/read/reader/auth_masking_reader.py: ########## @@ -0,0 +1,220 @@ +################################################################################ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +################################################################################ + +import json +from typing import Callable, Dict, List, Optional + +import pyarrow as pa +import pyarrow.compute as pc + +from pypaimon.common.predicate_json_parser import ( + _collect_all_field_refs_from_transform, + _paimon_type_to_arrow, +) +from pypaimon.read.reader.iface.record_batch_reader import RecordBatchReader + + +class RecordReaderToBatchAdapter(RecordBatchReader): + + def __init__(self, inner, schema: pa.Schema, chunk_size: int = 65536, include_row_kind: bool = False): + self._inner = inner + self._schema = schema + self._chunk_size = chunk_size + self._exhausted = False + self._pending_iterator = None + self._include_row_kind = include_row_kind + + def read_arrow_batch(self) -> Optional[pa.RecordBatch]: + if self._exhausted: + return None + row_tuples = [] + row_kinds = [] + while len(row_tuples) < self._chunk_size: + if self._pending_iterator is not None: + row = self._pending_iterator.next() + while row is not None: + row_tuples.append( + row.row_tuple[row.offset:row.offset + row.arity]) + if self._include_row_kind: + row_kinds.append(row.get_row_kind().to_string()) + if len(row_tuples) >= self._chunk_size: + return self._flush(row_tuples, row_kinds) + row = self._pending_iterator.next() + self._pending_iterator = None + + row_iterator = self._inner.read_batch() + if row_iterator is None: + self._exhausted = True + break + self._pending_iterator = row_iterator + + if not row_tuples: + return None + return self._flush(row_tuples, row_kinds) + + def _flush(self, row_tuples, row_kinds=None): + columns_data = list(zip(*row_tuples)) + pydict = { + name: list(col) + for name, col in zip(self._schema.names, columns_data) + } + batch = pa.RecordBatch.from_pydict(pydict, schema=self._schema) + if row_kinds: + row_kind_array = pa.array(row_kinds, type=pa.string()) + row_kind_field = pa.field("_row_kind", pa.string()) + new_schema = pa.schema([row_kind_field] + list(batch.schema)) + columns = [row_kind_array] + [batch.column(i) for i in range(batch.num_columns)] + batch = pa.RecordBatch.from_arrays(columns, schema=new_schema) + return batch + + def close(self): + self._inner.close() + + +class AuthFilterReader(RecordBatchReader): + + def __init__(self, inner_reader: RecordBatchReader, filter_fn: Callable[[pa.RecordBatch], pa.Array]): + self._inner = inner_reader + self._filter_fn = filter_fn + + def read_arrow_batch(self) -> Optional[pa.RecordBatch]: + batch = self._inner.read_arrow_batch() + if batch is None: + return None + mask = self._filter_fn(batch) + return batch.filter(mask) + + def close(self): + self._inner.close() + + +class AuthMaskingReader(RecordBatchReader): + + def __init__(self, inner_reader: RecordBatchReader, masking_rules: Dict[str, str], read_fields: List): + self._inner = inner_reader + self._masking_rules = masking_rules + self._read_fields = read_fields + self._parsed_rules = {col: json.loads(tj) for col, tj in masking_rules.items()} + read_field_names = {f.name for f in read_fields} + for col_name, transform in self._parsed_rules.items(): + for ref_name in _collect_all_field_refs_from_transform(transform): Review Comment: This validates references for every masking rule before checking whether the masked target column is actually projected. If REST returns a rule like `secret = FIELD_REF(email)` and the user reads only `id`, the Python reader raises because `email` is absent even though `secret` is not returned. Java skips masking rules whose target column is absent from the output row type before remapping inputs, so this should filter to projected target columns first. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
