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]

Reply via email to