FlightSqlParameters.java
// 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.TimestampTzLiteral;
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.TimeStampTzType;
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 TIMESTAMPSECTZ:
case TIMESTAMPMILLITZ:
case TIMESTAMPMICROTZ:
case TIMESTAMPNANOTZ:
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 TIMESTAMPSECTZ:
case TIMESTAMPMILLITZ:
case TIMESTAMPMICROTZ:
case TIMESTAMPNANOTZ:
// Arrow Java uses TZ vectors even for an empty annotation, which still means wall-clock time.
type = ((ArrowType.Timestamp) vector.getField().getType()).getTimezone().isEmpty()
? DateTimeV2Type.of(6) : TimeStampTzType.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());
if (type instanceof TimeStampTzType) {
// Zoned Arrow timestamps already encode a UTC instant. A DATETIMEV2 literal would
// reinterpret these fields in the session timezone when cast to TIMESTAMPTZ.
return new TimestampTzLiteral((TimeStampTzType) type, time.getYear(), time.getMonthValue(),
time.getDayOfMonth(), time.getHour(), time.getMinute(), time.getSecond(),
time.getNano() / 1000);
}
return new DateTimeV2Literal((DateTimeV2Type) type, time.getYear(), time.getMonthValue(),
time.getDayOfMonth(), time.getHour(), time.getMinute(), time.getSecond(),
time.getNano() / 1000);
}
}
private static void checkYear(int year) {
if (year < 0 || year > 9999) {
throw invalid("Date parameter is outside the supported year range 0000..9999");
}
}
private static FlightRuntimeException invalid(String message) {
return CallStatus.INVALID_ARGUMENT.withDescription(message).toRuntimeException();
}
}