This is an automated email from the ASF dual-hosted git repository.
yiguolei pushed a commit to branch branch-4.1
in repository https://gitbox.apache.org/repos/asf/doris.git
The following commit(s) were added to refs/heads/branch-4.1 by this push:
new bd4ce9fbc0b [fix](flight-sql) Support prepared query parameter binding
on branch-4.1 (#68768)
bd4ce9fbc0b is described below
commit bd4ce9fbc0bbbbcdb564c3acff89fe70e0a718f2
Author: Gabriel <[email protected]>
AuthorDate: Fri Oct 9 14:37:12 2026 +0800
[fix](flight-sql) Support prepared query parameter binding on branch-4.1
(#68768)
### What problem does this PR solve?
Flight SQL prepared queries containing `?` fail because Prepare rejects
placeholders and `acceptPutPreparedStatementQuery` returns
`UNIMPLEMENTED`. For example, ADBC cannot bind an integer to `SELECT
CAST(? AS BIGINT)` or execute a range predicate with two bound values.
This change implements single-row typed parameter binding on branch-4.1.
Prepare analyzes explicit casts and comparison columns to advertise
parameter types and a best-effort result schema. DoPut converts Arrow
values into detached Nereids literals, and execution binds those
literals to a fresh statement context. Result schemas are reanalyzed
after binding. Values are never interpolated into SQL text.
Arrow JDBC receives typed parameter metadata and a result schema before
binding, so it can construct parameter vectors and classify SELECT
correctly. Types that cannot be inferred remain unknown; JDBC callers
should use an explicit CAST for queries such as `SELECT ?`.
Timezone-bearing timestamp uploads preserve their UTC instants as
TIMESTAMPTZ literals; timestamps without a timezone retain wall-clock
semantics. Unsupported parameter vector types are not advertised as
bindable. Invalid SQL returns `INVALID_ARGUMENT` consistently through
Prepare and GetSchema.
Failed rebinds invalidate previous values. Handle ownership, concurrent
closure, upload versions, and per-binding/session memory limits protect
retained state. Multi-row uploads and unsupported types fail explicitly.
Queries requiring forwarding to the master FE reject bound parameters
because the forwarding protocol cannot carry typed bindings.
The PR also removes an obsolete duplicate JWT version property,
preserving the effective dependency version. The target branch already
contains the dependency declaration cleanup.
### Release note
Support single-row scalar parameter binding for Flight SQL prepared
queries, including repeated execution and NULL values.
### Check List (For Author)
- Test
- [x] Regression test
- [x] Unit Test
- [x] Manual test
- Behavior changed:
- [x] Yes. Parameterized Flight SQL queries can bind and execute
supported scalar values.
- Does this need documentation?
- [x] No. This implements the existing prepared-query protocol.
Validation:
- Full FE Checkstyle passed (`mvn -B clean checkstyle:check`).
- Focused compilation of all 11 changed Java source/test files passed
with `javac --release 8` using existing FE dependency artifacts. Java
8-compatible Guava copies preserve immutable parameter snapshots.
- 23 new tests and 33 existing Flight SQL schema tests passed, including
invalid-SQL status checks over Flight RPC and connection recovery.
- The new regression suite passed against a local FE service test
harness and a real BE, including actual JDBC `PreparedStatement`
execution, result fetching, rebinding, NULL, mixed types, range
predicates, and recovery after invalid uploads. TIMESTAMPTZ predicates
are checked under UTC, Asia/Shanghai and America/New_York, including a
daylight-saving transition.
- ADBC 1.11.0 with PyArrow 23.0.1 passed 18 concurrent connections and
540 parameterized queries, with result verification plus NULL and
connection-reuse checks. A further 72 timestamp result checks cover all
four Arrow units, three session timezones, two input timezone
annotations, negative epochs and NULL values.
The local integration harness uses one FE service and one real BE; a
full multi-FE cluster run and full FE build are left to CI.
### Check List (For Reviewer who merge this PR)
- [ ] Confirm the release note
- [ ] Confirm test cases
- [ ] Confirm document
- [ ] Add branch pick label
---
.../nereids/rules/analysis/ExpressionAnalyzer.java | 5 +-
.../java/org/apache/doris/qe/ConnectContext.java | 176 +++++-
.../java/org/apache/doris/qe/StmtExecutor.java | 6 +
.../arrowflight/DorisFlightSqlProducer.java | 132 +++-
.../arrowflight/FlightSqlConnectProcessor.java | 38 ++
.../service/arrowflight/FlightSqlParameters.java | 336 +++++++++++
.../service/arrowflight/FlightSqlQuerySchema.java | 65 +-
.../sessions/FlightSqlConnectPoolMgr.java | 3 +
.../arrowflight/DorisFlightSqlSchemaTest.java | 15 +-
.../arrowflight/FlightSqlParameterBindingTest.java | 41 ++
.../arrowflight/FlightSqlPreparedQueryTest.java | 666 +++++++++++++++++++++
.../test_prepared_query_parameters.groovy | 223 +++++++
12 files changed, 1648 insertions(+), 58 deletions(-)
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/ExpressionAnalyzer.java
b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/ExpressionAnalyzer.java
index 73049cf58c9..149f6101494 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/ExpressionAnalyzer.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/ExpressionAnalyzer.java
@@ -813,7 +813,10 @@ public class ExpressionAnalyzer extends
SubExprAnalyzer<ExpressionRewriteContext
private void registerPlaceholderIdToSlot(ComparisonPredicate cp,
ExpressionRewriteContext context, Expression left,
Expression right) {
if (ConnectContext.get() != null
- && ConnectContext.get().getCommand() ==
MysqlCommand.COM_STMT_EXECUTE) {
+ && (ConnectContext.get().getCommand() ==
MysqlCommand.COM_STMT_EXECUTE
+ // Flight Prepare needs the bound column type before
comparison coercion erases that constraint.
+ || (ConnectContext.get().getConnectType() ==
ConnectContext.ConnectType.ARROW_FLIGHT_SQL
+ &&
context.cascadesContext.getStatementContext().isPrepareStage()))) {
// Used to replace expression in ShortCircuit plan
if (cp.right() instanceof Placeholder && left instanceof
SlotReference) {
PlaceholderId id = ((Placeholder)
cp.right()).getPlaceholderId();
diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/ConnectContext.java
b/fe/fe-core/src/main/java/org/apache/doris/qe/ConnectContext.java
index 674c2a99548..e51a0031676 100644
--- a/fe/fe-core/src/main/java/org/apache/doris/qe/ConnectContext.java
+++ b/fe/fe-core/src/main/java/org/apache/doris/qe/ConnectContext.java
@@ -87,6 +87,7 @@ import org.apache.doris.transaction.TransactionEntry;
import org.apache.doris.transaction.TransactionStatus;
import com.google.common.base.Strings;
+import com.google.common.collect.ImmutableList;
import com.google.common.collect.Lists;
import com.google.common.collect.Maps;
import lombok.Getter;
@@ -156,12 +157,19 @@ public class ConnectContext {
// for arrow flight
protected volatile String peerIdentity;
private final Map<String, PreparedQuery> preparedQuerys = new HashMap<>();
+ private static final long MAX_FLIGHT_PARAMETER_BYTES = 16L * 1024 * 1024;
+ private long flightParameterBytes;
+ private boolean flightPreparedQueriesClosed;
private static class PreparedQuery {
private final String sql;
- private final String catalog;
- private final String database;
- private final Schema schema;
+ private String catalog;
+ private String database;
+ private Schema schema;
+ private int parameterCount;
+ private List<Literal> parameters;
+ private long bindingVersion;
+ private long parameterBytes;
private PreparedQuery(String sql, String catalog, String database,
Schema schema) {
this.sql = sql;
@@ -427,7 +435,7 @@ public class ConnectContext {
}
resetSessionVariable();
userVars = new HashMap<>();
- preparedQuerys.clear();
+ clearPreparedQueries();
preparedStatementContextMap.clear();
runningQuery = null;
queryId = null;
@@ -924,35 +932,157 @@ public class ConnectContext {
this.loginTime = System.currentTimeMillis();
}
- public synchronized void addPreparedQuery(String preparedStatementId,
String preparedQuery) {
- addPreparedQuery(preparedStatementId, preparedQuery, null);
+ public void addPreparedQuery(String preparedStatementId, String
preparedQuery) {
+ synchronized (preparedQuerys) {
+ addPreparedQuery(preparedStatementId, preparedQuery, null);
+ }
}
- public synchronized void addPreparedQuery(String preparedStatementId,
String preparedQuery, Schema schema) {
- preparedQuerys.put(preparedStatementId,
- new PreparedQuery(preparedQuery, getDefaultCatalog(),
getDatabase(), schema));
+ public void addPreparedQuery(String preparedStatementId, String
preparedQuery, Schema schema) {
+ synchronized (preparedQuerys) {
+ if (flightPreparedQueriesClosed) {
+ throw new IllegalStateException("Flight SQL session is
closed");
+ }
+ removePreparedQuery(preparedStatementId);
+ preparedQuerys.put(preparedStatementId,
+ new PreparedQuery(preparedQuery, getDefaultCatalog(),
getDatabase(), schema));
+ }
}
- public synchronized Schema getPreparedQuerySchema(String
preparedStatementId) {
- PreparedQuery query = preparedQuerys.get(preparedStatementId);
- return query == null ? null : query.schema;
+ public void addPreparedQuery(String id, String sql, Schema schema, int
parameterCount) {
+ synchronized (preparedQuerys) {
+ addPreparedQuery(id, sql, schema);
+ preparedQuerys.get(id).parameterCount = parameterCount;
+ }
}
- public synchronized String getPreparedQuery(String preparedStatementId) {
- PreparedQuery query = preparedQuerys.get(preparedStatementId);
- if (query == null) {
- return null;
+ public int getPreparedQueryParameterCount(String id) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(id);
+ return query == null ? -1 : query.parameterCount;
}
- // A handle must not execute unqualified SQL in a different namespace
than its advertised schema.
- if (!Objects.equals(query.catalog, getDefaultCatalog()) ||
!Objects.equals(query.database, getDatabase())) {
- preparedQuerys.remove(preparedStatementId);
- return null;
+ }
+
+ public List<Literal> getPreparedQueryParameters(String id) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(id);
+ return query == null ? null : query.parameters;
+ }
+ }
+
+ public long beginPreparedQueryBinding(String id) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(id);
+ if (query == null) {
+ return -1;
+ }
+ flightParameterBytes -= query.parameterBytes;
+ query.parameterBytes = 0;
+ query.parameters = null;
+ return ++query.bindingVersion;
}
- return query.sql;
}
- public synchronized void removePreparedQuery(String preparedStatementId) {
- preparedQuerys.remove(preparedStatementId);
+ public boolean isPreparedQueryBindingCurrent(String id, long version) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(id);
+ return query != null && query.bindingVersion == version;
+ }
+ }
+
+ public void setPreparedQueryParameters(String id, List<Literal>
parameters, Schema schema) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(id);
+ if (query == null) {
+ throw new IllegalStateException("Prepared statement expired");
+ }
+ long bytes = parameters == null ? 0 : parameters.size() * 128L;
+ if (parameters != null) {
+ for (Literal parameter : parameters) {
+ if (parameter.getValue() instanceof String) {
+ bytes += 2L * ((String) parameter.getValue()).length();
+ }
+ }
+ }
+ // A per-upload limit alone allows many handles to retain
unbounded parameter memory in one session.
+ if (flightParameterBytes - query.parameterBytes + bytes >
MAX_FLIGHT_PARAMETER_BYTES) {
+ throw new IllegalArgumentException("Prepared query parameters
exceed the 16 MiB session limit");
+ }
+ flightParameterBytes += bytes - query.parameterBytes;
+ query.parameterBytes = bytes;
+ // Keep a detached immutable snapshot using a Java 8-compatible
copy implementation.
+ query.parameters = parameters == null ? null :
ImmutableList.copyOf(parameters);
+ query.schema = schema;
+ }
+ }
+
+ public boolean setPreparedQueryParameters(String id, List<Literal>
parameters, Schema schema, long version) {
+ synchronized (preparedQuerys) {
+ if (!isPreparedQueryBindingCurrent(id, version)) {
+ return false;
+ }
+ setPreparedQueryParameters(id, parameters, schema);
+ return true;
+ }
+ }
+
+ public boolean refreshPreparedQueryNamespace(String id) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(id);
+ if (query == null) {
+ return false;
+ }
+ if (query.parameterCount == 0) {
+ query.catalog = getDefaultCatalog();
+ query.database = getDatabase();
+ }
+ return true;
+ }
+ }
+
+ public Schema getPreparedQuerySchema(String preparedStatementId) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(preparedStatementId);
+ return query == null ? null : query.schema;
+ }
+ }
+
+ public String getPreparedQuery(String preparedStatementId) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.get(preparedStatementId);
+ if (query == null) {
+ return null;
+ }
+ // A handle must not execute unqualified SQL in a different
namespace than its advertised schema.
+ if (!Objects.equals(query.catalog, getDefaultCatalog()) ||
!Objects.equals(query.database, getDatabase())) {
+ removePreparedQuery(preparedStatementId);
+ return null;
+ }
+ return query.sql;
+ }
+ }
+
+ public void closePreparedQueries() {
+ synchronized (preparedQuerys) {
+ flightPreparedQueriesClosed = true;
+ clearPreparedQueries();
+ }
+ }
+
+ public void clearPreparedQueries() {
+ synchronized (preparedQuerys) {
+ preparedQuerys.clear();
+ flightParameterBytes = 0;
+ }
+ }
+
+ public void removePreparedQuery(String preparedStatementId) {
+ synchronized (preparedQuerys) {
+ PreparedQuery query = preparedQuerys.remove(preparedStatementId);
+ if (query != null) {
+ flightParameterBytes -= query.parameterBytes;
+ }
+ }
}
public void setRunningQuery(String runningQuery) {
diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/StmtExecutor.java
b/fe/fe-core/src/main/java/org/apache/doris/qe/StmtExecutor.java
index 30a943bba1a..b1490b88529 100644
--- a/fe/fe-core/src/main/java/org/apache/doris/qe/StmtExecutor.java
+++ b/fe/fe-core/src/main/java/org/apache/doris/qe/StmtExecutor.java
@@ -1183,6 +1183,12 @@ public class StmtExecutor {
}
private void forwardToMaster() throws Exception {
+ // COM_QUERY forwarding carries SQL text but not Flight's typed
placeholder bindings.
+ if (context.getConnectType() == ConnectType.ARROW_FLIGHT_SQL
+ && !statementContext.getIdToPlaceholderRealExpr().isEmpty()) {
+ throw new UserException("Flight SQL queries with bound parameters
cannot be forwarded; "
+ + "connect to master FE");
+ }
masterOpExecutor = new MasterOpExecutor(originStmt, context,
redirectStatus, isQuery());
if (LOG.isDebugEnabled()) {
LOG.debug("need to transfer to Master. stmt: {}",
context.getStmtId());
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/DorisFlightSqlProducer.java
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/DorisFlightSqlProducer.java
index 2aae87d527b..13b3f506fb9 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/DorisFlightSqlProducer.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/DorisFlightSqlProducer.java
@@ -27,6 +27,7 @@ import org.apache.doris.common.util.DebugUtil;
import org.apache.doris.common.util.Util;
import org.apache.doris.mysql.MysqlCommand;
import org.apache.doris.nereids.glue.LogicalPlanAdapter;
+import org.apache.doris.nereids.trees.expressions.literal.Literal;
import org.apache.doris.qe.ConnectContext;
import org.apache.doris.qe.QueryState.MysqlStateType;
import org.apache.doris.qe.StmtExecutor;
@@ -186,10 +187,12 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
final StreamListener<Result> listener) {
executorService.submit(() -> {
try {
- String[] handleParts =
request.getPreparedStatementHandle().toStringUtf8().split(":");
- String executedPeerIdentity = handleParts[0];
- String preparedStatementId = handleParts[1];
-
flightSessionsManager.getConnectContext(executedPeerIdentity).removePreparedQuery(preparedStatementId);
+ String preparedStatementId = preparedStatementId(context,
request.getPreparedStatementHandle());
+ ConnectContext connection =
flightSessionsManager.getConnectContext(context.peerIdentity());
+ // Closing one handle serializes with its execution; closing a
session can still cancel immediately.
+ synchronized (connection) {
+ connection.removePreparedQuery(preparedStatementId);
+ }
} catch (final Throwable e) {
listener.onError(e);
return;
@@ -211,6 +214,11 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
private Pair<FlightInfo, StmtExecutor> executeQueryStatementLocked(String
peerIdentity,
ConnectContext connectContext, String query,
final FlightDescriptor descriptor) {
+ return executeQueryStatementLocked(peerIdentity, connectContext,
query, descriptor, null);
+ }
+
+ private Pair<FlightInfo, StmtExecutor> executeQueryStatementLocked(String
peerIdentity,
+ ConnectContext connectContext, String query, FlightDescriptor
descriptor, List<Literal> parameters) {
try {
Preconditions.checkState(null != connectContext);
Preconditions.checkState(!query.isEmpty());
@@ -222,7 +230,11 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
connectContext.getFlightSqlChannel().reset();
connectContext.clearFlightSqlEndpointsLocations();
try (FlightSqlConnectProcessor flightSQLConnectProcessor = new
FlightSqlConnectProcessor(connectContext)) {
- flightSQLConnectProcessor.handleQuery(query);
+ if (parameters == null) {
+ flightSQLConnectProcessor.handleQuery(query);
+ } else {
+ flightSQLConnectProcessor.handleQuery(query, parameters);
+ }
if (connectContext.getState().getStateType() ==
MysqlStateType.ERR) {
throw new RuntimeException("after executeQueryStatement
handleQuery");
}
@@ -348,8 +360,10 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
ConnectContext connection =
flightSessionsManager.getConnectContext(context.peerIdentity());
synchronized (connection) {
Pair<String, Schema> prepared = preparedQuery(connection, context,
command);
+ String preparedId = preparedStatementId(context,
command.getPreparedStatementHandle());
Pair<FlightInfo, StmtExecutor> result =
executeQueryStatementLocked(
- context.peerIdentity(), connection, prepared.getLeft(),
descriptor);
+ context.peerIdentity(), connection, prepared.getLeft(),
descriptor,
+ connection.getPreparedQueryParameters(preparedId));
FlightInfo info = result.getLeft();
String id = command.getPreparedStatementHandle().toStringUtf8()
.substring(context.peerIdentity().length() + 1);
@@ -361,7 +375,8 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
?
FlightSqlQuerySchema.matchesExecutionSchema(prepared.getRight(),
info.getSchema(),
((LogicalPlanAdapter)
executor.getParsedStmt()).getColLabels())
: prepared.getRight().equals(info.getSchema());
- if (!matches) {
+ // Session teardown can invalidate the handle during execution
without taking the connection monitor.
+ if (!matches || !connection.refreshPreparedQueryNamespace(id)) {
connection.removePreparedQuery(id);
try {
if (executor != null) {
@@ -377,9 +392,6 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
throw CallStatus.NOT_FOUND.withDescription("Prepared statement
schema changed; prepare again")
.toRuntimeException();
}
- // A successful USE/SWITCH in this statement may change its own
namespace. Rebind
- // only this handle; unrelated handles must still reject an
external namespace change.
- connection.addPreparedQuery(id, prepared.getLeft(),
prepared.getRight());
return info;
}
}
@@ -396,25 +408,27 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
final CallContext context, final FlightDescriptor descriptor) {
ConnectContext connection =
flightSessionsManager.getConnectContext(context.peerIdentity());
synchronized (connection) {
+ String id = preparedStatementId(context,
command.getPreparedStatementHandle());
+ String query = connection.getPreparedQuery(id);
+ if (query != null && connection.getPreparedQueryParameterCount(id)
> 0
+ && connection.getPreparedQueryParameters(id) == null) {
+ // Metadata discovery does not execute a query and must work
before values are bound.
+ return new SchemaResult(prepareQuerySchema(connection,
query).first);
+ }
return new SchemaResult(preparedQuery(connection, context,
command).getRight());
}
}
private Pair<String, Schema> preparedQuery(ConnectContext connection,
CallContext context,
CommandPreparedStatementQuery command) {
- String prefix = context.peerIdentity() + ":";
- String handle = command.getPreparedStatementHandle().toStringUtf8();
- if (!handle.startsWith(prefix)) {
- throw CallStatus.INVALID_ARGUMENT.withDescription("Invalid
prepared statement handle").toRuntimeException();
- }
- String id = handle.substring(prefix.length());
+ String id = preparedStatementId(context,
command.getPreparedStatementHandle());
String query = connection.getPreparedQuery(id);
if (query == null) {
throw CallStatus.NOT_FOUND
.withDescription("Prepared statement expired; prepare
again in the current namespace")
.toRuntimeException();
}
- Schema schema = analyzeQuerySchema(connection, query);
+ Schema schema = analyzeQuerySchema(connection, query,
connection.getPreparedQueryParameters(id));
// Execution reparses SQL using the current session. Never silently
replace the schema
// advertised by Prepare when settings such as sql_mode or time_zone
change its result.
if (!schema.equals(connection.getPreparedQuerySchema(id))) {
@@ -425,9 +439,22 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
return Pair.of(query, schema);
}
+ private String preparedStatementId(CallContext context, ByteString
handleBytes) {
+ String prefix = context.peerIdentity() + ":";
+ String handle = handleBytes.toStringUtf8();
+ if (!handle.startsWith(prefix) || handle.length() == prefix.length()) {
+ throw CallStatus.INVALID_ARGUMENT.withDescription("Invalid
prepared statement handle").toRuntimeException();
+ }
+ return handle.substring(prefix.length());
+ }
+
private Schema analyzeQuerySchema(ConnectContext context, String query) {
+ return analyzeQuerySchema(context, query, null);
+ }
+
+ private Schema analyzeQuerySchema(ConnectContext context, String query,
List<Literal> parameters) {
try {
- return FlightSqlQuerySchema.analyze(context, query);
+ return FlightSqlQuerySchema.analyze(context, query, parameters);
} catch (FlightRuntimeException e) {
throw e;
} catch (Exception e) {
@@ -436,6 +463,18 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
}
}
+ private org.apache.doris.common.Pair<Schema, Schema>
prepareQuerySchema(ConnectContext context, String query) {
+ try {
+ return FlightSqlQuerySchema.prepare(context, query);
+ } catch (FlightRuntimeException e) {
+ throw e;
+ } catch (Exception e) {
+ // Invalid SQL must report the same client error through Prepare
and GetSchema.
+ throw CallStatus.INVALID_ARGUMENT.withDescription("Cannot
determine query schema: " + e.getMessage())
+ .withCause(e).toRuntimeException();
+ }
+ }
+
@Override
public void close() throws Exception {
AutoCloseables.close(rootAllocator);
@@ -448,16 +487,18 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
private ActionCreatePreparedStatementResult
buildCreatePreparedStatementResult(ByteString handle,
Schema parameterSchema, Schema metaData) {
- Preconditions.checkState(!Objects.isNull(metaData));
final ByteString bytes = Objects.isNull(parameterSchema) ?
ByteString.EMPTY
: ByteString.copyFrom(serializeMetadata(parameterSchema));
return ActionCreatePreparedStatementResult.newBuilder()
-
.setDatasetSchema(ByteString.copyFrom(serializeMetadata(metaData))).setParameterSchema(bytes)
+ .setDatasetSchema(metaData == null ? ByteString.EMPTY
+ : ByteString.copyFrom(serializeMetadata(metaData)))
+ .setParameterSchema(bytes)
.setPreparedStatementHandle(handle).build();
}
@Override
- public void createPreparedStatement(final
ActionCreatePreparedStatementRequest request, final CallContext context,
+ public void createPreparedStatement(final
ActionCreatePreparedStatementRequest request,
+ final CallContext context,
final StreamListener<Result> listener) {
executorService.submit(() -> {
ConnectContext connectContext = null;
@@ -468,12 +509,15 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
String query = request.getQuery();
// ADBC ExecuteSchema reads this dataset schema directly
without calling GetSchema.
// Analyze before registering a handle so failed
preparation does not retain a query.
- Schema schema = analyzeQuerySchema(connectContext, query);
+ org.apache.doris.common.Pair<Schema, Schema> prepared =
prepareQuerySchema(
+ connectContext, query);
+ Schema schema = prepared.first;
preparedStatementId = UUID.randomUUID().toString();
ByteString handle =
ByteString.copyFromUtf8(context.peerIdentity() + ":" + preparedStatementId);
Result result = new
Result(Any.pack(buildCreatePreparedStatementResult(handle,
- new Schema(Collections.emptyList()),
schema)).toByteArray());
- connectContext.addPreparedQuery(preparedStatementId,
query, schema);
+ prepared.second, schema)).toByteArray());
+ connectContext.addPreparedQuery(preparedStatementId,
query, schema,
+ prepared.second.getFields().size());
listener.onNext(result);
listener.onCompleted();
}
@@ -530,8 +574,44 @@ public class DorisFlightSqlProducer implements
FlightSqlProducer, AutoCloseable
@Override
public Runnable
acceptPutPreparedStatementQuery(CommandPreparedStatementQuery command,
CallContext context,
FlightStream flightStream, StreamListener<PutResult> ackStream) {
- throw
CallStatus.UNIMPLEMENTED.withDescription("acceptPutPreparedStatementQuery
unimplemented")
- .toRuntimeException();
+ return () -> {
+ try {
+ ConnectContext connection =
flightSessionsManager.getConnectContext(context.peerIdentity());
+ String id = preparedStatementId(context,
command.getPreparedStatementHandle());
+ int count;
+ long version;
+ synchronized (connection) {
+ if (connection.getPreparedQuery(id) == null) {
+ throw CallStatus.NOT_FOUND.withDescription("Prepared
statement expired").toRuntimeException();
+ }
+ count = connection.getPreparedQueryParameterCount(id);
+ // A failed rebind must not leave the previous values
executable.
+ version = connection.beginPreparedQueryBinding(id);
+ if (version < 0) {
+ throw CallStatus.NOT_FOUND.withDescription("Prepared
statement expired").toRuntimeException();
+ }
+ }
+ List<Literal> parameters =
FlightSqlParameters.read(flightStream, count);
+ synchronized (connection) {
+ String query = connection.getPreparedQuery(id);
+ if (query == null ||
!connection.isPreparedQueryBindingCurrent(id, version)) {
+ throw CallStatus.CANCELLED.withDescription("Prepared
statement binding was superseded")
+ .toRuntimeException();
+ }
+ Schema schema = analyzeQuerySchema(connection, query,
parameters);
+ if (!connection.setPreparedQueryParameters(id, parameters,
schema, version)) {
+ throw CallStatus.CANCELLED.withDescription("Prepared
statement binding was superseded")
+ .toRuntimeException();
+ }
+ }
+ ackStream.onCompleted();
+ } catch (FlightRuntimeException e) {
+ ackStream.onError(e);
+ } catch (Exception e) {
+
ackStream.onError(CallStatus.INVALID_ARGUMENT.withDescription("Cannot bind
query parameters: "
+ + e.getMessage()).withCause(e).toRuntimeException());
+ }
+ };
}
@Override
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessor.java
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessor.java
index 5942424750c..87f00e763b8 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessor.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlConnectProcessor.java
@@ -18,16 +18,20 @@
package org.apache.doris.service.arrowflight;
import org.apache.doris.analysis.Expr;
+import org.apache.doris.analysis.StatementBase;
import org.apache.doris.common.ConnectionException;
import org.apache.doris.common.ErrorCode;
import org.apache.doris.common.ErrorReport;
import org.apache.doris.common.Status;
import org.apache.doris.common.util.DebugUtil;
import org.apache.doris.mysql.MysqlCommand;
+import org.apache.doris.nereids.glue.LogicalPlanAdapter;
+import org.apache.doris.nereids.trees.expressions.literal.Literal;
import org.apache.doris.proto.InternalService;
import org.apache.doris.proto.Types;
import org.apache.doris.qe.ConnectContext;
import org.apache.doris.qe.ConnectProcessor;
+import org.apache.doris.qe.SessionVariable;
import org.apache.doris.qe.StmtExecutor;
import org.apache.doris.rpc.BackendServiceProxy;
import org.apache.doris.rpc.RpcException;
@@ -61,6 +65,7 @@ import java.util.concurrent.TimeoutException;
public class FlightSqlConnectProcessor extends ConnectProcessor implements
AutoCloseable {
private static final Logger LOG =
LogManager.getLogger(FlightSqlConnectProcessor.class);
private Schema arrowSchema;
+ private List<Literal> parameters;
public FlightSqlConnectProcessor(ConnectContext context) {
super(context);
@@ -100,6 +105,39 @@ public class FlightSqlConnectProcessor extends
ConnectProcessor implements AutoC
super.handleQuery(query);
}
+ public void handleQuery(String query, List<Literal> parameters) throws
ConnectionException {
+ this.parameters = parameters;
+ try {
+ handleQuery(query);
+ } finally {
+ this.parameters = null;
+ }
+ }
+
+ @Override
+ protected List<StatementBase> parseWithFallback(String originStmt, String
convertedStmt,
+ SessionVariable sessionVariable) throws ConnectionException {
+ List<StatementBase> statements = super.parseWithFallback(originStmt,
convertedStmt, sessionVariable);
+ if (parameters != null && !parameters.isEmpty() && statements != null)
{
+ try {
+ if (statements.size() != 1 || !(statements.get(0) instanceof
LogicalPlanAdapter)) {
+ throw new IllegalArgumentException("Parameters require a
single query statement");
+ }
+ // Reparse on every execution so bindings cannot retain
planner state from a previous query.
+ FlightSqlParameters.bind(((LogicalPlanAdapter)
statements.get(0)).getStatementContext(), parameters);
+ } catch (RuntimeException e) {
+ handleQueryException(e, originStmt, null, null);
+ for (StatementBase statement : statements) {
+ if (statement instanceof LogicalPlanAdapter) {
+ ((LogicalPlanAdapter)
statement).getStatementContext().close();
+ }
+ }
+ return null;
+ }
+ }
+ return statements;
+ }
+
// TODO
// private void handleInitDb() {
// handleInitDb(fullDbName);
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlParameters.java
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlParameters.java
new file mode 100644
index 00000000000..59d21fc71ce
--- /dev/null
+++
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlParameters.java
@@ -0,0 +1,336 @@
+// 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();
+ }
+}
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java
index 678cab055a8..0ed2b77e92f 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/FlightSqlQuerySchema.java
@@ -27,6 +27,7 @@ import org.apache.doris.catalog.ScalarType;
import org.apache.doris.catalog.StructField;
import org.apache.doris.catalog.StructType;
import org.apache.doris.catalog.Type;
+import org.apache.doris.common.Pair;
import org.apache.doris.datasource.CatalogIf;
import org.apache.doris.datasource.es.EsExternalCatalog;
import org.apache.doris.datasource.lance.LanceExternalCatalog;
@@ -38,7 +39,11 @@ import org.apache.doris.nereids.glue.LogicalPlanAdapter;
import org.apache.doris.nereids.parser.NereidsParser;
import org.apache.doris.nereids.parser.SqlDialectHelper;
import org.apache.doris.nereids.rules.rewrite.CheckPrivileges;
+import org.apache.doris.nereids.trees.expressions.Placeholder;
import org.apache.doris.nereids.trees.expressions.Slot;
+import org.apache.doris.nereids.trees.expressions.literal.Literal;
+import org.apache.doris.nereids.trees.expressions.literal.NullLiteral;
+import org.apache.doris.nereids.trees.plans.PlaceholderId;
import org.apache.doris.nereids.trees.plans.Plan;
import org.apache.doris.nereids.trees.plans.PrepareCommandPlanner;
import org.apache.doris.nereids.trees.plans.commands.AlterTableCommand;
@@ -68,6 +73,8 @@ import
org.apache.doris.nereids.trees.plans.commands.insert.InsertOverwriteTable
import org.apache.doris.nereids.trees.plans.commands.merge.MergeIntoCommand;
import org.apache.doris.nereids.trees.plans.commands.use.SwitchCommand;
import org.apache.doris.nereids.trees.plans.commands.use.UseCommand;
+import org.apache.doris.nereids.types.DataType;
+import org.apache.doris.nereids.types.NullType;
import org.apache.doris.qe.ConnectContext;
import org.apache.doris.qe.QueryState;
import org.apache.doris.qe.ResultSetMetaData;
@@ -97,6 +104,19 @@ final class FlightSqlQuerySchema {
}
static Schema analyze(ConnectContext context, String query) throws
Exception {
+ return analyze(context, query, null);
+ }
+
+ static Schema analyze(ConnectContext context, String query, List<Literal>
parameters) throws Exception {
+ return analyze(context, query, parameters, false).first;
+ }
+
+ static Pair<Schema, Schema> prepare(ConnectContext context, String query)
throws Exception {
+ return analyze(context, query, null, true);
+ }
+
+ private static Pair<Schema, Schema> analyze(ConnectContext context, String
query,
+ List<Literal> parameters, boolean preparing) throws Exception {
synchronized (context) {
ConnectContext previousThreadContext = ConnectContext.get();
StatementContext previousStatement = context.getStatementContext();
@@ -149,12 +169,33 @@ final class FlightSqlQuerySchema {
StatementContext statementContext =
statement.getStatementContext();
context.setStatementContext(statementContext);
statementContext.setParsedStatement(statement);
- if (!statementContext.getPlaceholders().isEmpty()) {
- throw CallStatus.UNIMPLEMENTED.withDescription(
- "Flight SQL parameter binding is not
supported").toRuntimeException();
+ int parameterCount = statementContext.getPlaceholders().size();
+ if (parameterCount > FlightSqlParameters.MAX_PARAMETERS) {
+ throw CallStatus.INVALID_ARGUMENT.withDescription("Too
many query parameters (maximum 1024)")
+ .toRuntimeException();
}
- List<Field> fields = new ArrayList<>();
Plan plan = statement.getLogicalPlan();
+ Map<PlaceholderId, DataType> parameterTypes = new HashMap<>();
+ if (parameterCount > 0) {
+ if (statements.size() != 1 || plan instanceof Command) {
+ throw CallStatus.INVALID_ARGUMENT.withDescription(
+ "Parameters require a single query
statement").toRuntimeException();
+ }
+ if (preparing) {
+ statementContext.setPrepareStage(true);
+
parameterTypes.putAll(FlightSqlParameters.inferCastTypes(plan));
+ parameters = new ArrayList<>();
+ for (Placeholder placeholder :
statementContext.getPlaceholders()) {
+ // Unknown values must not acquire the synthetic
STRING type used by MySQL Prepare.
+ parameters.add(new
NullLiteral(parameterTypes.getOrDefault(
+ placeholder.getPlaceholderId(),
NullType.INSTANCE)));
+ }
+ }
+ }
+ if (parameters != null || parameterCount > 0) {
+ FlightSqlParameters.bind(statementContext, parameters);
+ }
+ List<Field> fields = new ArrayList<>();
if (plan instanceof Command) {
resolveNamespace(context, plan, scopedDatabases);
ResultSetMetaData metadata = commandMetadata(context,
(Command) plan);
@@ -194,12 +235,26 @@ final class FlightSqlQuerySchema {
Plan analyzed = cascades.getRewritePlan();
// PrepareCommandPlanner stops before the rewrite phase
that normally checks privileges.
new CheckPrivileges().rewriteRoot(analyzed,
cascades.getCurrentJobContext());
+ if (preparing) {
+ statementContext.getIdToComparisonSlot().forEach((id,
slot) ->
+ parameterTypes.putIfAbsent(id,
slot.getDataType()));
+ }
for (Slot slot : analyzed.getOutput()) {
fields.add(field(slot.getName(),
slot.getDataType().toCatalogDataType(), slot.nullable(),
true,
context.getSessionVariable().getTimeZone()));
}
}
- return new Schema(fields);
+ List<Field> parameterFields = new ArrayList<>();
+ for (int i = 0; preparing && i < parameterCount; i++) {
+ DataType type = parameterTypes.getOrDefault(
+
statementContext.getPlaceholders().get(i).getPlaceholderId(),
NullType.INSTANCE);
+ Field parameterField = field(String.valueOf(i),
type.toCatalogDataType(), true,
+ true, context.getSessionVariable().getTimeZone());
+ // Result schemas support more Arrow types than parameter
uploads do.
+
parameterFields.add(FlightSqlParameters.supportsParameterType(parameterField.getType())
+ ? parameterField :
Field.nullable(String.valueOf(i), new ArrowType.Null()));
+ }
+ return Pair.of(new Schema(fields), new
Schema(parameterFields));
} finally {
try {
List<AutoCloseable> resources = new ArrayList<>();
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/sessions/FlightSqlConnectPoolMgr.java
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/sessions/FlightSqlConnectPoolMgr.java
index c8854507e00..63d01a87a8a 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/sessions/FlightSqlConnectPoolMgr.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/service/arrowflight/sessions/FlightSqlConnectPoolMgr.java
@@ -57,6 +57,9 @@ public class FlightSqlConnectPoolMgr extends ConnectPoolMgr {
@Override
public void unregisterConnection(ConnectContext ctx) {
+ // Use the short-lived prepared-state lock so teardown never waits for
a running query
+ // holding the connection monitor before it can cancel that query.
+ ctx.closePreparedQueries();
// All Flight SQL session teardown paths (idle/query timeout, bearer
token expiry, and
// explicit CloseSession) reach here. Release channel-cached Arrow
results before removing
// the context from the pool.
diff --git
a/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/DorisFlightSqlSchemaTest.java
b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/DorisFlightSqlSchemaTest.java
index 70fb507e502..86fea280901 100644
---
a/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/DorisFlightSqlSchemaTest.java
+++
b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/DorisFlightSqlSchemaTest.java
@@ -182,11 +182,20 @@ public class DorisFlightSqlSchemaTest extends
TestWithFeService {
@Test
void invalidSqlIsRejectedAndConnectionRecovers() throws Exception {
- for (String sql : Arrays.asList("SELEC 1", "SELECT missing_column FROM
schema_input")) {
- Assertions.assertThrows(FlightRuntimeException.class, () ->
schema(sql));
+ for (String sql : Arrays.asList("SELEC 1", "SELECT missing_column FROM
schema_input",
+ "SELECT CAST(? AS BIGINT) FROM missing_schema_table")) {
+ FlightRuntimeException schemaFailure = Assertions.assertThrows(
+ FlightRuntimeException.class, () -> schema(sql));
+ Assertions.assertEquals(FlightStatusCode.INVALID_ARGUMENT,
schemaFailure.status().code());
Mockito.clearInvocations(connectContext);
-
Assertions.assertThrows(java.util.concurrent.ExecutionException.class, () ->
prepare(sql));
+ java.util.concurrent.ExecutionException prepareFailure =
Assertions.assertThrows(
+ java.util.concurrent.ExecutionException.class, () ->
prepare(sql));
+ Assertions.assertTrue(prepareFailure.getCause() instanceof
FlightRuntimeException);
+ Assertions.assertEquals(FlightStatusCode.INVALID_ARGUMENT,
+ ((FlightRuntimeException)
prepareFailure.getCause()).status().code());
Mockito.verify(connectContext,
Mockito.never()).addPreparedQuery(Mockito.anyString(), Mockito.anyString(),
Mockito.any());
+ Mockito.verify(connectContext, Mockito.never()).addPreparedQuery(
+ Mockito.anyString(), Mockito.anyString(), Mockito.any(),
Mockito.anyInt());
Assertions.assertEquals("id", preparedSchema("SELECT id FROM
schema_input")
.getFields().get(0).getName());
}
diff --git
a/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlParameterBindingTest.java
b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlParameterBindingTest.java
new file mode 100644
index 00000000000..1f6c7cc2a7f
--- /dev/null
+++
b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlParameterBindingTest.java
@@ -0,0 +1,41 @@
+// 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.service.arrowflight.sessions.FlightSessionsManager;
+
+import org.apache.arrow.flight.FlightProducer.CallContext;
+import org.apache.arrow.flight.FlightProducer.StreamListener;
+import org.apache.arrow.flight.FlightStream;
+import org.apache.arrow.flight.Location;
+import
org.apache.arrow.flight.sql.impl.FlightSql.CommandPreparedStatementQuery;
+import org.junit.Assert;
+import org.junit.Test;
+import org.mockito.Mockito;
+
+public class FlightSqlParameterBindingTest {
+ @Test
+ public void acceptsParameterUpload() throws Exception {
+ try (DorisFlightSqlProducer producer = new DorisFlightSqlProducer(
+ Location.forGrpcInsecure("127.0.0.1", 0),
Mockito.mock(FlightSessionsManager.class))) {
+ Assert.assertNotNull(producer.acceptPutPreparedStatementQuery(
+ CommandPreparedStatementQuery.getDefaultInstance(),
Mockito.mock(CallContext.class),
+ Mockito.mock(FlightStream.class),
Mockito.mock(StreamListener.class)));
+ }
+ }
+}
diff --git
a/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlPreparedQueryTest.java
b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlPreparedQueryTest.java
new file mode 100644
index 00000000000..aa6d37582bb
--- /dev/null
+++
b/fe/fe-core/src/test/java/org/apache/doris/service/arrowflight/FlightSqlPreparedQueryTest.java
@@ -0,0 +1,666 @@
+// 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.analysis.StatementBase;
+import org.apache.doris.nereids.StatementContext;
+import org.apache.doris.nereids.glue.LogicalPlanAdapter;
+import org.apache.doris.nereids.parser.NereidsParser;
+import org.apache.doris.nereids.trees.expressions.literal.BigIntLiteral;
+import org.apache.doris.nereids.trees.expressions.literal.DateTimeLiteral;
+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.StringLiteral;
+import org.apache.doris.nereids.trees.expressions.literal.TimestampTzLiteral;
+import org.apache.doris.nereids.types.DateTimeV2Type;
+import org.apache.doris.nereids.types.StringType;
+import org.apache.doris.nereids.types.TimeStampTzType;
+import org.apache.doris.qe.ConnectContext;
+import org.apache.doris.qe.StmtExecutor;
+import org.apache.doris.service.arrowflight.sessions.FlightSessionsManager;
+import org.apache.doris.service.arrowflight.sessions.FlightSqlConnectContext;
+import org.apache.doris.utframe.TestWithFeService;
+
+import com.google.common.base.Strings;
+import com.google.protobuf.Any;
+import com.google.protobuf.ByteString;
+import org.apache.arrow.flight.Action;
+import org.apache.arrow.flight.AsyncPutListener;
+import org.apache.arrow.flight.CallHeaders;
+import org.apache.arrow.flight.FlightClient;
+import org.apache.arrow.flight.FlightDescriptor;
+import org.apache.arrow.flight.FlightProducer.CallContext;
+import org.apache.arrow.flight.FlightRuntimeException;
+import org.apache.arrow.flight.FlightServer;
+import org.apache.arrow.flight.FlightStatusCode;
+import org.apache.arrow.flight.FlightStream;
+import org.apache.arrow.flight.Location;
+import org.apache.arrow.flight.auth2.CallHeaderAuthenticator;
+import org.apache.arrow.flight.sql.FlightSqlClient;
+import
org.apache.arrow.flight.sql.impl.FlightSql.ActionClosePreparedStatementRequest;
+import
org.apache.arrow.flight.sql.impl.FlightSql.ActionCreatePreparedStatementRequest;
+import
org.apache.arrow.flight.sql.impl.FlightSql.ActionCreatePreparedStatementResult;
+import
org.apache.arrow.flight.sql.impl.FlightSql.CommandPreparedStatementQuery;
+import org.apache.arrow.memory.RootAllocator;
+import org.apache.arrow.vector.BigIntVector;
+import org.apache.arrow.vector.DateDayVector;
+import org.apache.arrow.vector.DecimalVector;
+import org.apache.arrow.vector.Float8Vector;
+import org.apache.arrow.vector.IntVector;
+import org.apache.arrow.vector.TimeStampMicroVector;
+import org.apache.arrow.vector.TimeStampVector;
+import org.apache.arrow.vector.VarBinaryVector;
+import org.apache.arrow.vector.VarCharVector;
+import org.apache.arrow.vector.VectorSchemaRoot;
+import org.apache.arrow.vector.types.pojo.ArrowType;
+import org.apache.arrow.vector.types.pojo.Field;
+import org.apache.arrow.vector.types.pojo.Schema;
+import org.junit.jupiter.api.Assertions;
+import org.junit.jupiter.api.Test;
+import org.mockito.MockedConstruction;
+import org.mockito.Mockito;
+
+import java.lang.reflect.InvocationTargetException;
+import java.lang.reflect.Method;
+import java.math.BigDecimal;
+import java.nio.charset.StandardCharsets;
+import java.sql.Connection;
+import java.sql.DriverManager;
+import java.sql.PreparedStatement;
+import java.sql.Types;
+import java.util.Arrays;
+import java.util.List;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.Executors;
+import java.util.concurrent.TimeUnit;
+
+public class FlightSqlPreparedQueryTest extends TestWithFeService {
+ private RootAllocator allocator;
+ private FlightServer server;
+ private FlightClient client;
+ private FlightSqlClient sqlClient;
+ private DorisFlightSqlProducer producer;
+ private FlightSqlConnectContext flightContext;
+
+ @Override
+ protected void runBeforeAll() throws Exception {
+ allocator = new RootAllocator();
+ flightContext = new FlightSqlConnectContext("parameter-peer");
+ flightContext.setEnv(connectContext.getEnv());
+
flightContext.setCurrentUserIdentity(connectContext.getCurrentUserIdentity());
+ flightContext.setSessionVariable(connectContext.getSessionVariable());
+ FlightSessionsManager sessions =
Mockito.mock(FlightSessionsManager.class);
+
Mockito.when(sessions.getConnectContext("parameter-peer")).thenReturn(flightContext);
+ producer = new
DorisFlightSqlProducer(Location.forGrpcInsecure("127.0.0.1", 0), sessions);
+ server = FlightServer.builder(allocator,
Location.forGrpcInsecure("127.0.0.1", 0), producer)
+ .headerAuthenticator(headers -> new
CallHeaderAuthenticator.AuthResult() {
+ @Override
+ public String getPeerIdentity() {
+ return "parameter-peer";
+ }
+
+ @Override
+ public void appendToOutgoingHeaders(CallHeaders headers) {
+ headers.insert("authorization", "Bearer
parameter-peer");
+ }
+ }).build().start();
+ client = FlightClient.builder(allocator, server.getLocation()).build();
+ sqlClient = new FlightSqlClient(client);
+ }
+
+ @Override
+ protected void runAfterAll() throws Exception {
+ sqlClient.close();
+ server.close();
+ java.lang.reflect.Field executor =
DorisFlightSqlProducer.class.getDeclaredField("executorService");
+ executor.setAccessible(true);
+ ((ExecutorService) executor.get(producer)).shutdownNow();
+ producer.close();
+ flightContext.getFlightSqlChannel().close();
+ allocator.close();
+ }
+
+ private CommandPreparedStatementQuery prepare(String sql) throws Exception
{
+ Action action = new Action("CreatePreparedStatement", Any.pack(
+
ActionCreatePreparedStatementRequest.newBuilder().setQuery(sql).build()).toByteArray());
+ ActionCreatePreparedStatementResult prepared =
Any.parseFrom(client.doAction(action).next().getBody())
+ .unpack(ActionCreatePreparedStatementResult.class);
+ Assertions.assertFalse(prepared.getDatasetSchema().isEmpty());
+ return CommandPreparedStatementQuery.newBuilder()
+
.setPreparedStatementHandle(prepared.getPreparedStatementHandle()).build();
+ }
+
+ private void bind(CommandPreparedStatementQuery command, VectorSchemaRoot
root) {
+ AsyncPutListener listener = new AsyncPutListener();
+ FlightClient.ClientStreamListener writer = client.startPut(
+ FlightDescriptor.command(Any.pack(command).toByteArray()),
root, listener);
+ writer.putNext();
+ writer.completed();
+ listener.getResult();
+ }
+
+ private String id(CommandPreparedStatementQuery command) {
+ return
command.getPreparedStatementHandle().toStringUtf8().substring("parameter-peer:".length());
+ }
+
+ private Schema schema(CommandPreparedStatementQuery command) {
+ return
client.getSchema(FlightDescriptor.command(Any.pack(command).toByteArray())).getSchema();
+ }
+
+ @Test
+ public void jdbcPrepareAdvertisesResultAndParameterTypes() throws
Exception {
+ Class.forName("org.apache.arrow.driver.jdbc.ArrowFlightJdbcDriver");
+ try (Connection connection =
DriverManager.getConnection("jdbc:arrow-flight-sql://127.0.0.1:"
+ + server.getPort() + "?useEncryption=false", "fixture",
"fixture");
+ PreparedStatement statement = connection.prepareStatement(
+ "SELECT CAST(? AS BIGINT) AS n, CAST(? AS STRING) AS
s, CAST(? AS DOUBLE) AS d")) {
+ Assertions.assertEquals(3,
statement.getMetaData().getColumnCount());
+ for (int i = 0; i < 3; i++) {
+ int type = new int[] {Types.BIGINT, Types.VARCHAR,
Types.DOUBLE}[i];
+ Assertions.assertEquals(type,
statement.getMetaData().getColumnType(i + 1));
+ Assertions.assertEquals(type,
statement.getParameterMetaData().getParameterType(i + 1));
+ }
+ }
+ }
+
+ @Test
+ public void jdbcPrepareInfersRangePredicateTypes() throws Exception {
+ Class.forName("org.apache.arrow.driver.jdbc.ArrowFlightJdbcDriver");
+ try (Connection connection =
DriverManager.getConnection("jdbc:arrow-flight-sql://127.0.0.1:"
+ + server.getPort() + "?useEncryption=false", "fixture",
"fixture");
+ PreparedStatement statement = connection.prepareStatement(
+ "SELECT number FROM numbers(\"number\"=\"10\") WHERE
number >= ? AND ? > number")) {
+ Assertions.assertEquals(Types.BIGINT,
statement.getParameterMetaData().getParameterType(1));
+ Assertions.assertEquals(Types.BIGINT,
statement.getParameterMetaData().getParameterType(2));
+ Assertions.assertEquals(Types.BIGINT,
statement.getMetaData().getColumnType(1));
+ }
+ }
+
+ @Test
+ public void jdbcPreparePreservesInnerCastsAndSubqueryTypes() throws
Exception {
+ Class.forName("org.apache.arrow.driver.jdbc.ArrowFlightJdbcDriver");
+ try (Connection connection =
DriverManager.getConnection("jdbc:arrow-flight-sql://127.0.0.1:"
+ + server.getPort() + "?useEncryption=false", "fixture",
"fixture");
+ PreparedStatement statement = connection.prepareStatement(
+ "SELECT CAST(CAST(? AS BIGINT) AS STRING), (SELECT
CAST(? AS DOUBLE))")) {
+ Assertions.assertEquals(Types.BIGINT,
statement.getParameterMetaData().getParameterType(1));
+ Assertions.assertEquals(Types.DOUBLE,
statement.getParameterMetaData().getParameterType(2));
+ Assertions.assertEquals(Types.VARCHAR,
statement.getMetaData().getColumnType(1));
+ Assertions.assertEquals(Types.DOUBLE,
statement.getMetaData().getColumnType(2));
+ }
+ }
+
+ @Test
+ public void unknownParameterTypesAreNotInvented() throws Exception {
+ try (FlightSqlClient.PreparedStatement statement =
sqlClient.prepare("SELECT ? AS value")) {
+ Assertions.assertEquals(new ArrowType.Null(),
+
statement.getParameterSchema().getFields().get(0).getType());
+ Assertions.assertEquals(new ArrowType.Null(),
+
statement.getResultSetSchema().getFields().get(0).getType());
+ }
+ }
+
+ @Test
+ public void invalidSqlReportsClientErrorOverFlight() throws Exception {
+ for (String query : Arrays.asList("SELEC 1", "SELECT
definitely_missing_schema_column",
+ "SELECT CAST(? AS BIGINT) FROM missing_parameter_table")) {
+ FlightRuntimeException prepareFailure =
Assertions.assertThrows(FlightRuntimeException.class,
+ () -> sqlClient.prepare(query));
+ Assertions.assertEquals(FlightStatusCode.INVALID_ARGUMENT,
prepareFailure.status().code());
+ FlightRuntimeException schemaFailure =
Assertions.assertThrows(FlightRuntimeException.class,
+ () -> sqlClient.getExecuteSchema(query));
+ Assertions.assertEquals(FlightStatusCode.INVALID_ARGUMENT,
schemaFailure.status().code());
+ try (FlightSqlClient.PreparedStatement statement =
sqlClient.prepare("SELECT 1 AS id")) {
+ Assertions.assertEquals("id",
statement.getResultSetSchema().getFields().get(0).getName());
+ }
+ }
+ }
+
+ @Test
+ public void unsupportedParameterVectorsDoNotRestrictResultTypes() throws
Exception {
+ try (FlightSqlClient.PreparedStatement statement =
sqlClient.prepare("SELECT CAST(? AS ARRAY<INT>)")) {
+ Assertions.assertEquals(new ArrowType.Null(),
+
statement.getParameterSchema().getFields().get(0).getType());
+ Assertions.assertEquals(new ArrowType.List(),
+
statement.getResultSetSchema().getFields().get(0).getType());
+ }
+ try (FlightSqlClient.PreparedStatement statement = sqlClient.prepare(
+ "SELECT CAST(CAST(? AS STRING) AS ARRAY<INT>)")) {
+ Assertions.assertEquals(new ArrowType.Utf8(),
+
statement.getParameterSchema().getFields().get(0).getType());
+ Assertions.assertEquals(new ArrowType.List(),
+
statement.getResultSetSchema().getFields().get(0).getType());
+ }
+ }
+
+ @Test
+ public void bindsAndRebindsIntegerQueryOverFlight() throws Exception {
+ CommandPreparedStatementQuery command = prepare("SELECT CAST(? AS
BIGINT) AS value");
+ Assertions.assertEquals(new ArrowType.Int(64, true),
schema(command).getFields().get(0).getType());
+ Assertions.assertThrows(FlightRuntimeException.class,
+ () ->
client.getInfo(FlightDescriptor.command(Any.pack(command).toByteArray())));
+ try (BigIntVector value = new BigIntVector("parameter", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(value)) {
+ for (long expected : new long[] {42, -7, Long.MAX_VALUE}) {
+ value.setSafe(0, expected);
+ root.setRowCount(1);
+ bind(command, root);
+ Assertions.assertEquals(new ArrowType.Int(64, true),
schema(command).getFields().get(0).getType());
+ Assertions.assertEquals(expected,
flightContext.getPreparedQueryParameters(id(command)).get(0).getValue());
+ }
+ } finally {
+ flightContext.removePreparedQuery(id(command));
+ }
+ }
+
+ @Test
+ public void bindsMixedTypesAndPreservesQuotedText() throws Exception {
+ String query = "SELECT CAST(? AS BIGINT) AS n, CAST(? AS STRING) AS
text, CAST(? AS DOUBLE) AS fraction";
+ CommandPreparedStatementQuery command = prepare(query);
+ try (BigIntVector n = new BigIntVector("n", allocator);
+ VarCharVector text = new VarCharVector("text", allocator);
+ Float8Vector fraction = new Float8Vector("fraction",
allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(n, text,
fraction)) {
+ String expected = "中文 ' ? \\ ; SELECT 2";
+ n.setSafe(0, 17);
+ text.setSafe(0, expected.getBytes(StandardCharsets.UTF_8));
+ fraction.setSafe(0, 2.5);
+ root.setRowCount(1);
+ bind(command, root);
+ List<Literal> parameters =
flightContext.getPreparedQueryParameters(id(command));
+ Assertions.assertEquals(17L, parameters.get(0).getValue());
+ Assertions.assertEquals(expected, parameters.get(1).getValue());
+ Assertions.assertEquals(2.5, parameters.get(2).getValue());
+ Assertions.assertEquals(3, schema(command).getFields().size());
+ // Exercise the execution parser hook without requiring a BE in
the FE unit test fixture.
+ try (FlightSqlConnectProcessor processor = new
FlightSqlConnectProcessor(flightContext)) {
+ java.lang.reflect.Field bindings =
FlightSqlConnectProcessor.class.getDeclaredField("parameters");
+ bindings.setAccessible(true);
+ bindings.set(processor, parameters);
+ List<StatementBase> statements =
processor.parseWithFallback(query, query,
+ flightContext.getSessionVariable());
+ StatementContext statement = ((LogicalPlanAdapter)
statements.get(0))
+ .getStatementContext();
+ try {
+ for (int i = 0; i < parameters.size(); i++) {
+ Assertions.assertEquals(parameters.get(i),
statement.getIdToPlaceholderRealExpr().get(
+
statement.getPlaceholders().get(i).getPlaceholderId()));
+ }
+ } finally {
+ statement.close();
+ }
+ }
+ } finally {
+ flightContext.removePreparedQuery(id(command));
+ }
+ }
+
+ @Test
+ public void failedRebindInvalidatesOldValuesAndConnectionRecovers() throws
Exception {
+ CommandPreparedStatementQuery command = prepare("SELECT ? AS value");
+ try (IntVector value = new IntVector("value", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(value)) {
+ value.setSafe(0, 1);
+ root.setRowCount(1);
+ bind(command, root);
+ value.setSafe(1, 2);
+ root.setRowCount(2);
+ FlightRuntimeException failure =
Assertions.assertThrows(FlightRuntimeException.class,
+ () -> bind(command, root));
+ Assertions.assertEquals(FlightStatusCode.UNIMPLEMENTED,
failure.status().code());
+
Assertions.assertNull(flightContext.getPreparedQueryParameters(id(command)));
+ Assertions.assertThrows(FlightRuntimeException.class,
+ () ->
client.getInfo(FlightDescriptor.command(Any.pack(command).toByteArray())));
+ root.setRowCount(1);
+ value.setSafe(0, 3);
+ bind(command, root);
+ Assertions.assertEquals(new ArrowType.Int(32, true),
schema(command).getFields().get(0).getType());
+ value.setNull(0);
+ bind(command, root);
+
Assertions.assertTrue(flightContext.getPreparedQueryParameters(id(command)).get(0)
instanceof NullLiteral);
+ } finally {
+ flightContext.removePreparedQuery(id(command));
+ }
+ Assertions.assertNotNull(sqlClient.getExecuteSchema("SELECT 1"));
+ }
+
+ @Test
+ public void convertsDetachedValuesAndTypedNulls() {
+ try (IntVector value = new IntVector("value", allocator);
+ VarCharVector text = new VarCharVector("text", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(value, text)) {
+ value.setSafe(0, 12);
+ text.setNull(0);
+ root.setRowCount(1);
+ List<Literal> literals = FlightSqlParameters.convert(root);
+ value.setSafe(0, 99);
+ Assertions.assertEquals(12, literals.get(0).getValue());
+ Assertions.assertTrue(literals.get(1) instanceof NullLiteral);
+ Assertions.assertEquals(StringType.INSTANCE,
literals.get(1).getDataType());
+ }
+ }
+
+ @Test
+ public void validatesWholeUploadInsteadOfUsingFirstBatch() {
+ try (IntVector value = new IntVector("value", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(value)) {
+ value.setSafe(0, 1);
+ root.setRowCount(1);
+ FlightStream stream = Mockito.mock(FlightStream.class);
+ Mockito.when(stream.next()).thenReturn(true, true, false);
+ Mockito.when(stream.getRoot()).thenReturn(root);
+ Assertions.assertThrows(Exception.class, () ->
FlightSqlParameters.read(stream, 1));
+ Mockito.verify(stream, Mockito.times(3)).next();
+ FlightStream empty = Mockito.mock(FlightStream.class);
+ Assertions.assertThrows(Exception.class, () ->
FlightSqlParameters.read(empty, 1));
+ Assertions.assertEquals(Arrays.asList(),
FlightSqlParameters.read(empty, 0));
+ }
+ }
+
+ @Test
+ public void limitsRetainedParametersAcrossHandlesAndReleasesTheirBudget() {
+ ConnectContext connection = new ConnectContext();
+ // Test sources also target Java 8, where String.repeat is unavailable.
+ List<Literal> parameters = Arrays.asList(
+ new StringLiteral(Strings.repeat("x", 512 * 1024)));
+ for (int i = 0; i < 15; i++) {
+ connection.addPreparedQuery("p" + i, "SELECT ?", null, 1);
+ connection.setPreparedQueryParameters("p" + i, parameters, null);
+ }
+ connection.addPreparedQuery("overflow", "SELECT ?", null, 1);
+ Assertions.assertThrows(IllegalArgumentException.class,
+ () -> connection.setPreparedQueryParameters("overflow",
parameters, null));
+ connection.removePreparedQuery("p0");
+ connection.setPreparedQueryParameters("overflow", parameters, null);
+ connection.beginPreparedQueryBinding("overflow");
+ connection.addPreparedQuery("replacement", "SELECT ?", null, 1);
+ connection.setPreparedQueryParameters("replacement", parameters, null);
+ long upload = connection.beginPreparedQueryBinding("replacement");
+ connection.clearPreparedQueries();
+
Assertions.assertFalse(connection.isPreparedQueryBindingCurrent("replacement",
upload));
+ Assertions.assertNull(connection.getPreparedQuery("replacement"));
+ connection.addPreparedQuery("fresh", "SELECT ?", null, 1);
+ connection.setPreparedQueryParameters("fresh", parameters, null);
+ }
+
+ @Test
+ public void rejectsInvalidValuesAndUnsupportedTypes() {
+ try (VarCharVector text = new VarCharVector("text", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(text)) {
+ text.setSafe(0, new byte[] {(byte) 0xc3, (byte) 0x28});
+ root.setRowCount(1);
+ Assertions.assertThrows(FlightRuntimeException.class, () ->
FlightSqlParameters.convert(root));
+ }
+ try (Float8Vector number = new Float8Vector("number", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(number)) {
+ number.setSafe(0, Double.POSITIVE_INFINITY);
+ root.setRowCount(1);
+ Assertions.assertThrows(FlightRuntimeException.class, () ->
FlightSqlParameters.convert(root));
+ }
+ try (VarBinaryVector binary =
+ new VarBinaryVector("binary", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(binary)) {
+ binary.setSafe(0, new byte[] {0, 1});
+ root.setRowCount(1);
+ FlightRuntimeException error =
Assertions.assertThrows(FlightRuntimeException.class,
+ () -> FlightSqlParameters.convert(root));
+ Assertions.assertEquals(FlightStatusCode.UNIMPLEMENTED,
error.status().code());
+ }
+ }
+
+ @Test
+ public void rejectsForwardingInsteadOfLosingTypedBindings() throws
Exception {
+ flightContext.setThreadLocalInfo();
+ LogicalPlanAdapter statement = (LogicalPlanAdapter) new NereidsParser()
+ .parseSQL("SELECT ?",
flightContext.getSessionVariable()).get(0);
+ try {
+ FlightSqlParameters.bind(statement.getStatementContext(),
Arrays.asList(
+ new BigIntLiteral(7)));
+ StmtExecutor executor = new StmtExecutor(flightContext, statement);
+ Method forward = StmtExecutor.class
+ .getDeclaredMethod("forwardToMaster");
+ forward.setAccessible(true);
+ InvocationTargetException failure = Assertions.assertThrows(
+ InvocationTargetException.class, () ->
forward.invoke(executor));
+
Assertions.assertTrue(failure.getCause().getMessage().contains("connect to
master FE"));
+ } finally {
+ statement.getStatementContext().close();
+ ConnectContext.remove();
+ }
+ }
+
+ @Test
+ public void preservesDecimalAndTemporalPrecision() {
+ try (DecimalVector decimal =
+ new DecimalVector("decimal", allocator, 20, 6);
+ DateDayVector date = new DateDayVector("date", allocator);
+ TimeStampMicroVector timestamp =
+ new TimeStampMicroVector("timestamp", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(decimal, date,
timestamp)) {
+ BigDecimal expected = new BigDecimal("12345678901234.567890");
+ decimal.setSafe(0, expected);
+ date.setSafe(0, -1);
+ timestamp.setSafe(0, -1);
+ root.setRowCount(1);
+ List<Literal> values = FlightSqlParameters.convert(root);
+ Assertions.assertEquals(expected, values.get(0).getValue());
+ Assertions.assertEquals("1969-12-31",
values.get(1).getStringValue());
+ Assertions.assertEquals("1969-12-31 23:59:59.999999",
values.get(2).getStringValue());
+ date.setSafe(0, Integer.MAX_VALUE);
+ Assertions.assertThrows(FlightRuntimeException.class, () ->
FlightSqlParameters.convert(root));
+ }
+ }
+
+ @Test
+ public void preservesTimestampSemanticsAndNulls() {
+ for (String zone : Arrays.asList(null, "", "UTC", "Asia/Shanghai",
"America/New_York")) {
+ for (org.apache.arrow.vector.types.TimeUnit unit :
org.apache.arrow.vector.types.TimeUnit.values()) {
+ ArrowType.Timestamp type = new ArrowType.Timestamp(unit, zone);
+ long nanosPerUnit =
java.util.concurrent.TimeUnit.valueOf(unit.name() + "S").toNanos(1);
+ try (VectorSchemaRoot root = VectorSchemaRoot.create(
+ new Schema(Arrays.asList(Field.nullable("value",
type))), allocator)) {
+ root.allocateNew();
+ TimeStampVector vector = (TimeStampVector)
root.getVector(0);
+ for (long value : new long[] {-1000, 0, 1000}) {
+ vector.setSafe(0, value);
+ root.setRowCount(1);
+ Literal literal =
FlightSqlParameters.convert(root).get(0);
+ Assertions.assertEquals(zone != null &&
!zone.isEmpty(), literal instanceof TimestampTzLiteral);
+
Assertions.assertEquals(java.time.LocalDateTime.of(1970, 1, 1, 0, 0)
+ .plusNanos(value * nanosPerUnit),
((DateTimeLiteral) literal).toJavaDateType());
+ }
+ vector.setNull(0);
+ Literal literal = FlightSqlParameters.convert(root).get(0);
+ Assertions.assertTrue(literal instanceof NullLiteral);
+ Assertions.assertEquals(zone != null && !zone.isEmpty()
+ ? TimeStampTzType.of(6) : DateTimeV2Type.of(6),
literal.getDataType());
+ }
+ }
+ }
+ }
+
+ @Test
+ public void rejectsTimezoneTimestampPrecisionLossAndOverflow() {
+ for (org.apache.arrow.vector.types.TimeUnit unit : Arrays.asList(
+ org.apache.arrow.vector.types.TimeUnit.NANOSECOND,
org.apache.arrow.vector.types.TimeUnit.SECOND)) {
+ try (VectorSchemaRoot root = VectorSchemaRoot.create(new
Schema(Arrays.asList(
+ Field.nullable("value", new ArrowType.Timestamp(unit,
"UTC")))), allocator)) {
+ root.allocateNew();
+ TimeStampVector vector = (TimeStampVector) root.getVector(0);
+ long[] values = unit ==
org.apache.arrow.vector.types.TimeUnit.NANOSECOND
+ ? new long[] {-1, 1, 1001} : new long[]
{Long.MIN_VALUE, Long.MAX_VALUE};
+ for (long value : values) {
+ vector.setSafe(0, value);
+ root.setRowCount(1);
+ FlightStream stream = Mockito.mock(FlightStream.class);
+ Mockito.when(stream.next()).thenReturn(true, false);
+ Mockito.when(stream.getRoot()).thenReturn(root);
+ FlightRuntimeException failure =
Assertions.assertThrows(FlightRuntimeException.class,
+ () -> FlightSqlParameters.read(stream, 1));
+ Assertions.assertEquals(FlightStatusCode.INVALID_ARGUMENT,
failure.status().code());
+ }
+ }
+ }
+ }
+
+ @Test
+ public void bindsTimezoneTimestampInNonUtcSession() throws Exception {
+ String previousZone = flightContext.getSessionVariable().getTimeZone();
+ try {
+ for (String zone : Arrays.asList("Asia/Shanghai",
"America/New_York")) {
+ flightContext.getSessionVariable().setTimeZone(zone);
+ try (FlightSqlClient.PreparedStatement statement =
sqlClient.prepare("SELECT CAST(? AS TIMESTAMPTZ(6))")) {
+ Assertions.assertEquals(new ArrowType.Timestamp(
+
org.apache.arrow.vector.types.TimeUnit.MICROSECOND, zone),
+
statement.getParameterSchema().getFields().get(0).getType());
+ }
+ CommandPreparedStatementQuery command = prepare("SELECT CAST(?
AS TIMESTAMPTZ(6))");
+ try (VectorSchemaRoot root = VectorSchemaRoot.create(new
Schema(Arrays.asList(Field.nullable("value",
+ new
ArrowType.Timestamp(org.apache.arrow.vector.types.TimeUnit.MICROSECOND,
"UTC")))), allocator)) {
+ root.allocateNew();
+ TimeStampVector vector = (TimeStampVector)
root.getVector(0);
+ for (long value : new long[] {0, -1, 1234567}) {
+ vector.setSafe(0, value);
+ root.setRowCount(1);
+ bind(command, root);
+ TimestampTzLiteral literal = (TimestampTzLiteral)
flightContext
+
.getPreparedQueryParameters(id(command)).get(0);
+
Assertions.assertEquals(java.time.LocalDateTime.of(1970, 1, 1, 0, 0)
+ .plusNanos(value * 1000),
literal.toJavaDateType());
+ Assertions.assertEquals(new ArrowType.Timestamp(
+
org.apache.arrow.vector.types.TimeUnit.MICROSECOND, zone),
+ schema(command).getFields().get(0).getType());
+ }
+ vector.setNull(0);
+ bind(command, root);
+ Assertions.assertEquals(TimeStampTzType.of(6),
flightContext
+
.getPreparedQueryParameters(id(command)).get(0).getDataType());
+ } finally {
+ client.doAction(new Action("ClosePreparedStatement",
Any.pack(
+ ActionClosePreparedStatementRequest.newBuilder()
+
.setPreparedStatementHandle(command.getPreparedStatementHandle()).build())
+ .toByteArray())).forEachRemaining(result -> { });
+ }
+ }
+ } finally {
+ flightContext.getSessionVariable().setTimeZone(previousZone);
+ }
+ }
+
+ @Test
+ public void rejectsForeignAndClosedHandles() throws Exception {
+ CommandPreparedStatementQuery command = prepare("SELECT ?");
+ CommandPreparedStatementQuery foreign =
CommandPreparedStatementQuery.newBuilder()
+
.setPreparedStatementHandle(ByteString.copyFromUtf8("another-peer:handle")).build();
+ FlightRuntimeException failure =
Assertions.assertThrows(FlightRuntimeException.class, () -> schema(foreign));
+ Assertions.assertEquals(FlightStatusCode.INVALID_ARGUMENT,
failure.status().code());
+ ActionClosePreparedStatementRequest close =
+ ActionClosePreparedStatementRequest.newBuilder()
+
.setPreparedStatementHandle(command.getPreparedStatementHandle()).build();
+ client.doAction(new Action("ClosePreparedStatement",
Any.pack(close).toByteArray())).forEachRemaining(r -> { });
+ FlightRuntimeException closed =
Assertions.assertThrows(FlightRuntimeException.class, () -> schema(command));
+ Assertions.assertEquals(FlightStatusCode.NOT_FOUND,
closed.status().code());
+ }
+
+ @Test
+ public void bindsRangePredicatesAndRejectsWrongParameterCount() throws
Exception {
+ CommandPreparedStatementQuery command = prepare("SELECT number FROM
numbers(\"number\"=\"10\") "
+ + "WHERE number >= ? AND number < ? ORDER BY number");
+ try (BigIntVector lower = new BigIntVector("lower", allocator);
+ BigIntVector upper = new BigIntVector("upper", allocator);
+ VectorSchemaRoot root = VectorSchemaRoot.of(lower, upper)) {
+ lower.setSafe(0, 2);
+ upper.setSafe(0, 5);
+ root.setRowCount(1);
+ bind(command, root);
+ Assertions.assertEquals(new ArrowType.Int(64, true),
schema(command).getFields().get(0).getType());
+ }
+ try (IntVector only = new IntVector("only", allocator);
+ VectorSchemaRoot wrong = VectorSchemaRoot.of(only)) {
+ only.setSafe(0, 1);
+ wrong.setRowCount(1);
+ FlightRuntimeException error =
Assertions.assertThrows(FlightRuntimeException.class,
+ () -> bind(command, wrong));
+ Assertions.assertEquals(FlightStatusCode.INVALID_ARGUMENT,
error.status().code());
+
Assertions.assertNull(flightContext.getPreparedQueryParameters(id(command)));
+ } finally {
+ flightContext.removePreparedQuery(id(command));
+ }
+ }
+
+ @Test
+ public void sessionCloseDoesNotWaitForExecutionMonitor() throws Exception {
+ ConnectContext connection = new ConnectContext();
+ connection.addPreparedQuery("p", "SELECT ?", null, 1);
+ long upload = connection.beginPreparedQueryBinding("p");
+ ExecutorService executor = Executors.newSingleThreadExecutor();
+ try {
+ synchronized (connection) {
+ executor.submit(connection::closePreparedQueries).get(5,
TimeUnit.SECONDS);
+ }
+
Assertions.assertFalse(connection.isPreparedQueryBindingCurrent("p", upload));
+ Assertions.assertThrows(IllegalStateException.class,
+ () -> connection.addPreparedQuery("new", "SELECT ?", null,
1));
+ } finally {
+ executor.shutdownNow();
+ }
+ }
+
+ @Test
+ public void closedHandleCannotCommitAnUploadOrLeaveUnreachableResults()
throws Exception {
+ ConnectContext connection = new ConnectContext();
+ connection.addPreparedQuery("p", "SELECT ?", null, 1);
+ long version = connection.beginPreparedQueryBinding("p");
+ connection.closePreparedQueries();
+ Assertions.assertFalse(connection.setPreparedQueryParameters("p",
Arrays.asList(new BigIntLiteral(1)),
+ null, version));
+ Assertions.assertEquals(-1,
connection.getPreparedQueryParameterCount("p"));
+ Assertions.assertNull(connection.getPreparedQueryParameters("p"));
+ Assertions.assertEquals(-1, connection.beginPreparedQueryBinding("p"));
+ String query = "SET query_timeout = 300";
+ flightContext.getState().reset();
+ flightContext.addPreparedQuery("close-race", query,
FlightSqlQuerySchema.analyze(flightContext, query));
+ CallContext call = Mockito.mock(CallContext.class);
+ Mockito.when(call.peerIdentity()).thenReturn("parameter-peer");
+ CommandPreparedStatementQuery command =
CommandPreparedStatementQuery.newBuilder()
+
.setPreparedStatementHandle(ByteString.copyFromUtf8("parameter-peer:close-race")).build();
+ StmtExecutor deferred = Mockito.mock(StmtExecutor.class);
+ boolean previousLocal = flightContext.isReturnResultFromLocal();
+ try (MockedConstruction<FlightSqlConnectProcessor> ignored =
Mockito.mockConstruction(
+ FlightSqlConnectProcessor.class, (processor, context) -> {
+ Mockito.doAnswer(invocation -> {
+ flightContext.clearPreparedQueries();
+ flightContext.setReturnResultFromLocal(true);
+ flightContext.addFlightSqlDeferredExecutor(deferred);
+ return null;
+ }).when(processor).handleQuery(Mockito.anyString());
+ })) {
+ FlightRuntimeException error =
Assertions.assertThrows(FlightRuntimeException.class,
+ () -> producer.getFlightInfoPreparedStatement(
+ command, call,
FlightDescriptor.command(Any.pack(command).toByteArray())));
+ Assertions.assertEquals(FlightStatusCode.NOT_FOUND,
error.status().code());
+ Assertions.assertEquals(0,
flightContext.getFlightSqlChannel().resultNum());
+ Mockito.verify(deferred).finalizeArrowFlightQuery();
+ } finally {
+ flightContext.setReturnResultFromLocal(previousLocal);
+ flightContext.closeFlightSqlDeferredExecutors();
+ flightContext.getFlightSqlChannel().reset();
+ }
+ }
+
+}
diff --git
a/regression-test/suites/arrow_flight_sql_p0/test_prepared_query_parameters.groovy
b/regression-test/suites/arrow_flight_sql_p0/test_prepared_query_parameters.groovy
new file mode 100644
index 00000000000..e41a3bbbbaa
--- /dev/null
+++
b/regression-test/suites/arrow_flight_sql_p0/test_prepared_query_parameters.groovy
@@ -0,0 +1,223 @@
+// 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 org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.CallOption
+import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.CallOptions
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.CloseSessionRequest
+import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.FlightClient
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.FlightRuntimeException
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.FlightStatusCode
+import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.Location
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.flight.sql.FlightSqlClient
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.memory.RootAllocator
+import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.vector.BigIntVector
+import org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.vector.Float8Vector
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.vector.TimeStampMicroTZVector
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.vector.VarCharVector
+import
org.apache.arrow.driver.jdbc.shaded.org.apache.arrow.vector.VectorSchemaRoot
+
+import java.sql.DriverManager
+import java.sql.Types
+
+suite("test_prepared_query_parameters", "arrow_flight_sql") {
+ def config = context.config.otherConfigs
+ def location =
Location.forGrpcInsecure(config.get("extArrowFlightSqlHost"),
+ Integer.parseInt(config.get("extArrowFlightSqlPort")))
+ Class.forName("org.apache.arrow.driver.jdbc.ArrowFlightJdbcDriver")
+ def jdbcUrl =
"jdbc:arrow-flight-sql://${config.get('extArrowFlightSqlHost')}:" +
+ "${config.get('extArrowFlightSqlPort')}?useEncryption=false"
+ DriverManager.getConnection(jdbcUrl, config.get("extArrowFlightSqlUser"),
+ config.get("extArrowFlightSqlPassword")).withCloseable {
connection ->
+ // Exercise JDBC's schema-driven binder and query classifier, not only
direct vector uploads.
+ connection.prepareStatement("SELECT CAST(? AS BIGINT) AS
value").withCloseable { statement ->
+ assertEquals(1, statement.getMetaData().getColumnCount())
+ assertEquals(Types.BIGINT,
statement.getParameterMetaData().getParameterType(1))
+ [42L, -7L, Long.MAX_VALUE].each { value ->
+ statement.setLong(1, value)
+ assertTrue(statement.execute())
+ statement.getResultSet().withCloseable { result ->
+ assertTrue(result.next())
+ assertEquals(value, result.getLong(1))
+ assertTrue(!result.next())
+ }
+ }
+ statement.setNull(1, Types.BIGINT)
+ statement.executeQuery().withCloseable { result ->
+ assertTrue(result.next())
+ assertEquals(null, result.getObject(1))
+ assertTrue(!result.next())
+ }
+ }
+ connection.prepareStatement("SELECT CAST(? AS BIGINT), CAST(? AS
STRING), CAST(? AS DOUBLE)")
+ .withCloseable { statement ->
+ statement.setLong(1, 17L)
+ statement.setString(2, "quoted ' ? 中文")
+ statement.setDouble(3, 2.5d)
+ statement.executeQuery().withCloseable { result ->
+ assertTrue(result.next())
+ assertEquals(17L, result.getLong(1))
+ assertEquals("quoted ' ? 中文", result.getString(2))
+ assertEquals(2.5d, result.getDouble(3))
+ assertTrue(!result.next())
+ }
+ }
+ connection.prepareStatement('SELECT number FROM numbers("number"="10")
'
+ + 'WHERE number >= ? AND ? > number ORDER BY
number').withCloseable { statement ->
+ [[2L, 5L], [6L, 9L]].each { bounds ->
+ statement.setLong(1, bounds[0])
+ statement.setLong(2, bounds[1])
+ statement.executeQuery().withCloseable { result ->
+ def rows = []
+ while (result.next()) {
+ rows.add(result.getLong(1))
+ }
+ assertEquals((bounds[0]..<bounds[1]).toList(), rows)
+ }
+ }
+ }
+ }
+ new RootAllocator().withCloseable { allocator ->
+ FlightClient.builder(allocator, location).build().withCloseable {
flight ->
+ def token =
flight.authenticateBasicToken(config.get("extArrowFlightSqlUser"),
+ config.get("extArrowFlightSqlPassword")).get()
+ CallOption[] options = [token, CallOptions.timeout(30,
java.util.concurrent.TimeUnit.SECONDS)]
+ def client = new FlightSqlClient(flight)
+ def fetch = { info ->
+ def rows = []
+ info.getEndpoints().each { endpoint ->
+ def resultLocation = endpoint.getLocations().isEmpty() ?
location : endpoint.getLocations()[0]
+ FlightClient.builder(allocator,
resultLocation).build().withCloseable { resultClient ->
+ resultClient.getStream(endpoint.getTicket(),
options).withCloseable { stream ->
+ while (stream.next()) {
+ def root = stream.getRoot()
+ for (int row = 0; row < root.getRowCount();
row++) {
+ rows.add(root.getFieldVectors().collect {
vector ->
+ def value = vector.getObject(row)
+ value == null ? null : value.toString()
+ })
+ }
+ }
+ }
+ }
+ }
+ rows
+ }
+ try {
+ def prepared = client.prepare("SELECT CAST(? AS BIGINT) AS
value", options)
+ try {
+ new BigIntVector("value", allocator).withCloseable { value
->
+ VectorSchemaRoot.of(value).withCloseable { root ->
+ prepared.setParameters(root)
+ [42L, -7L, Long.MAX_VALUE].each { expected ->
+ value.setSafe(0, expected)
+ root.setRowCount(1)
+ assertEquals([[expected.toString()]],
fetch(prepared.execute(options)))
+ }
+ value.setNull(0)
+ assertEquals([[null]],
fetch(prepared.execute(options)))
+ value.setSafe(0, 1L)
+ value.setSafe(1, 2L)
+ root.setRowCount(2)
+ try {
+ prepared.execute(options)
+ assertTrue(false, "Multiple parameter rows
must not silently execute only the first")
+ } catch (FlightRuntimeException e) {
+ assertEquals(FlightStatusCode.UNIMPLEMENTED,
e.status().code())
+ }
+ root.setRowCount(1)
+ value.setSafe(0, 3L)
+ assertEquals([["3"]],
fetch(prepared.execute(options)))
+ }
+ }
+ } finally {
+ // Close must use the same authenticated session that owns
the prepared handle.
+ prepared.close(options)
+ }
+ prepared = client.prepare("SELECT CAST(? AS BIGINT), CAST(? AS
STRING), CAST(? AS DOUBLE)", options)
+ try {
+ VectorSchemaRoot.of(new BigIntVector("n", allocator), new
VarCharVector("s", allocator),
+ new Float8Vector("f", allocator)).withCloseable {
root ->
+ String text = "quoted ' text ? with Unicode 中文"
+ root.getVector(0).setSafe(0, 17L)
+ root.getVector(1).setSafe(0, text.getBytes("UTF-8"))
+ root.getVector(2).setSafe(0, 2.5d)
+ root.setRowCount(1)
+ prepared.setParameters(root)
+ assertEquals([["17", text, "2.5"]],
fetch(prepared.execute(options)))
+ }
+ } finally {
+ prepared.close(options)
+ }
+ prepared = client.prepare('SELECT number FROM
numbers("number"="10") '
+ + 'WHERE number >= ? AND number < ? ORDER BY number',
options)
+ try {
+ VectorSchemaRoot.of(new BigIntVector("lower", allocator),
+ new BigIntVector("upper",
allocator)).withCloseable { root ->
+ root.setRowCount(1)
+ prepared.setParameters(root)
+ [[2L, 5L], [6L, 9L]].each { bounds ->
+ root.getVector(0).setSafe(0, bounds[0])
+ root.getVector(1).setSafe(0, bounds[1])
+ root.setRowCount(1)
+ assertEquals((bounds[0]..<bounds[1]).collect {
[it.toString()] },
+ fetch(prepared.execute(options)))
+ }
+ }
+ } finally {
+ prepared.close(options)
+ }
+ def originalZone = fetch(client.execute("SHOW VARIABLES LIKE
'time_zone'", options))[0][1]
+ try {
+ ["UTC", "Asia/Shanghai", "America/New_York"].each {
sessionZone ->
+ fetch(client.execute("SET
time_zone='${sessionZone}'".toString(), options))
+ prepared = client.prepare('SELECT number FROM
numbers("number"="3") '
+ + 'WHERE CAST(? AS TIMESTAMPTZ(6)) = CAST(? AS
TIMESTAMPTZ(6)) ORDER BY number', options)
+ try {
+ assertEquals(sessionZone,
+
prepared.getParameterSchema().getFields()[0].getType().getTimezone())
+ ["UTC", "America/New_York"].each { inputZone ->
+ VectorSchemaRoot.of(new
TimeStampMicroTZVector("instant", allocator, inputZone),
+ new VarCharVector("expected",
allocator)).withCloseable { root ->
+ prepared.setParameters(root)
+ // The timezone annotation must not shift
the epoch or change predicate matches.
+ [[0L, "1970-01-01 00:00:00+00:00"],
+ [-1L, "1969-12-31
23:59:59.999999+00:00"],
+ [1234567L, "1970-01-01
00:00:01.234567+00:00"],
+ [1710054000000000L, "2024-03-10
07:00:00+00:00"]].each { value ->
+ root.getVector(0).setSafe(0, value[0])
+ root.getVector(1).setSafe(0,
value[1].getBytes("UTF-8"))
+ root.setRowCount(1)
+ assertEquals([["0"], ["1"], ["2"]],
fetch(prepared.execute(options)))
+ }
+ root.getVector(0).setNull(0)
+ assertEquals([],
fetch(prepared.execute(options)))
+ }
+ }
+ } finally {
+ prepared.close(options)
+ }
+ }
+ } finally {
+ fetch(client.execute("SET
time_zone='${originalZone}'".toString(), options))
+ }
+ assertEquals([["1"]], fetch(client.execute("SELECT 1",
options)))
+ } finally {
+ flight.closeSession(new CloseSessionRequest(), options)
+ }
+ }
+ }
+}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]