github-actions[bot] commented on code in PR #68768: URL: https://github.com/apache/doris/pull/68768#discussion_r4217619395
########## fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlParameters.java: ########## @@ -0,0 +1,315 @@ +// 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. + +package org.apache.doris.service.arrowflight; + +import org.apache.doris.nereids.StatementContext; +import org.apache.doris.nereids.trees.expressions.Cast; +import org.apache.doris.nereids.trees.expressions.Placeholder; +import org.apache.doris.nereids.trees.expressions.SubqueryExpr; +import org.apache.doris.nereids.trees.expressions.literal.BigIntLiteral; +import org.apache.doris.nereids.trees.expressions.literal.BooleanLiteral; +import org.apache.doris.nereids.trees.expressions.literal.DateTimeV2Literal; +import org.apache.doris.nereids.trees.expressions.literal.DateV2Literal; +import org.apache.doris.nereids.trees.expressions.literal.DecimalV3Literal; +import org.apache.doris.nereids.trees.expressions.literal.DoubleLiteral; +import org.apache.doris.nereids.trees.expressions.literal.FloatLiteral; +import org.apache.doris.nereids.trees.expressions.literal.IntegerLiteral; +import org.apache.doris.nereids.trees.expressions.literal.Literal; +import org.apache.doris.nereids.trees.expressions.literal.NullLiteral; +import org.apache.doris.nereids.trees.expressions.literal.SmallIntLiteral; +import org.apache.doris.nereids.trees.expressions.literal.StringLiteral; +import org.apache.doris.nereids.trees.expressions.literal.TinyIntLiteral; +import org.apache.doris.nereids.trees.plans.PlaceholderId; +import org.apache.doris.nereids.trees.plans.Plan; +import org.apache.doris.nereids.types.BigIntType; +import org.apache.doris.nereids.types.BooleanType; +import org.apache.doris.nereids.types.DataType; +import org.apache.doris.nereids.types.DateTimeV2Type; +import org.apache.doris.nereids.types.DateV2Type; +import org.apache.doris.nereids.types.DecimalV3Type; +import org.apache.doris.nereids.types.DoubleType; +import org.apache.doris.nereids.types.FloatType; +import org.apache.doris.nereids.types.IntegerType; +import org.apache.doris.nereids.types.NullType; +import org.apache.doris.nereids.types.SmallIntType; +import org.apache.doris.nereids.types.StringType; +import org.apache.doris.nereids.types.TinyIntType; + +import org.apache.arrow.flight.CallStatus; +import org.apache.arrow.flight.FlightRuntimeException; +import org.apache.arrow.flight.FlightStream; +import org.apache.arrow.vector.DateDayVector; +import org.apache.arrow.vector.FieldVector; +import org.apache.arrow.vector.TimeStampVector; +import org.apache.arrow.vector.VarCharVector; +import org.apache.arrow.vector.VectorSchemaRoot; +import org.apache.arrow.vector.types.Types; +import org.apache.arrow.vector.types.pojo.ArrowType; + +import java.math.BigDecimal; +import java.nio.ByteBuffer; +import java.nio.charset.CharacterCodingException; +import java.nio.charset.StandardCharsets; +import java.time.LocalDate; +import java.time.LocalDateTime; +import java.time.ZoneOffset; +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +/** Converts one parameter row into detached values consumed by Nereids placeholder analysis. */ +final class FlightSqlParameters { + private static final long MAX_PARAMETER_BYTES = 1024 * 1024; + static final int MAX_PARAMETERS = 1024; + + private FlightSqlParameters() { + } + + static Map<PlaceholderId, DataType> inferCastTypes(Plan plan) { + Map<PlaceholderId, DataType> types = new HashMap<>(); + plan.foreach(node -> { + ((Plan) node).getExpressions().forEach(expression -> expression.foreach(child -> { + if (child instanceof Cast && ((Cast) child).child() instanceof Placeholder) { + // Only the innermost explicit cast constrains the value supplied for this placeholder. + types.put(((Placeholder) ((Cast) child).child()).getPlaceholderId(), ((Cast) child).getDataType()); + } else if (child instanceof SubqueryExpr) { + types.putAll(inferCastTypes(((SubqueryExpr) child).getQueryPlan())); + } + })); + }); + return types; + } + + static boolean supportsParameterType(ArrowType type) { + switch (Types.getMinorTypeForArrowType(type)) { + case NULL: + case BIT: + case TINYINT: + case SMALLINT: + case INT: + case BIGINT: + case FLOAT4: + case FLOAT8: + case VARCHAR: + case DATEDAY: + case TIMESTAMPSEC: + case TIMESTAMPMILLI: + case TIMESTAMPMICRO: + case TIMESTAMPNANO: + case DECIMAL: + return true; + default: + return false; + } + } + + static void bind(StatementContext context, List<Literal> parameters) { + if (parameters == null || parameters.size() != context.getPlaceholders().size()) { + throw CallStatus.INVALID_ARGUMENT.withDescription("Bind all query parameters before execution") + .toRuntimeException(); + } + for (int i = 0; i < parameters.size(); i++) { + context.getIdToPlaceholderRealExpr().put( + context.getPlaceholders().get(i).getPlaceholderId(), parameters.get(i)); + } + } + + static List<Literal> read(FlightStream stream, int parameterCount) { + List<Literal> parameters = null; + FlightRuntimeException failure = null; + while (stream.next()) { + // Consume the upload even after validation fails so clients can receive the final gRPC status. + if (failure != null) { + continue; + } + try { + VectorSchemaRoot root = stream.getRoot(); + if (root.getFieldVectors().size() != parameterCount) { + throw invalid("Parameter count does not match the prepared query"); + } + if (root.getRowCount() == 0) { + continue; + } + if (root.getRowCount() != 1 || parameters != null) { + throw CallStatus.UNIMPLEMENTED.withDescription("Only one parameter row per binding is supported") + .toRuntimeException(); + } + parameters = convert(root); + } catch (FlightRuntimeException e) { + failure = e; + } catch (RuntimeException e) { + failure = invalid("Invalid parameter value: " + e.getMessage()); + } + } + if (failure != null) { + throw failure; + } + if (parameters == null) { + if (parameterCount == 0) { + return Collections.emptyList(); + } + throw invalid("Parameter upload contains no row"); + } + return parameters; + } + + static List<Literal> convert(VectorSchemaRoot root) { + if (root.getFieldVectors().size() > MAX_PARAMETERS) { + throw invalid("Too many query parameters (maximum 1024)"); + } + long bytes = 0; + List<Literal> parameters = new ArrayList<>(); + for (FieldVector vector : root.getFieldVectors()) { + bytes += vector.getBufferSize(); + if (bytes > MAX_PARAMETER_BYTES) { + throw invalid("Query parameters exceed the 1 MiB binding limit"); + } + parameters.add(literal(vector)); + } + return parameters; + } + + private static Literal literal(FieldVector vector) { + if (vector.getField().getDictionary() != null) { + throw CallStatus.UNIMPLEMENTED.withDescription("Dictionary encoded parameters are not supported") + .toRuntimeException(); + } + DataType type; + switch (vector.getMinorType()) { + case NULL: + type = NullType.INSTANCE; + break; + case BIT: + type = BooleanType.INSTANCE; + break; + case TINYINT: + type = TinyIntType.INSTANCE; + break; + case SMALLINT: + type = SmallIntType.INSTANCE; + break; + case INT: + type = IntegerType.INSTANCE; + break; + case BIGINT: + type = BigIntType.INSTANCE; + break; + case FLOAT4: + type = FloatType.INSTANCE; + break; + case FLOAT8: + type = DoubleType.INSTANCE; + break; + case VARCHAR: + type = StringType.INSTANCE; + break; + case DATEDAY: + type = DateV2Type.INSTANCE; + break; + case TIMESTAMPSEC: + case TIMESTAMPMILLI: + case TIMESTAMPMICRO: + case TIMESTAMPNANO: + type = DateTimeV2Type.of(6); + break; + case DECIMAL: + ArrowType.Decimal decimal = (ArrowType.Decimal) vector.getField().getType(); + if (decimal.getScale() < 0 || decimal.getScale() > decimal.getPrecision()) { + throw invalid("Unsupported decimal parameter scale"); + } + type = DecimalV3Type.createDecimalV3Type(decimal.getPrecision(), decimal.getScale()); + break; + default: + throw CallStatus.UNIMPLEMENTED.withDescription( + "Unsupported query parameter type: " + vector.getField().getType()).toRuntimeException(); + } + if (vector.isNull(0)) { + return new NullLiteral(type); + } + Object value = vector.getObject(0); + switch (vector.getMinorType()) { + case BIT: return BooleanLiteral.of((Boolean) value); + case TINYINT: return new TinyIntLiteral(((Number) value).byteValue()); + case SMALLINT: return new SmallIntLiteral(((Number) value).shortValue()); + case INT: return new IntegerLiteral(((Number) value).intValue()); + case BIGINT: return new BigIntLiteral(((Number) value).longValue()); + case FLOAT4: + case FLOAT8: + double number = ((Number) value).doubleValue(); + if (!Double.isFinite(number)) { + throw invalid("Non-finite floating point parameters are not supported"); + } + return type instanceof FloatType ? new FloatLiteral((float) number) : new DoubleLiteral(number); + case VARCHAR: + try { + // Reject malformed UTF-8 instead of silently replacing bytes in a bound predicate. + return new StringLiteral(StandardCharsets.UTF_8.newDecoder() + .decode(ByteBuffer.wrap(((VarCharVector) vector).get(0))).toString()); + } catch (CharacterCodingException e) { + throw invalid("String parameter is not valid UTF-8"); + } + case DECIMAL: return new DecimalV3Literal((DecimalV3Type) type, (BigDecimal) value); + case DATEDAY: + LocalDate date = LocalDate.ofEpochDay(((DateDayVector) vector).get(0)); + checkYear(date.getYear()); + return new DateV2Literal(date.getYear(), date.getMonthValue(), date.getDayOfMonth()); + default: + long timestamp = ((TimeStampVector) vector).get(0); + ArrowType.Timestamp timestampType = (ArrowType.Timestamp) vector.getField().getType(); + long units; + switch (timestampType.getUnit()) { + case SECOND: + units = 1; + break; + case MILLISECOND: + units = 1000; + break; + case MICROSECOND: + units = 1000000; + break; + case NANOSECOND: + units = 1000000000; + break; + default: + throw invalid("Unsupported timestamp unit"); + } + long nanos = Math.floorMod(timestamp, units) * (1000000000 / units); + if (nanos % 1000 != 0) { + throw invalid("Timestamp parameter exceeds microsecond precision"); + } + LocalDateTime time = LocalDateTime.ofEpochSecond( + Math.floorDiv(timestamp, units), (int) nanos, ZoneOffset.UTC); + checkYear(time.getYear()); + return new DateTimeV2Literal((DateTimeV2Type) type, time.getYear(), time.getMonthValue(), Review Comment: [P2] Preserve timezone-aware Arrow timestamp instants when binding. For a non-UTC session, `SELECT CAST(? AS TIMESTAMPTZ)` advertises a timezone-bearing Arrow timestamp, whose integer is already relative to the UTC epoch. This branch makes UTC wall-clock fields into a `DateTimeV2Literal`; the subsequent DATETIMEV2-to-TIMESTAMPTZ cast interprets those fields in the session timezone and shifts the instant again. For example, epoch zero in an Asia/Shanghai session becomes 1969-12-31 16:00:00 UTC, so TIMESTAMPTZ predicates can miss rows. Keep timezone-bearing uploads as an instant (or explicitly convert from UTC) and test a non-UTC session. -- 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] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
