From 812a7a1a94e937c17006aa26a9a6b0d3fdd9e04d Mon Sep 17 00:00:00 2001 From: morrySnow Date: Sat, 12 Sep 2026 02:02:12 +0800 Subject: [PATCH 1/6] [fix](nereids) Enforce fixed key predicates in prepared point queries ### What problem does this PR solve? Issue Number: None Related PR: None Problem Summary: Server-side prepared point queries could reuse a cached scan plan after a restrictive row policy added a fixed equality on the same key column as a placeholder. The direct path rediscovered key values from mutable translated conjuncts and rewrote every predicate sharing that column name, so a later binding could overwrite the policy literal and contaminate subsequent executions. Freeze placeholder bindings and exact fixed literals into an immutable key template, derive a typed tuple per execution, return a schema-correct empty batch before pruning or backend RPC for NULL or conflicting bindings, and fall back to normal planning for inexact coercions or unprovable predicates. ### Release note Prepared point queries now preserve fixed row-policy key constraints across repeated server-side executions and safely fall back when a predicate cannot be represented as an exact physical lookup key. ### Check List (For Author) - Test: - Unit Test: ExecuteCommandTest, ShortCircuitQueryContextTest, and PointQueryExecutorTest (21 tests). - Regression test: prepared_point_query_row_policy. - Build/checkstyle: DISABLE_BUILD_UI=ON ./build.sh --fe. - Behavior changed: Yes. Prepared point queries now return an empty result before tablet lookup when a bound key conflicts with an exact fixed predicate or is NULL; unprovable predicates use normal planning. - Does this need documentation: No. --- .../doris/nereids/StatementContext.java | 62 +++- .../rules/analysis/ExpressionAnalyzer.java | 40 ++- ...calResultSinkToShortCircuitPointQuery.java | 3 + .../trees/plans/commands/ExecuteCommand.java | 6 +- .../apache/doris/planner/OlapScanNode.java | 12 + .../org/apache/doris/planner/ScanNode.java | 23 ++ .../apache/doris/qe/PointQueryExecutor.java | 89 ++---- .../doris/qe/ShortCircuitQueryContext.java | 288 ++++++++++++++++++ .../org/apache/doris/qe/StmtExecutor.java | 31 +- .../plans/commands/ExecuteCommandTest.java | 44 ++- .../doris/qe/PointQueryExecutorTest.java | 19 ++ .../qe/ShortCircuitQueryContextTest.java | 185 +++++++++++ .../prepared_point_query_row_policy.groovy | 155 ++++++++++ 13 files changed, 879 insertions(+), 78 deletions(-) create mode 100644 regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java index 64670a7398b41a..fba49791c36056 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java @@ -55,6 +55,7 @@ import org.apache.doris.nereids.trees.expressions.Slot; import org.apache.doris.nereids.trees.expressions.SlotReference; import org.apache.doris.nereids.trees.expressions.StatementScopeIdGenerator; +import org.apache.doris.nereids.trees.expressions.literal.Literal; import org.apache.doris.nereids.trees.plans.ObjectId; import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.nereids.trees.plans.Plan; @@ -189,10 +190,16 @@ public enum TableFrom { private final IdGenerator placeHolderIdGenerator = PlaceholderId.createGenerator(); // relation id to placeholders for prepared statement, ordered by placeholder id private final Map idToPlaceholderRealExpr = new TreeMap<>(); - // map placeholder id to comparison slot, which will used to replace conjuncts - // directly + // Map placeholder id to the physical key slot used by the immutable point-query template. private final Map idToComparisonSlot = new TreeMap<>(); + // Equality literals that were written as constants in the statement or injected by a + // security policy. They are deliberately separate from placeholder bindings: a prepared + // point query must never replace a fixed predicate merely because it references the same + // column as a placeholder. + private final List pointQueryFixedKeyConstraints = new ArrayList<>(); + private boolean pointQueryFixedKeyConstraintsComplete = true; + // collect all hash join conditions to compute node connectivity in join graph private final List joinFilters = new ArrayList<>(); @@ -282,6 +289,9 @@ public enum TableFrom { private ShortCircuitQueryContext shortCircuitQueryContext; + // Built afresh for one EXECUTE. Never copied into the next StatementContext. + private ShortCircuitQueryContext.PointQueryExecutionContext pointQueryExecutionContext; + private FormatOptions formatOptions = FormatOptions.getDefault(); private Set plannerHooks = new HashSet<>(); @@ -435,8 +445,8 @@ public StatementContext createNextExecuteContext() { next.cteIdGenerator.resetId(cteIdGenerator.getCurrentId()); next.talbeIdGenerator.resetId(talbeIdGenerator.getCurrentId()); next.placeHolderIdGenerator.resetId(placeHolderIdGenerator.getCurrentId()); - // Placeholder bindings of this EXECUTE, and the comparison-slot registry used to replace - // conjuncts on the cached short-circuit plan without re-planning. + // Copy this EXECUTE's placeholder values and the stable placeholder-to-key registry. + // Fixed constraints and bound key tuples remain local to the context that owns them. next.idToPlaceholderRealExpr.putAll(idToPlaceholderRealExpr); next.idToComparisonSlot.putAll(idToComparisonSlot); next.placeholders = new ArrayList<>(placeholders); @@ -688,6 +698,15 @@ public void setShortCircuitQueryContext(ShortCircuitQueryContext shortCircuitQue this.shortCircuitQueryContext = shortCircuitQueryContext; } + public ShortCircuitQueryContext.PointQueryExecutionContext getPointQueryExecutionContext() { + return pointQueryExecutionContext; + } + + public void setPointQueryExecutionContext( + ShortCircuitQueryContext.PointQueryExecutionContext pointQueryExecutionContext) { + this.pointQueryExecutionContext = pointQueryExecutionContext; + } + public Optional getSqlCacheContext() { return Optional.ofNullable(sqlCacheContext); } @@ -837,6 +856,22 @@ public Map getIdToComparisonSlot() { return idToComparisonSlot; } + public void addPointQueryFixedKeyConstraint(SlotReference slot, Literal literal) { + pointQueryFixedKeyConstraints.add(new PointQueryFixedKeyConstraint(slot, literal)); + } + + public List getPointQueryFixedKeyConstraints() { + return pointQueryFixedKeyConstraints; + } + + public void markPointQueryFixedKeyConstraintsIncomplete() { + pointQueryFixedKeyConstraintsComplete = false; + } + + public boolean arePointQueryFixedKeyConstraintsComplete() { + return pointQueryFixedKeyConstraintsComplete; + } + public Map, Group>>> getCteIdToConsumerGroup() { return cteIdToConsumerGroup; } @@ -1665,4 +1700,23 @@ public boolean isDelete() { public void setIsDelete(boolean del) { isDelete = del; } + + /** A fixed equality operand and the exact bound slot it constrains. */ + public static class PointQueryFixedKeyConstraint { + private final SlotReference slot; + private final Literal literal; + + public PointQueryFixedKeyConstraint(SlotReference slot, Literal literal) { + this.slot = Objects.requireNonNull(slot); + this.literal = Objects.requireNonNull(literal); + } + + public SlotReference getSlot() { + return slot; + } + + public Literal getLiteral() { + return literal; + } + } } 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 f8fd0910ae12db..f8c9950eefa81d 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 @@ -92,6 +92,7 @@ import org.apache.doris.nereids.trees.expressions.typecoercion.ImplicitCastInputTypes; import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.nereids.trees.plans.Plan; +import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalJoin; import org.apache.doris.nereids.trees.plans.logical.LogicalPlan; import org.apache.doris.nereids.types.ArrayType; @@ -920,13 +921,11 @@ public Expression visitPlaceholder(Placeholder placeholder, ExpressionRewriteCon return visit(realExpr, context); } - // Register prepared statement placeholder id to related slot in comparison predicate. - // Used to replace expression in ShortCircuit plan + // Register each prepared-statement placeholder with its point-query key slot. private void registerPlaceholderIdToSlot(ComparisonPredicate cp, ExpressionRewriteContext context, Expression left, Expression right) { if (ConnectContext.get() != null && ConnectContext.get().getCommand() == MysqlCommand.COM_STMT_EXECUTE) { - // Used to replace expression in ShortCircuit plan if (cp.right() instanceof Placeholder && left instanceof SlotReference) { PlaceholderId id = ((Placeholder) cp.right()).getPlaceholderId(); context.cascadesContext.getStatementContext().getIdToComparisonSlot().put(id, (SlotReference) left); @@ -941,12 +940,43 @@ private void registerPlaceholderIdToSlot(ComparisonPredicate cp, public Expression visitComparisonPredicate(ComparisonPredicate cp, ExpressionRewriteContext context) { Expression left = cp.left().accept(this, context); Expression right = cp.right().accept(this, context); - // Used to replace expression in ShortCircuit plan registerPlaceholderIdToSlot(cp, context, left, right); + ComparisonPredicate original = cp; cp = (ComparisonPredicate) cp.withChildren(left, right); - return isEqualityBetweenJoinChildren(cp) + Expression analyzed = isEqualityBetweenJoinChildren(cp) ? TypeCoercionUtils.processJoinComparisonPredicate(cp) : TypeCoercionUtils.processComparisonPredicate(cp); + registerPointQueryFixedKeyConstraint(original, analyzed, context); + return analyzed; + } + + /** + * Keep fixed equality values distinct from prepared-statement placeholders. The point-query + * executor used to rediscover both from the translated scan conjuncts and then update every + * predicate sharing a column name. Once a row policy adds {@code key = constant}, that loses + * provenance and turns the policy constant into caller-controlled state. + */ + private void registerPointQueryFixedKeyConstraint(ComparisonPredicate original, + Expression analyzed, ExpressionRewriteContext context) { + if (!(currentPlan instanceof LogicalFilter) + || !(original instanceof EqualTo) + || original.left() instanceof Placeholder + || original.right() instanceof Placeholder + || !(analyzed instanceof EqualTo)) { + return; + } + Expression left = analyzed.child(0); + Expression right = analyzed.child(1); + if (left instanceof SlotReference && right instanceof Literal) { + context.cascadesContext.getStatementContext().addPointQueryFixedKeyConstraint( + (SlotReference) left, (Literal) right); + } else { + // The logical short-circuit shape checker currently peels one Cast from the key. + // A cast can be lossy (for example CAST(INT AS CHAR(1))), so it is not evidence for + // an exact physical lookup key. Keep the normal plan unless the fixed predicate is + // literally Slot = Literal. + context.cascadesContext.getStatementContext().markPointQueryFixedKeyConstraintsIncomplete(); + } } private boolean isEqualityBetweenJoinChildren(ComparisonPredicate comparisonPredicate) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java index 51bdc44b66b9c6..3a651ad2001b53 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java @@ -103,6 +103,9 @@ boolean scanMatchShortCircuitCondition(LogicalOlapScan olapScan) { // set short circuit flag and return the original plan private Plan shortCircuit(Plan root, OlapTable olapTable, Set conjuncts, StatementContext statementContext) { + if (!statementContext.arePointQueryFixedKeyConstraintsComplete()) { + return root; + } // All key columns in conjuncts Set colNames = Sets.newHashSet(); for (Expression expr : conjuncts) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java index ee09d2a7d1aa64..fedeab96956e59 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java @@ -162,8 +162,10 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { // statementContext.getShortCircuitQueryContext(), and the fallback (building one from a // null planner, since this path skips planning) would NPE. statementContext.setShortCircuitQueryContext(preparedStmtCtx.shortCircuitQueryContext.get()); - PointQueryExecutor.directExecuteShortCircuitQuery(executor, preparedStmtCtx, statementContext); - return; + if (PointQueryExecutor.directExecuteShortCircuitQuery( + executor, preparedStmtCtx, statementContext)) { + return; + } } if (ctx.getSessionVariable().enableGroupCommitFullPrepare) { if (preparedStmtCtx.groupCommitPlanner.isPresent()) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java index fcbd915100261e..f9dc65012605c8 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java @@ -1136,6 +1136,18 @@ public List lazyEvaluateRangeLocations() throws UserExcepti selectedIndexId = olapTable.getBaseIndexId(); // Only key columns computeColumnsFilter(olapTable.getBaseSchemaKeyColumns(), olapTable.getPartitionInfo()); + return evaluatePointQueryRangeLocations(); + } + + // Prepared point queries use execution-owned values so the cached conjunct template stays immutable. + public List lazyEvaluateRangeLocations( + Map keyValues) throws UserException { + selectedIndexId = olapTable.getBaseIndexId(); + computePointQueryColumnFilters(keyValues); + return evaluatePointQueryRangeLocations(); + } + + private List evaluatePointQueryRangeLocations() throws UserException { computePartitionInfo(); scanBackendIds.clear(); selectionHint = null; diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java index 98d9056e1af08c..3be660eb17dc8a 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java @@ -73,6 +73,7 @@ import org.apache.logging.log4j.Logger; import java.util.ArrayList; +import java.util.Collections; import java.util.HashSet; import java.util.LinkedHashSet; import java.util.List; @@ -201,6 +202,28 @@ public void computeColumnsFilter(List columns, PartitionInfo partitionsI } } + /** + * Build point-query pruning state from one execution's immutable key tuple. This avoids + * changing cached scan conjuncts (which are shared by every EXECUTE of a prepared handle) + * while still making partition and distribution pruning use the current parameter values. + */ + protected void computePointQueryColumnFilters(Map keyValues) { + columnFilters.clear(); + columnNameToRange.clear(); + for (Map.Entry entry : keyValues.entrySet()) { + LiteralExpr literal = entry.getValue(); + PartitionColumnFilter partitionFilter = new PartitionColumnFilter(); + partitionFilter.setLowerBound(literal, true); + partitionFilter.setUpperBound(literal, true); + columnFilters.put(entry.getKey(), partitionFilter); + + ColumnBound bound = ColumnBound.of(literal); + ColumnRange columnRange = ColumnRange.create(); + columnRange.intersect(Collections.singletonList(Range.closed(bound, bound))); + columnNameToRange.put(entry.getKey(), columnRange); + } + } + public void computeColumnsFilter() { // for load scan node, table is null // partitionsInfo maybe null for other scan node, eg: ExternalScanNode... diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java b/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java index e60ae03020742d..7042a776e33e6c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java @@ -17,14 +17,10 @@ package org.apache.doris.qe; -import org.apache.doris.analysis.BinaryPredicate; import org.apache.doris.analysis.Expr; -import org.apache.doris.analysis.ExprToSqlVisitor; import org.apache.doris.analysis.ExprToThriftVisitor; import org.apache.doris.analysis.LiteralExpr; import org.apache.doris.analysis.LiteralExprUtils; -import org.apache.doris.analysis.SlotRef; -import org.apache.doris.analysis.ToSqlParams; import org.apache.doris.catalog.Column; import org.apache.doris.catalog.Env; import org.apache.doris.catalog.OlapTable; @@ -35,13 +31,11 @@ import org.apache.doris.common.UserException; import org.apache.doris.mysql.MysqlCommand; import org.apache.doris.nereids.StatementContext; -import org.apache.doris.nereids.exceptions.AnalysisException; -import org.apache.doris.nereids.trees.expressions.SlotReference; -import org.apache.doris.nereids.trees.expressions.literal.Literal; -import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.planner.OlapScanNode; import org.apache.doris.proto.InternalService; import org.apache.doris.proto.InternalService.KeyTuple; +import org.apache.doris.qe.ShortCircuitQueryContext.PointQueryExecutionContext; +import org.apache.doris.qe.ShortCircuitQueryContext.PointQueryExecutionContext.Decision; import org.apache.doris.rpc.BackendServiceProxy; import org.apache.doris.rpc.RpcException; import org.apache.doris.rpc.TCustomProtocolFactory; @@ -56,7 +50,6 @@ import com.google.common.base.Preconditions; import com.google.common.base.Strings; import com.google.common.collect.Lists; -import com.google.common.collect.Maps; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; import org.apache.thrift.TDeserializer; @@ -68,8 +61,6 @@ import java.util.HashSet; import java.util.Iterator; import java.util.List; -import java.util.Map; -import java.util.Map.Entry; import java.util.Set; import java.util.concurrent.ExecutionException; import java.util.concurrent.Future; @@ -89,10 +80,13 @@ public class PointQueryExecutor implements CoordInterface { private List snapshotVisibleVersions; private final ShortCircuitQueryContext shortCircuitQueryContext; + private final PointQueryExecutionContext executionContext; - public PointQueryExecutor(ShortCircuitQueryContext ctx, int maxMessageSize) { + public PointQueryExecutor(ShortCircuitQueryContext ctx, + PointQueryExecutionContext executionContext, int maxMessageSize) { ctx.sanitize(); this.shortCircuitQueryContext = ctx; + this.executionContext = executionContext; this.maxMsgSizeOfResultReceiver = maxMessageSize; } @@ -116,7 +110,7 @@ private void updateCloudPartitionVersions() throws RpcException { void setScanRangeLocations() throws Exception { OlapScanNode scanNode = shortCircuitQueryContext.scanNode; // compute scan range - List locations = scanNode.lazyEvaluateRangeLocations(); + List locations = scanNode.lazyEvaluateRangeLocations(executionContext.getKeyValues()); Preconditions.checkNotNull(locations); if (scanNode.getScanTabletIds().isEmpty()) { return; @@ -151,52 +145,28 @@ static boolean shouldShuffleCandidateBackends(OlapScanNode scanNode) { } // execute query without analyze & plan - public static void directExecuteShortCircuitQuery(StmtExecutor executor, + public static boolean directExecuteShortCircuitQuery(StmtExecutor executor, PreparedStatementContext preparedStmtCtx, StatementContext statementContext) throws Exception { Preconditions.checkNotNull(preparedStmtCtx.shortCircuitQueryContext); ShortCircuitQueryContext shortCircuitQueryContext = preparedStmtCtx.shortCircuitQueryContext.get(); - // update conjuncts - Map colNameToConjunct = Maps.newHashMap(); - for (Entry entry : statementContext.getIdToComparisonSlot().entrySet()) { - String colName = entry.getValue().getOriginalColumn().get().getName(); - Expr conjunctVal = ((Literal) statementContext.getIdToPlaceholderRealExpr() - .get(entry.getKey())).toLegacyLiteral(); - colNameToConjunct.put(colName, conjunctVal); + PointQueryExecutionContext executionContext = + shortCircuitQueryContext.createPointQueryExecutionContext(statementContext); + if (executionContext.getDecision() == Decision.FALLBACK) { + // The copied prepared StatementContext still carries the previous execution's + // short-circuit flag. Clear all fast-path state before normal planning; otherwise + // planner construction can treat an unplanned statement as a point-query plan and + // try to build a ShortCircuitQueryContext from an empty scan-node list. + statementContext.setShortCircuitQuery(false); + statementContext.setShortCircuitQueryContext(null); + statementContext.setPointQueryExecutionContext(null); + return false; } - if (colNameToConjunct.size() != preparedStmtCtx.command.placeholderCount()) { - throw new AnalysisException("Mismatched conjuncts values size with prepared" - + "statement parameters size, expected " - + preparedStmtCtx.command.placeholderCount() - + ", but meet " + colNameToConjunct.size()); - } - updateScanNodeConjuncts(shortCircuitQueryContext.scanNode, colNameToConjunct); + statementContext.setPointQueryExecutionContext(executionContext); // short circuit plan and execution executor.executeAndSendResult(false, false, shortCircuitQueryContext.analzyedQuery, executor.getContext().getResultSender(), null, null); - } - - private static void updateScanNodeConjuncts(OlapScanNode scanNode, - Map colNameToConjunct) { - for (Expr conjunct : scanNode.getConjuncts()) { - BinaryPredicate binaryPredicate = (BinaryPredicate) conjunct; - SlotRef slot = null; - int updateChildIdx = 0; - if (binaryPredicate.getChild(0) instanceof LiteralExpr) { - slot = (SlotRef) binaryPredicate.getChildWithoutCast(1); - } else if (binaryPredicate.getChild(1) instanceof LiteralExpr) { - slot = (SlotRef) binaryPredicate.getChildWithoutCast(0); - updateChildIdx = 1; - } else { - Preconditions.checkState(false, "Should contains literal in " - + binaryPredicate.accept(ExprToSqlVisitor.INSTANCE, ToSqlParams.WITH_TABLE)); - } - // not a placeholder to replace - if (!colNameToConjunct.containsKey(slot.getColumnName())) { - continue; - } - binaryPredicate.setChild(updateChildIdx, colNameToConjunct.get(slot.getColumnName())); - } + return true; } public void setTimeout(long timeoutMs) { @@ -206,21 +176,13 @@ public void setTimeout(long timeoutMs) { void addKeyTuples( InternalService.PTabletKeyLookupRequest.Builder requestBuilder) throws TException { // TODO handle IN predicates - Map columnExpr = Maps.newHashMap(); KeyTuple.Builder kBuilder = KeyTuple.newBuilder(); - for (Expr expr : shortCircuitQueryContext.scanNode.getConjuncts()) { - BinaryPredicate predicate = (BinaryPredicate) expr; - Expr left = predicate.getChild(0); - Expr right = predicate.getChild(1); - SlotRef columnSlot = left.unwrapSlotRef(); - columnExpr.put(columnSlot.getColumnName(), right); - } // Serialize each literal expr as TExprNode bytes for typed value transfer. // BE deserializes the TExprNode and uses DataType::get_field() to extract // typed Field values directly, avoiding string parsing. TSerializer serializer = new TSerializer(); for (Column column : shortCircuitQueryContext.scanNode.getOlapTable().getBaseSchemaKeyColumns()) { - Expr literalExpr = columnExpr.get(column.getName()); + Expr literalExpr = executionContext.getKeyValues().get(column.getName()); // Ensure the literal type matches the column type for proper TExprNode // deserialization on BE side. Prepared statement parameters may have // mismatched types (e.g., setBigDecimal for INT column produces a @@ -260,6 +222,13 @@ public void cancel(Status cancelReason) { @Override public RowBatch getNext() throws Exception { + // A NULL placeholder or a fixed security constraint that disagrees with the bound + // key makes the full WHERE predicate false/unknown. Return before tablet pruning, + // cloud version lookup, or any BE RPC. + if (executionContext.getDecision() == Decision.EMPTY) { + return new RowBatch(); + } + Preconditions.checkState(executionContext.getDecision() == Decision.LOOKUP); setScanRangeLocations(); // No partition/tablet found return emtpy row batch if (candidateBackends == null || candidateBackends.isEmpty()) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java index 99496be25c9b78..88b6435bc378e5 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java @@ -20,9 +20,25 @@ import org.apache.doris.analysis.DescriptorToThriftConverter; import org.apache.doris.analysis.Expr; import org.apache.doris.analysis.ExprToThriftVisitor; +import org.apache.doris.analysis.LiteralExpr; +import org.apache.doris.analysis.LiteralExprUtils; import org.apache.doris.analysis.Queriable; +import org.apache.doris.catalog.Column; import org.apache.doris.catalog.OlapTable; import org.apache.doris.catalog.Type; +import org.apache.doris.nereids.NereidsPlanner; +import org.apache.doris.nereids.StatementContext; +import org.apache.doris.nereids.StatementContext.PointQueryFixedKeyConstraint; +import org.apache.doris.nereids.rules.expression.rules.FoldConstantRuleOnFE; +import org.apache.doris.nereids.trees.expressions.EqualTo; +import org.apache.doris.nereids.trees.expressions.Expression; +import org.apache.doris.nereids.trees.expressions.Placeholder; +import org.apache.doris.nereids.trees.expressions.SlotReference; +import org.apache.doris.nereids.trees.expressions.literal.BooleanLiteral; +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.util.TypeCoercionUtils; import org.apache.doris.planner.OlapScanNode; import org.apache.doris.planner.Planner; import org.apache.doris.thrift.TExpr; @@ -37,9 +53,12 @@ import org.apache.thrift.TSerializer; import java.util.ArrayList; +import java.util.Collections; +import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Objects; +import java.util.TreeMap; import java.util.UUID; import java.util.stream.Collectors; @@ -64,6 +83,7 @@ public class ShortCircuitQueryContext { public final OlapScanNode scanNode; public final Queriable analzyedQuery; + private final PointQueryKeyTemplate pointQueryKeyTemplate; // Serialized mysql Field, this could avoid serialize mysql field each time sendFields. // Since, serialize fields is too heavy when table is wide Map serializedFields = Maps.newHashMap(); @@ -87,6 +107,12 @@ List getReturnTypes() { } public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery) throws TException { + this(planner, analzyedQuery, + planner instanceof NereidsPlanner ? ((NereidsPlanner) planner).getStatementContext() : null); + } + + public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, + StatementContext statementContext) throws TException { this.planner = planner; this.serializedDescTable = ByteString.copyFrom( new TSerializer().serialize(DescriptorToThriftConverter.toThrift(planner.getDescTable()))); @@ -118,6 +144,7 @@ public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery) throws this.schemaVersion = this.tbl.getBaseSchemaVersion(); this.partitionTopologyVersion = this.tbl.getPartitionTopologyVersion(); this.analzyedQuery = analzyedQuery; + this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(this.scanNode, statementContext); } @VisibleForTesting @@ -135,6 +162,24 @@ public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery) throws this.partitionTopologyVersion = tbl.getPartitionTopologyVersion(); this.scanNode = null; this.analzyedQuery = null; + this.pointQueryKeyTemplate = PointQueryKeyTemplate.unsupported(); + } + + @VisibleForTesting + ShortCircuitQueryContext(OlapScanNode scanNode, StatementContext statementContext) { + this.planner = null; + this.serializedDescTable = ByteString.EMPTY; + this.serializedOutputExpr = ByteString.EMPTY; + this.serializedQueryOptions = ByteString.EMPTY; + this.cacheID = UUID.randomUUID(); + this.scanNode = scanNode; + this.tbl = scanNode.getOlapTable(); + this.tableName = scanNode.getTableNameInPlan(); + this.schemaVersion = tbl.getBaseSchemaVersion(); + this.fileCacheQueryLimitBytes = -1; + this.partitionTopologyVersion = tbl.getPartitionTopologyVersion(); + this.analzyedQuery = null; + this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(scanNode, statementContext); } public boolean isReusable(ConnectContext ctx) { @@ -152,4 +197,247 @@ public void sanitize() { Preconditions.checkNotNull(tbl); Preconditions.checkNotNull(tableName); } + + /** Build state owned by one execution without modifying the cached plan or scan conjuncts. */ + public PointQueryExecutionContext createPointQueryExecutionContext(StatementContext statementContext) { + return pointQueryKeyTemplate.bind(statementContext); + } + + private static class PointQueryKeyTemplate { + private final List keyColumns; + private final List placeholderBindings; + private final List> fixedConstraints; + private final boolean complete; + + private PointQueryKeyTemplate(List keyColumns, + List placeholderBindings, + List> fixedConstraints, boolean complete) { + this.keyColumns = Collections.unmodifiableList(new ArrayList<>(keyColumns)); + this.placeholderBindings = Collections.unmodifiableList(new ArrayList<>(placeholderBindings)); + List> immutableConstraints = new ArrayList<>(fixedConstraints.size()); + for (List constraints : fixedConstraints) { + immutableConstraints.add(Collections.unmodifiableList(new ArrayList<>(constraints))); + } + this.fixedConstraints = Collections.unmodifiableList(immutableConstraints); + this.complete = complete; + } + + private static PointQueryKeyTemplate unsupported() { + return new PointQueryKeyTemplate(Collections.emptyList(), Collections.emptyList(), + Collections.emptyList(), false); + } + + private static PointQueryKeyTemplate create(OlapScanNode scanNode, StatementContext statementContext) { + if (statementContext == null) { + return unsupported(); + } + List keyColumns = scanNode.getOlapTable().getBaseSchemaKeyColumns(); + if (keyColumns.isEmpty()) { + return new PointQueryKeyTemplate(keyColumns, Collections.emptyList(), + Collections.emptyList(), true); + } + if (!statementContext.arePointQueryFixedKeyConstraintsComplete()) { + return unsupported(); + } + + Map keyOrdinals = new TreeMap<>(String.CASE_INSENSITIVE_ORDER); + List> fixedConstraints = new ArrayList<>(keyColumns.size()); + for (int ordinal = 0; ordinal < keyColumns.size(); ordinal++) { + keyOrdinals.put(keyColumns.get(ordinal).getName(), ordinal); + fixedConstraints.add(new ArrayList<>()); + } + + List placeholderBindings = new ArrayList<>(); + for (Map.Entry entry + : statementContext.getIdToComparisonSlot().entrySet()) { + SlotReference slot = entry.getValue(); + if (!slot.getOriginalColumn().isPresent()) { + return unsupported(); + } + Integer ordinal = keyOrdinals.get(slot.getOriginalColumn().get().getName()); + if (ordinal == null) { + return unsupported(); + } + placeholderBindings.add(new PlaceholderKeyBinding(entry.getKey(), ordinal, slot)); + } + + List placeholders = statementContext.getPlaceholders(); + if (placeholderBindings.size() != placeholders.size()) { + return unsupported(); + } + for (Placeholder placeholder : placeholders) { + if (!statementContext.getIdToComparisonSlot().containsKey(placeholder.getPlaceholderId())) { + return unsupported(); + } + } + + for (PointQueryFixedKeyConstraint constraint + : statementContext.getPointQueryFixedKeyConstraints()) { + SlotReference slot = constraint.getSlot(); + if (!slot.getOriginalColumn().isPresent()) { + return unsupported(); + } + Integer ordinal = keyOrdinals.get(slot.getOriginalColumn().get().getName()); + if (ordinal != null) { + fixedConstraints.get(ordinal).add(constraint.getLiteral()); + } else if (!Column.DELETE_SIGN.equals(slot.getOriginalColumn().get().getName())) { + return unsupported(); + } + } + + boolean[] covered = new boolean[keyColumns.size()]; + for (PlaceholderKeyBinding binding : placeholderBindings) { + covered[binding.keyOrdinal] = true; + } + for (int ordinal = 0; ordinal < fixedConstraints.size(); ordinal++) { + covered[ordinal] |= !fixedConstraints.get(ordinal).isEmpty(); + } + for (boolean keyCovered : covered) { + if (!keyCovered) { + return unsupported(); + } + } + return new PointQueryKeyTemplate(keyColumns, placeholderBindings, fixedConstraints, true); + } + + private PointQueryExecutionContext bind(StatementContext statementContext) { + if (!complete || statementContext == null) { + return PointQueryExecutionContext.fallback(); + } + List> valuesByKey = new ArrayList<>(fixedConstraints.size()); + for (List constraints : fixedConstraints) { + valuesByKey.add(new ArrayList<>(constraints)); + } + for (PlaceholderKeyBinding binding : placeholderBindings) { + Expression value = statementContext.getIdToPlaceholderRealExpr().get(binding.placeholderId); + if (!(value instanceof Literal)) { + return PointQueryExecutionContext.fallback(); + } + Literal typedValue = coerceComparisonLiteral(binding.slot, (Literal) value); + if (typedValue == null) { + return PointQueryExecutionContext.fallback(); + } + if (typedValue instanceof NullLiteral) { + return PointQueryExecutionContext.empty(); + } + valuesByKey.get(binding.keyOrdinal).add(typedValue); + } + + Map keyValues = new LinkedHashMap<>(); + for (int ordinal = 0; ordinal < keyColumns.size(); ordinal++) { + List values = valuesByKey.get(ordinal); + if (values.isEmpty()) { + return PointQueryExecutionContext.fallback(); + } + Literal representative = values.get(0); + if (representative instanceof NullLiteral) { + return PointQueryExecutionContext.empty(); + } + for (int i = 1; i < values.size(); i++) { + Boolean equal = sqlEquals(representative, values.get(i)); + if (equal == null) { + return PointQueryExecutionContext.fallback(); + } + if (!equal) { + return PointQueryExecutionContext.empty(); + } + } + LiteralExpr physicalValue = toPhysicalKeyLiteral(representative, keyColumns.get(ordinal)); + if (physicalValue == null) { + return PointQueryExecutionContext.fallback(); + } + keyValues.put(keyColumns.get(ordinal).getName(), physicalValue); + } + return PointQueryExecutionContext.lookup(keyValues); + } + + private static Literal coerceComparisonLiteral(SlotReference slot, Literal value) { + try { + Expression comparison = TypeCoercionUtils.processComparisonPredicate(new EqualTo(slot, value)); + Expression comparisonSlot = comparison.child(0); + // A cast on the physical key can change equality semantics (for example INT 1 + // compared with string '01'). Normal planning must evaluate such comparisons. + return comparisonSlot instanceof SlotReference && comparison.child(1) instanceof Literal + ? (Literal) comparison.child(1) : null; + } catch (Exception e) { + return null; + } + } + + private static Boolean sqlEquals(Literal left, Literal right) { + if (left instanceof NullLiteral || right instanceof NullLiteral) { + return false; + } + try { + Expression comparison = TypeCoercionUtils.processComparisonPredicate(new EqualTo(left, right)); + Expression result = FoldConstantRuleOnFE.evaluateWithoutContext(comparison); + return result instanceof BooleanLiteral ? ((BooleanLiteral) result).getValue() : null; + } catch (Exception e) { + return null; + } + } + + private static LiteralExpr toPhysicalKeyLiteral(Literal literal, Column column) { + try { + LiteralExpr legacyLiteral = literal.toLegacyLiteral(); + Type columnType = column.getType(); + if (!columnType.equals(legacyLiteral.getType()) + && !columnType.matchesType(legacyLiteral.getType())) { + legacyLiteral = LiteralExprUtils.createLiteral(legacyLiteral.getStringValue(), columnType); + } + return legacyLiteral; + } catch (Exception e) { + return null; + } + } + } + + private static class PlaceholderKeyBinding { + private final PlaceholderId placeholderId; + private final int keyOrdinal; + private final SlotReference slot; + + private PlaceholderKeyBinding(PlaceholderId placeholderId, int keyOrdinal, SlotReference slot) { + this.placeholderId = placeholderId; + this.keyOrdinal = keyOrdinal; + this.slot = slot; + } + } + + /** Immutable outcome and typed key tuple for exactly one point-query execution. */ + public static class PointQueryExecutionContext { + public enum Decision { + LOOKUP, + EMPTY, + FALLBACK + } + + private final Decision decision; + private final Map keyValues; + + private PointQueryExecutionContext(Decision decision, Map keyValues) { + this.decision = decision; + this.keyValues = Collections.unmodifiableMap(new LinkedHashMap<>(keyValues)); + } + + public static PointQueryExecutionContext lookup(Map keyValues) { + return new PointQueryExecutionContext(Decision.LOOKUP, keyValues); + } + + public static PointQueryExecutionContext empty() { + return new PointQueryExecutionContext(Decision.EMPTY, Collections.emptyMap()); + } + + public static PointQueryExecutionContext fallback() { + return new PointQueryExecutionContext(Decision.FALLBACK, Collections.emptyMap()); + } + + public Decision getDecision() { + return decision; + } + + public Map getKeyValues() { + return keyValues; + } + } } 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 7f71454ec166c5..7e7e4a83718cc5 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 @@ -1536,21 +1536,40 @@ public void executeAndSendResult(boolean isOutfileQuery, boolean isSendFields, if (statementContext.isShortCircuitQuery()) { ShortCircuitQueryContext shortCircuitQueryContext = statementContext.getShortCircuitQueryContext(); if (shortCircuitQueryContext == null) { - shortCircuitQueryContext = new ShortCircuitQueryContext(planner, (Queriable) parsedStmt); + shortCircuitQueryContext = new ShortCircuitQueryContext( + planner, (Queriable) parsedStmt, statementContext); // ExecuteCommand publishes this same context after a successful first prepared execution. statementContext.setShortCircuitQueryContext(shortCircuitQueryContext); } - coordBase = new PointQueryExecutor(shortCircuitQueryContext, - context.getSessionVariable().getMaxMsgSizeOfResultReceiver()); - context.getState().setIsQuery(true); - } else if (planner instanceof NereidsPlanner && ((NereidsPlanner) planner).getDistributedPlans() != null) { + ShortCircuitQueryContext.PointQueryExecutionContext pointQueryExecutionContext = + statementContext.getPointQueryExecutionContext(); + if (pointQueryExecutionContext == null) { + pointQueryExecutionContext = shortCircuitQueryContext + .createPointQueryExecutionContext(statementContext); + statementContext.setPointQueryExecutionContext(pointQueryExecutionContext); + } + if (pointQueryExecutionContext.getDecision() + == ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK) { + // The physical plan is still a valid normal plan. If an execution value cannot be + // safely reduced to an exact typed key, use the Coordinator instead of failing the + // statement or guessing a lookup key. + statementContext.setShortCircuitQuery(false); + statementContext.setShortCircuitQueryContext(null); + } else { + coordBase = new PointQueryExecutor(shortCircuitQueryContext, pointQueryExecutionContext, + context.getSessionVariable().getMaxMsgSizeOfResultReceiver()); + context.getState().setIsQuery(true); + } + } + if (coordBase == null + && planner instanceof NereidsPlanner && ((NereidsPlanner) planner).getDistributedPlans() != null) { coord = new NereidsCoordinator(context, (NereidsPlanner) planner, context.getStatsErrorEstimator()); profile.addExecutionProfile(coord.getExecutionProfile()); QeProcessorImpl.INSTANCE.registerQuery(context.queryId(), new QueryInfo(context, originStmt.originStmt, coord)); coordBase = coord; - } else { + } else if (coordBase == null) { coord = EnvFactory.getInstance().createCoordinator( context, planner, context.getStatsErrorEstimator()); profile.addExecutionProfile(coord.getExecutionProfile()); diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java index 7905b31b5efeb1..1fde66587cde6d 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java @@ -42,6 +42,7 @@ import org.apache.doris.qe.PreparedStatementContext; import org.apache.doris.qe.SessionVariable; import org.apache.doris.qe.ShortCircuitQueryContext; +import org.apache.doris.qe.ShortCircuitQueryContext.PointQueryExecutionContext; import org.apache.doris.qe.StmtExecutor; import org.apache.doris.thrift.TQueryOptions; @@ -275,12 +276,14 @@ public void testFastPathInstallsCachedShortCircuitContextAcrossExecutions() thro OlapTable table = Mockito.spy(new OlapTable()); Mockito.doReturn("tbl").when(table).getName(); Mockito.doReturn(10).when(table).getBaseSchemaVersion(); + Mockito.doReturn(Collections.emptyList()).when(table).getBaseSchemaKeyColumns(); Mockito.when(scanNode.getPointQueryProjectList()).thenReturn(Collections.emptyList()); Mockito.when(scanNode.getOlapTable()).thenReturn(table); Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); Mockito.when(scanNode.getConjuncts()).thenReturn(Collections.emptyList()); Mockito.when(planner.getScanNodes()).thenReturn(Collections.singletonList(scanNode)); - ShortCircuitQueryContext cachedPlan = new ShortCircuitQueryContext(planner, Mockito.mock(Queriable.class)); + ShortCircuitQueryContext cachedPlan = new ShortCircuitQueryContext( + planner, Mockito.mock(Queriable.class), statementContext); preparedStatement.shortCircuitQueryContext = Optional.of(cachedPlan); StmtExecutor executor = Mockito.mock(StmtExecutor.class); @@ -305,6 +308,45 @@ public void testFastPathInstallsCachedShortCircuitContextAcrossExecutions() thro Mockito.any(), Mockito.any(), Mockito.any(), Mockito.any()); } + @Test + public void testUnsafePointKeyFallsBackToNormalPreparedExecution() throws Exception { + String sql = "select 1"; + LogicalPlan logicalPlan = new NereidsParser().parseSingle(sql); + + ConnectContext connectContext = Mockito.mock(ConnectContext.class); + StatementContext statementContext = new StatementContext(); + statementContext.setShortCircuitQuery(true); + PrepareCommand prepareCommand = new PrepareCommand( + "stmt", logicalPlan, Collections.emptyList(), new OriginStatement(sql, 0)); + PreparedStatementContext preparedStatement = new PreparedStatementContext( + prepareCommand, connectContext, statementContext, "stmt"); + ShortCircuitQueryContext cachedPlan = Mockito.mock(ShortCircuitQueryContext.class); + Mockito.when(cachedPlan.isReusable(connectContext)).thenReturn(true); + Mockito.when(cachedPlan.createPointQueryExecutionContext(Mockito.any(StatementContext.class))) + .thenReturn(PointQueryExecutionContext.fallback()); + preparedStatement.shortCircuitQueryContext = Optional.of(cachedPlan); + + StmtExecutor executor = Mockito.mock(StmtExecutor.class); + Mockito.when(connectContext.getPreparedStementContext("stmt")).thenReturn(preparedStatement); + SessionVariable sessionVariable = new SessionVariable(); + sessionVariable.enableGroupCommitFullPrepare = false; + Mockito.when(connectContext.getSessionVariable()).thenReturn(sessionVariable); + Mockito.when(connectContext.getStatementContext()).thenReturn(statementContext); + Mockito.when(executor.getContext()).thenReturn(connectContext); + + new ExecuteCommand("stmt", prepareCommand, statementContext).run(connectContext, executor); + + Mockito.verify(executor).execute(); + Mockito.verify(executor, Mockito.never()).executeAndSendResult(Mockito.anyBoolean(), Mockito.anyBoolean(), + Mockito.any(), Mockito.any(), Mockito.any(), Mockito.any()); + Assertions.assertFalse(preparedStatement.shortCircuitQueryContext.isPresent(), + "an unsafe typed key must discard direct reuse and run the normal planner"); + Assertions.assertFalse(preparedStatement.getStatementContext().isShortCircuitQuery(), + "normal planning must not inherit the cached execution's short-circuit flag"); + Assertions.assertNull(preparedStatement.getStatementContext().getShortCircuitQueryContext()); + Assertions.assertNull(preparedStatement.getStatementContext().getPointQueryExecutionContext()); + } + private String resolveNextSnapshot(TableScanParams scanParams, AtomicInteger snapshotId) { return scanParams.getOrResolveMapParams(ignored -> ImmutableMap.of( "scan.snapshot-id", String.valueOf(snapshotId.incrementAndGet()))) diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java index fcf3bc20c6e220..e0cef9e3302d96 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java @@ -17,6 +17,8 @@ package org.apache.doris.qe; +import org.apache.doris.catalog.OlapTable; +import org.apache.doris.nereids.StatementContext; import org.apache.doris.planner.OlapScanNode; import org.junit.jupiter.api.Assertions; @@ -34,4 +36,21 @@ public void testCandidateBackendsShuffleDependsOnQuerySelectionOrder() { Mockito.when(scanNode.isScanBackendOrderBySelection()).thenReturn(true); Assertions.assertFalse(PointQueryExecutor.shouldShuffleCandidateBackends(scanNode)); } + + @Test + public void testEmptyDecisionReturnsBeforeTabletPruning() throws Exception { + OlapTable table = Mockito.mock(OlapTable.class); + Mockito.when(table.getBaseSchemaKeyColumns()).thenReturn(java.util.Collections.emptyList()); + OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); + Mockito.when(scanNode.getOlapTable()).thenReturn(table); + Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); + ShortCircuitQueryContext queryContext = new ShortCircuitQueryContext(scanNode, new StatementContext()); + PointQueryExecutor executor = new PointQueryExecutor(queryContext, + ShortCircuitQueryContext.PointQueryExecutionContext.empty(), 1024); + + Mockito.clearInvocations(scanNode); + Assertions.assertNotNull(executor.getNext()); + // lazyEvaluateRangeLocations is the first operation that can resolve a tablet and lead to a BE RPC. + Mockito.verifyNoInteractions(scanNode); + } } diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java index b43f0157073f91..ddc0c3d4c92671 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java @@ -27,6 +27,14 @@ import org.apache.doris.catalog.PrimitiveType; import org.apache.doris.catalog.RandomDistributionInfo; import org.apache.doris.catalog.SinglePartitionInfo; +import org.apache.doris.nereids.StatementContext; +import org.apache.doris.nereids.trees.expressions.Placeholder; +import org.apache.doris.nereids.trees.expressions.SlotReference; +import org.apache.doris.nereids.trees.expressions.StatementScopeIdGenerator; +import org.apache.doris.nereids.trees.expressions.literal.DecimalLiteral; +import org.apache.doris.nereids.trees.expressions.literal.IntegerLiteral; +import org.apache.doris.nereids.trees.expressions.literal.NullLiteral; +import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.planner.OlapScanNode; import org.apache.doris.planner.Planner; import org.apache.doris.thrift.TQueryOptions; @@ -37,6 +45,7 @@ import org.junit.jupiter.api.Test; import org.mockito.Mockito; +import java.math.BigDecimal; import java.util.Collections; import java.util.List; @@ -48,6 +57,14 @@ private OlapTable table(String name, int schemaVersion) { return table; } + private OlapTable pointQueryTable(List keyColumns) { + OlapTable table = Mockito.mock(OlapTable.class); + Mockito.when(table.getName()).thenReturn("tbl"); + Mockito.when(table.getBaseSchemaKeyColumns()).thenReturn(keyColumns); + Mockito.when(table.getBaseSchemaVersion()).thenReturn(1); + return table; + } + private ConnectContext connectContext(long fileCacheQueryLimitBytes) { ConnectContext ctx = new ConnectContext(); SessionVariable sessionVariable = new SessionVariable(); @@ -117,4 +134,172 @@ public void testSerializedQueryOptionsKeepBitmapOpCountVersion() throws Exceptio Assertions.assertTrue(serializedQueryOptions.isSetNewVersionBitmapOpCount()); Assertions.assertTrue(serializedQueryOptions.isNewVersionBitmapOpCount()); } + + @Test + public void testPreparedKeyTemplateKeepsFixedConstraintsAcrossExecutions() { + Column parameterKey = new Column("parameter_key", PrimitiveType.INT); + parameterKey.setIsKey(true); + Column policyKey = new Column("policy_key", PrimitiveType.INT); + policyKey.setIsKey(true); + List schema = List.of(parameterKey, policyKey); + OlapTable table = pointQueryTable(schema); + SlotReference parameterSlot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, parameterKey, Collections.emptyList()); + SlotReference policySlot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, policyKey, Collections.emptyList()); + + PlaceholderId placeholderId = new PlaceholderId(0); + StatementContext templateContext = new StatementContext(); + templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); + templateContext.getIdToComparisonSlot().put(placeholderId, parameterSlot); + // This models a restrictive policy on the same key as the placeholder, plus a + // policy-fixed column in a composite key. + templateContext.addPointQueryFixedKeyConstraint(parameterSlot, new IntegerLiteral(1)); + templateContext.addPointQueryFixedKeyConstraint(policySlot, new IntegerLiteral(9)); + + OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); + Mockito.when(scanNode.getOlapTable()).thenReturn(table); + Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); + ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode, templateContext); + + StatementContext first = execution(placeholderId, new IntegerLiteral(1)); + ShortCircuitQueryContext.PointQueryExecutionContext firstExecution = + cached.createPointQueryExecutionContext(first); + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, + firstExecution.getDecision()); + Assertions.assertEquals("1", firstExecution.getKeyValues().get("parameter_key").getStringValue()); + Assertions.assertEquals("9", firstExecution.getKeyValues().get("policy_key").getStringValue()); + + StatementContext second = execution(placeholderId, new IntegerLiteral(2)); + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.EMPTY, + cached.createPointQueryExecutionContext(second).getDecision()); + + // Reusing the same prepared handle with 1 -> 2 -> 1 must not contaminate the template. + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, + cached.createPointQueryExecutionContext(first).getDecision()); + + StatementContext nullValue = execution(placeholderId, new NullLiteral()); + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.EMPTY, + cached.createPointQueryExecutionContext(nullValue).getDecision()); + Mockito.verify(scanNode, Mockito.never()).getConjuncts(); + } + + @Test + public void testFixedOnlyKeyTemplate() { + Column key = new Column("k", PrimitiveType.INT); + key.setIsKey(true); + OlapTable table = pointQueryTable(Collections.singletonList(key)); + SlotReference slot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); + StatementContext templateContext = new StatementContext(); + templateContext.addPointQueryFixedKeyConstraint(slot, new IntegerLiteral(7)); + ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode(table), templateContext); + + ShortCircuitQueryContext.PointQueryExecutionContext execution = + cached.createPointQueryExecutionContext(new StatementContext()); + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, + execution.getDecision()); + Assertions.assertEquals("7", execution.getKeyValues().get("k").getStringValue()); + } + + @Test + public void testPlaceholderOnlyKeyTemplate() { + Column key = new Column("k", PrimitiveType.INT); + key.setIsKey(true); + OlapTable table = pointQueryTable(Collections.singletonList(key)); + SlotReference slot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); + PlaceholderId placeholderId = new PlaceholderId(0); + StatementContext templateContext = new StatementContext(); + templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); + templateContext.getIdToComparisonSlot().put(placeholderId, slot); + ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode(table), templateContext); + + ShortCircuitQueryContext.PointQueryExecutionContext execution = + cached.createPointQueryExecutionContext(execution(placeholderId, new IntegerLiteral(8))); + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, + execution.getDecision()); + Assertions.assertEquals("8", execution.getKeyValues().get("k").getStringValue()); + } + + @Test + public void testInexactPhysicalKeyFallsBack() { + Column key = new Column("k", PrimitiveType.INT); + key.setIsKey(true); + OlapTable table = pointQueryTable(Collections.singletonList(key)); + SlotReference slot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); + PlaceholderId placeholderId = new PlaceholderId(0); + StatementContext templateContext = new StatementContext(); + templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); + templateContext.getIdToComparisonSlot().put(placeholderId, slot); + OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); + Mockito.when(scanNode.getOlapTable()).thenReturn(table); + Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); + ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode, templateContext); + + StatementContext execution = execution(placeholderId, new DecimalLiteral(new BigDecimal("1.2"))); + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK, + cached.createPointQueryExecutionContext(execution).getDecision()); + } + + @Test + public void testNonSlotFixedConstraintFallsBack() { + Column key = new Column("k", PrimitiveType.INT); + key.setIsKey(true); + OlapTable table = pointQueryTable(Collections.singletonList(key)); + SlotReference slot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); + PlaceholderId placeholderId = new PlaceholderId(0); + StatementContext templateContext = new StatementContext(); + templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); + templateContext.getIdToComparisonSlot().put(placeholderId, slot); + // ExpressionAnalyzer uses this marker for a fixed predicate such as + // CAST(k AS CHAR(1)) = '1', whose cast cannot identify an exact physical key. + templateContext.markPointQueryFixedKeyConstraintsIncomplete(); + OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); + Mockito.when(scanNode.getOlapTable()).thenReturn(table); + Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); + ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode, templateContext); + + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK, + cached.createPointQueryExecutionContext( + execution(placeholderId, new IntegerLiteral(1))).getDecision()); + } + + @Test + public void testFixedNonKeyConstraintFallsBack() { + Column key = new Column("k", PrimitiveType.INT); + key.setIsKey(true); + Column value = new Column("v", PrimitiveType.INT); + OlapTable table = pointQueryTable(Collections.singletonList(key)); + SlotReference keySlot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); + SlotReference valueSlot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, value, Collections.emptyList()); + PlaceholderId placeholderId = new PlaceholderId(0); + StatementContext templateContext = new StatementContext(); + templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); + templateContext.getIdToComparisonSlot().put(placeholderId, keySlot); + templateContext.addPointQueryFixedKeyConstraint(valueSlot, new IntegerLiteral(1)); + ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode(table), templateContext); + + Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK, + cached.createPointQueryExecutionContext( + execution(placeholderId, new IntegerLiteral(1))).getDecision()); + } + + private OlapScanNode scanNode(OlapTable table) { + OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); + Mockito.when(scanNode.getOlapTable()).thenReturn(table); + Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); + return scanNode; + } + + private StatementContext execution(PlaceholderId placeholderId, + org.apache.doris.nereids.trees.expressions.Expression value) { + StatementContext context = new StatementContext(); + context.getIdToPlaceholderRealExpr().put(placeholderId, value); + return context; + } } diff --git a/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy new file mode 100644 index 00000000000000..3b7b4d7ad88e0c --- /dev/null +++ b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy @@ -0,0 +1,155 @@ +// 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. + +suite("prepared_point_query_row_policy", "p0") { + def dbName = context.config.getDbNameByFile(context.file) + def policyName = "prepared_point_query_key_guard" + def user = "prepared_point_query_policy_user" + def password = "Prepared_policy_123!" + + sql "DROP ROW POLICY IF EXISTS ${policyName} ON ${dbName}.prepared_point_query_row_policy FOR ${user}" + sql "DROP USER IF EXISTS ${user}" + sql "DROP TABLE IF EXISTS prepared_point_query_row_policy" + sql """ + CREATE TABLE prepared_point_query_row_policy ( + tenant_id INT NOT NULL, + item_id INT NOT NULL, + value VARCHAR(32) + ) ENGINE=OLAP + UNIQUE KEY(tenant_id, item_id) + DISTRIBUTED BY HASH(tenant_id, item_id) BUCKETS 3 + PROPERTIES ( + "replication_num" = "1", + "light_schema_change" = "true", + "store_row_column" = "true", + "enable_unique_key_merge_on_write" = "true" + ) + """ + sql """ + INSERT INTO prepared_point_query_row_policy VALUES + (1, 10, 'allowed'), (1, 20, 'other'), (2, 10, 'hidden'), (10, 10, 'cast-match') + """ + sql "CREATE USER ${user} IDENTIFIED BY '${password}'" + sql "GRANT SELECT_PRIV ON internal.${dbName}.prepared_point_query_row_policy TO ${user}" + sql """ + CREATE ROW POLICY ${policyName} ON ${dbName}.prepared_point_query_row_policy + AS RESTRICTIVE TO ${user} USING (tenant_id = 1) + """ + + if (isCloudMode()) { + def clusters = sql "SHOW CLUSTERS" + assertTrue(!clusters.isEmpty()) + sql "GRANT USAGE_PRIV ON CLUSTER `${clusters[0][0]}` TO ${user}" + } + sql "SET GLOBAL enable_server_side_prepared_statement = true" + sql "SYNC" + + String url = getServerPrepareJdbcUrl(context.config.jdbcUrl, dbName) + String explainUrl = url.replace("useServerPrepStmts=true", "useServerPrepStmts=false") + connect(user, password, explainUrl) { + def explainRows = sql """ + EXPLAIN SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + FROM prepared_point_query_row_policy WHERE tenant_id = 1 AND item_id = 10 + """ + assertTrue(explainRows.toString().contains("SHORT-CIRCUIT")) + } + + connect(user, password, url) { + def prepared = prepareStatement """ + SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + FROM prepared_point_query_row_policy + WHERE tenant_id = ? AND item_id = ? + """ + assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) + + def readRows = { Integer tenant, int item -> + if (tenant == null) { + prepared.setNull(1, java.sql.Types.INTEGER) + } else { + prepared.setInt(1, tenant) + } + prepared.setInt(2, item) + def rows = [] + prepared.executeQuery().withCloseable { result -> + assertEquals(3, result.getMetaData().getColumnCount()) + assertEquals("tenant_id", result.getMetaData().getColumnLabel(1)) + assertEquals("item_id", result.getMetaData().getColumnLabel(2)) + assertEquals("value", result.getMetaData().getColumnLabel(3)) + while (result.next()) { + rows.add([result.getInt(1), result.getInt(2), result.getString(3)]) + } + } + return rows + } + + assertEquals([[1, 10, "allowed"]], readRows(1, 10)) + assertEquals([], readRows(2, 10)) + assertEquals([[1, 10, "allowed"]], readRows(1, 10)) + assertEquals([], readRows(null, 10)) + + // A non-integral parameter cannot be represented by the physical INT lookup key. It is + // evaluated by the normal planner, exercising direct-reuse FALLBACK without an error. + prepared.setBigDecimal(1, new BigDecimal("1.2")) + prepared.setInt(2, 10) + def fallbackRows = [] + prepared.executeQuery().withCloseable { result -> + while (result.next()) { + fallbackRows.add([result.getInt(1), result.getInt(2), result.getString(3)]) + } + } + assertEquals([], fallbackRows) + prepared.close() + } + + // A fixed predicate on a cast key is not an exact physical-key constraint. Keep it on + // the normal path unless a future proof can establish that the cast is lossless/injective. + sql "DROP ROW POLICY IF EXISTS ${policyName} ON ${dbName}.prepared_point_query_row_policy FOR ${user}" + sql """ + CREATE ROW POLICY ${policyName} ON ${dbName}.prepared_point_query_row_policy + AS RESTRICTIVE TO ${user} USING (CAST(tenant_id AS CHAR(1)) = '1') + """ + sql "SYNC" + connect(user, password, explainUrl) { + def explainRows = sql """ + EXPLAIN SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + FROM prepared_point_query_row_policy WHERE tenant_id = 1 AND item_id = 10 + """ + assertFalse(explainRows.toString().contains("SHORT-CIRCUIT")) + } + + connect(user, password, url) { + def prepared = prepareStatement """ + SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + FROM prepared_point_query_row_policy + WHERE tenant_id = ? AND item_id = ? + """ + assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) + prepared.setInt(1, 10) + prepared.setInt(2, 10) + def rows = [] + prepared.executeQuery().withCloseable { result -> + while (result.next()) { + rows.add([result.getInt(1), result.getInt(2), result.getString(3)]) + } + } + prepared.close() + assertEquals([[10, 10, "cast-match"]], rows) + } + + sql "DROP ROW POLICY IF EXISTS ${policyName} ON ${dbName}.prepared_point_query_row_policy FOR ${user}" + sql "DROP USER IF EXISTS ${user}" +} From be8a0aaa28ec54994b1c4db1162dbbf16f7a2716 Mon Sep 17 00:00:00 2001 From: morrySnow Date: Sat, 12 Sep 2026 00:51:13 +0800 Subject: [PATCH 2/6] [fix](fe) Revalidate security before prepared point-query reuse Issue Number: None Related PR: None Problem Summary: Server-side prepared point queries retain a direct short-circuit execution context after their first execution. Reusing that context bypassed the normal privilege and policy analysis passes, so a SELECT revocation or a changed row-filter or data-mask policy could leave the retained plan authorized with stale decisions. Record the planning identity, checked columns, complete row-filter answers, and data-mask answers (including negative answers), then revalidate them before every direct reuse. Missing dependencies and authorization-source failures reject reuse and fall back to normal planning, where the authoritative check returns the standard result or error. Prepared point queries now honor current SELECT privileges and row-filter/data-mask policies on every execution. - Test: Unit tests and regression test - `SecurityDependencyContextTest`, `ExecuteCommandTest`, and `ShortCircuitQueryContextTest` - `prepared_short_circuit_security_refresh` - Behavior changed: Yes. A stale prepared point-query context is rebuilt after security decisions change. - Does this need documentation: No --- .../nereids/SecurityDependencyContext.java | 202 ++++++++++++++++++ .../doris/nereids/StatementContext.java | 7 + .../rules/rewrite/CheckPrivileges.java | 1 + .../plans/logical/LogicalCheckPolicy.java | 5 + .../doris/qe/ShortCircuitQueryContext.java | 27 ++- .../SecurityDependencyContextTest.java | 175 +++++++++++++++ .../plans/commands/ExecuteCommandTest.java | 32 +++ .../qe/ShortCircuitQueryContextTest.java | 17 +- ...repared_short_circuit_security_refresh.out | 15 ++ ...ared_short_circuit_security_refresh.groovy | 131 ++++++++++++ 10 files changed, 610 insertions(+), 2 deletions(-) create mode 100644 fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java create mode 100644 fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java create mode 100644 regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out create mode 100644 regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java new file mode 100644 index 00000000000000..5ec32baa64f13d --- /dev/null +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java @@ -0,0 +1,202 @@ +// 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.nereids; + +import org.apache.doris.analysis.UserIdentity; +import org.apache.doris.authorization.DataMaskSpec; +import org.apache.doris.authorization.RowFilterSpec; +import org.apache.doris.catalog.DatabaseIf; +import org.apache.doris.catalog.Env; +import org.apache.doris.catalog.TableIf; +import org.apache.doris.common.UserException; +import org.apache.doris.datasource.CatalogIf; +import org.apache.doris.nereids.SqlCacheContext.FullColumnName; +import org.apache.doris.nereids.SqlCacheContext.FullTableName; +import org.apache.doris.nereids.rules.analysis.UserAuthentication; +import org.apache.doris.qe.ConnectContext; + +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableSet; +import com.google.common.collect.Maps; +import org.apache.commons.collections4.CollectionUtils; + +import java.util.LinkedHashMap; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.Set; + +/** + * Security decisions which an analyzed plan depends on. + * + *

Unlike {@link SqlCacheContext}, this context exists independently of the SQL result-cache switch. A prepared + * short-circuit plan can otherwise outlive the privilege and data-policy decisions made while it was analyzed. + * Callers record both positive and negative policy answers so that adding a policy invalidates a plan which was + * built before that policy existed. + */ +public class SecurityDependencyContext { + private final UserIdentity userIdentity; + private final Map> checkedPrivileges = Maps.newLinkedHashMap(); + private final Map> rowPolicies = Maps.newLinkedHashMap(); + private final Map> dataMaskPolicies = Maps.newLinkedHashMap(); + private boolean complete; + + /** SecurityDependencyContext */ + public SecurityDependencyContext(UserIdentity userIdentity) { + this.userIdentity = userIdentity; + this.complete = userIdentity != null; + } + + /** Record the columns whose SELECT privilege was checked while the plan was analyzed. */ + public synchronized void addCheckedPrivilege(TableIf table, Set usedColumns) { + Optional tableName = qualifiedName(table); + if (!tableName.isPresent()) { + complete = false; + return; + } + Set existing = checkedPrivileges.get(tableName.get()); + if (existing == null) { + checkedPrivileges.put(tableName.get(), ImmutableSet.copyOf(usedColumns)); + } else { + checkedPrivileges.put(tableName.get(), ImmutableSet.builder() + .addAll(existing).addAll(usedColumns).build()); + } + } + + /** Record the complete row-filter answer, including an empty answer. */ + public synchronized void setRowPolicies( + String catalog, String database, String table, List policies) { + rowPolicies.put(new FullTableName(catalog, database, table), ImmutableList.copyOf(policies)); + } + + /** Record the mask answer for a column, including the absence of a mask. */ + public synchronized void addDataMask( + String catalog, String database, String table, String column, Optional mask) { + dataMaskPolicies.put(new FullColumnName( + catalog, database, table, column.toLowerCase(Locale.ROOT)), mask); + } + + /** Freeze the decisions used by a completed plan before storing them in a reusable context. */ + public synchronized SecurityDependencyContext snapshot() { + SecurityDependencyContext snapshot = new SecurityDependencyContext(userIdentity); + snapshot.complete = complete; + for (Map.Entry> entry : checkedPrivileges.entrySet()) { + snapshot.checkedPrivileges.put(entry.getKey(), ImmutableSet.copyOf(entry.getValue())); + } + for (Map.Entry> entry : rowPolicies.entrySet()) { + snapshot.rowPolicies.put(entry.getKey(), ImmutableList.copyOf(entry.getValue())); + } + snapshot.dataMaskPolicies.putAll(dataMaskPolicies); + return snapshot; + } + + /** Freeze the decisions for a prepared short-circuit plan, failing closed if authorization was not recorded. */ + public synchronized SecurityDependencyContext snapshotForShortCircuit() { + SecurityDependencyContext snapshot = snapshot(); + if (checkedPrivileges.isEmpty()) { + snapshot.complete = false; + } + return snapshot; + } + + /** + * Revalidate every security decision before a cached plan bypasses analysis. + * + *

A false result does not deny the statement itself. It rejects only the cached plan, after which the normal + * planning path performs the authoritative checks and returns the usual user-facing error when access was + * revoked. Authorization-source failures also reject reuse, so this fast path always fails closed. + */ + public synchronized boolean isValid(ConnectContext connectContext) { + if (!complete || connectContext == null) { + return false; + } + try { + if (!Objects.equals(userIdentity, connectContext.getCurrentUserIdentity())) { + return false; + } + Env env = connectContext.getEnv(); + for (Map.Entry> entry : checkedPrivileges.entrySet()) { + TableIf table = findTable(env, entry.getKey()); + if (table == null) { + return false; + } + UserAuthentication.checkPermission(table, connectContext, entry.getValue()); + } + for (Map.Entry> entry : rowPolicies.entrySet()) { + FullTableName table = entry.getKey(); + List current = env.getAccessManager().evalRowFilterPolicies( + userIdentity, table.catalog, table.db, table.table); + if (!CollectionUtils.isEqualCollection(entry.getValue(), current)) { + return false; + } + } + return dataMasksAreValid(env); + } catch (UserException | RuntimeException e) { + return false; + } + } + + private boolean dataMasksAreValid(Env env) { + Map> columnsByTable = new LinkedHashMap<>(); + for (FullColumnName column : dataMaskPolicies.keySet()) { + columnsByTable.computeIfAbsent(new FullTableName(column.catalog, column.db, column.table), + table -> new LinkedHashSet<>()).add(column.column); + } + for (Map.Entry> entry : columnsByTable.entrySet()) { + FullTableName table = entry.getKey(); + Map current = env.getAccessManager().evalDataMaskPolicies( + userIdentity, table.catalog, table.db, table.table, entry.getValue()); + for (String column : entry.getValue()) { + Optional currentMask = Optional.ofNullable( + current.get(column.toLowerCase(Locale.ROOT))); + if (!Objects.equals(dataMaskPolicies.get( + new FullColumnName(table.catalog, table.db, table.table, column)), currentMask)) { + return false; + } + } + } + return true; + } + + private Optional qualifiedName(TableIf table) { + if (table == null) { + return Optional.empty(); + } + DatabaseIf database = table.getDatabase(); + if (database == null || database.getCatalog() == null) { + return Optional.empty(); + } + return Optional.of(new FullTableName( + database.getCatalog().getName(), database.getFullName(), table.getName())); + } + + private TableIf findTable(Env env, FullTableName fullTableName) { + CatalogIf> catalog = env.getCatalogMgr().getCatalog(fullTableName.catalog); + if (catalog == null) { + return null; + } + Optional> database = catalog.getDb(fullTableName.db); + if (!database.isPresent()) { + return null; + } + return database.get().getTable(fullTableName.table).orElse(null); + } +} diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java index fba49791c36056..32a121914ba46f 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java @@ -184,6 +184,7 @@ public enum TableFrom { private final Map rewrittenCteConsumer = new HashMap<>(); private final Set viewDdlSqlSet = Sets.newHashSet(); private final SqlCacheContext sqlCacheContext; + private final SecurityDependencyContext securityDependencyContext; // generate for next id for prepared statement's placeholders, which is // connection level @@ -393,6 +394,8 @@ private StatementContext(ConnectContext connectContext, OriginStatement originSt this.connectContext = connectContext; this.originStatement = originStatement; exprIdGenerator = ExprId.createGenerator(initialId); + this.securityDependencyContext = new SecurityDependencyContext( + connectContext == null ? null : connectContext.getCurrentUserIdentity()); if (connectContext != null && connectContext.getSessionVariable() != null) { if (CacheAnalyzer.canUseSqlCache(connectContext.getSessionVariable())) { // cannot set the queryId here because the queryId for the current query is set @@ -711,6 +714,10 @@ public Optional getSqlCacheContext() { return Optional.ofNullable(sqlCacheContext); } + public SecurityDependencyContext getSecurityDependencyContext() { + return securityDependencyContext; + } + public boolean isDpHyp() { return isDpHyp; } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/CheckPrivileges.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/CheckPrivileges.java index bce4862aeac94b..1ea1e217887d68 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/CheckPrivileges.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/CheckPrivileges.java @@ -144,6 +144,7 @@ private void checkColumnPrivileges(TableIf table, Set usedColumns) { throw new AnalysisException(e.getMessage(), e); } StatementContext statementContext = cascadesContext.getStatementContext(); + statementContext.getSecurityDependencyContext().addCheckedPrivilege(table, usedColumns); Optional sqlCacheContext = statementContext.getSqlCacheContext(); if (sqlCacheContext.isPresent()) { sqlCacheContext.get().addCheckPrivilegeTablesOrViews(table, usedColumns); diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java index 40d013ce37c1a2..5e387bc482ef70 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java @@ -25,6 +25,7 @@ import org.apache.doris.datasource.CatalogIf; import org.apache.doris.mysql.privilege.AccessControllerManager; import org.apache.doris.nereids.CascadesContext; +import org.apache.doris.nereids.SecurityDependencyContext; import org.apache.doris.nereids.SqlCacheContext; import org.apache.doris.nereids.StatementContext; import org.apache.doris.nereids.analyzer.UnboundAlias; @@ -188,6 +189,7 @@ public RelatedPolicy findPolicy(LogicalPlan logicalPlan, CascadesContext cascade = ImmutableList.builderWithExpectedSize(logicalPlan.getOutput().size()); StatementContext statementContext = cascadesContext.getStatementContext(); + SecurityDependencyContext securityDependencyContext = statementContext.getSecurityDependencyContext(); Optional sqlCacheContext = statementContext.getSqlCacheContext(); boolean hasDataMask = false; // One question for the whole relation rather than one per column: that is what the contract offers @@ -218,6 +220,8 @@ public RelatedPolicy findPolicy(LogicalPlan logicalPlan, CascadesContext cascade if (sqlCacheContext.isPresent()) { sqlCacheContext.get().addDataMaskPolicy(ctlName, dbName, tableName, slot.getName(), dataMaskPolicy); } + securityDependencyContext.addDataMask( + ctlName, dbName, tableName, slot.getName(), dataMaskPolicy); } List rowPolicies = accessManager.evalRowFilterPolicies( @@ -225,6 +229,7 @@ public RelatedPolicy findPolicy(LogicalPlan logicalPlan, CascadesContext cascade if (sqlCacheContext.isPresent()) { sqlCacheContext.get().setRowFilterPolicy(ctlName, dbName, tableName, rowPolicies); } + securityDependencyContext.setRowPolicies(ctlName, dbName, tableName, rowPolicies); return new RelatedPolicy( Optional.ofNullable(CollectionUtils.isEmpty(rowPolicies) diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java index 88b6435bc378e5..e991aaa21c9ed9 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java @@ -27,6 +27,7 @@ import org.apache.doris.catalog.OlapTable; import org.apache.doris.catalog.Type; import org.apache.doris.nereids.NereidsPlanner; +import org.apache.doris.nereids.SecurityDependencyContext; import org.apache.doris.nereids.StatementContext; import org.apache.doris.nereids.StatementContext.PointQueryFixedKeyConstraint; import org.apache.doris.nereids.rules.expression.rules.FoldConstantRuleOnFE; @@ -80,6 +81,7 @@ public class ShortCircuitQueryContext { public final String tableName; private final long fileCacheQueryLimitBytes; private final long partitionTopologyVersion; + private final SecurityDependencyContext securityDependencyContext; public final OlapScanNode scanNode; public final Queriable analzyedQuery; @@ -113,6 +115,18 @@ public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery) throws public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, StatementContext statementContext) throws TException { + this(planner, analzyedQuery, statementContext, + statementContext == null ? null : statementContext.getSecurityDependencyContext()); + } + + @VisibleForTesting + public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, + SecurityDependencyContext securityDependencyContext) throws TException { + this(planner, analzyedQuery, null, securityDependencyContext); + } + + private ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, + StatementContext statementContext, SecurityDependencyContext securityDependencyContext) throws TException { this.planner = planner; this.serializedDescTable = ByteString.copyFrom( new TSerializer().serialize(DescriptorToThriftConverter.toThrift(planner.getDescTable()))); @@ -145,11 +159,19 @@ public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.partitionTopologyVersion = this.tbl.getPartitionTopologyVersion(); this.analzyedQuery = analzyedQuery; this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(this.scanNode, statementContext); + this.securityDependencyContext = securityDependencyContext == null + ? null : securityDependencyContext.snapshotForShortCircuit(); } @VisibleForTesting ShortCircuitQueryContext(OlapTable tbl, String tableName, int schemaVersion, long fileCacheQueryLimitBytes) { + this(tbl, tableName, schemaVersion, fileCacheQueryLimitBytes, null); + } + + @VisibleForTesting + ShortCircuitQueryContext(OlapTable tbl, String tableName, int schemaVersion, + long fileCacheQueryLimitBytes, SecurityDependencyContext securityDependencyContext) { this.planner = null; this.serializedDescTable = ByteString.EMPTY; this.serializedOutputExpr = ByteString.EMPTY; @@ -163,6 +185,7 @@ public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.scanNode = null; this.analzyedQuery = null; this.pointQueryKeyTemplate = PointQueryKeyTemplate.unsupported(); + this.securityDependencyContext = securityDependencyContext; } @VisibleForTesting @@ -180,6 +203,7 @@ public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.partitionTopologyVersion = tbl.getPartitionTopologyVersion(); this.analzyedQuery = null; this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(scanNode, statementContext); + this.securityDependencyContext = null; } public boolean isReusable(ConnectContext ctx) { @@ -187,7 +211,8 @@ public boolean isReusable(ConnectContext ctx) { && this.tbl.getBaseSchemaVersion() == this.schemaVersion && Objects.equals(this.tableName, this.tbl.getName()) && this.fileCacheQueryLimitBytes == ctx.getSessionVariable().fileCacheQueryLimitBytes - && this.tbl.getPartitionTopologyVersion() == this.partitionTopologyVersion; + && this.tbl.getPartitionTopologyVersion() == this.partitionTopologyVersion + && (securityDependencyContext == null || securityDependencyContext.isValid(ctx)); } public void sanitize() { diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java new file mode 100644 index 00000000000000..42e628b5bcf0a7 --- /dev/null +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java @@ -0,0 +1,175 @@ +// 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.nereids; + +import org.apache.doris.analysis.UserIdentity; +import org.apache.doris.authorization.DataMaskSpec; +import org.apache.doris.authorization.RowFilterSpec; +import org.apache.doris.catalog.DatabaseIf; +import org.apache.doris.catalog.Env; +import org.apache.doris.catalog.TableIf; +import org.apache.doris.common.UserException; +import org.apache.doris.datasource.CatalogIf; +import org.apache.doris.datasource.CatalogMgr; +import org.apache.doris.mysql.privilege.AccessControllerManager; +import org.apache.doris.mysql.privilege.PrivPredicate; +import org.apache.doris.qe.ConnectContext; +import org.apache.doris.qe.SessionVariable; + +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentMatchers; +import org.mockito.Mockito; + +import java.util.Optional; + +public class SecurityDependencyContextTest { + private static final UserIdentity USER = UserIdentity.createAnalyzedUserIdentWithIp("reader", "%"); + private static final String CATALOG = "internal"; + private static final String DATABASE = "db"; + private static final String TABLE = "tbl"; + private static final String COLUMN = "value"; + + @Test + public void testUnchangedPoliciesAreValid() { + RowFilterSpec rowFilter = RowFilterSpec.restrictive("row:1", "tenant_id = 1"); + DataMaskSpec dataMask = new DataMaskSpec("mask:1", "null"); + SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of(rowFilter)); + dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.of(dataMask)); + + ConnectContext connectContext = contextWithPolicies( + ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1")), + ImmutableMap.of(COLUMN, new DataMaskSpec("mask:1", "null"))); + + Assertions.assertTrue(dependencies.snapshot().isValid(connectContext)); + } + + @Test + public void testAddedRowPolicyInvalidatesNegativeSnapshot() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of()); + ConnectContext connectContext = contextWithPolicies( + ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1")), ImmutableMap.of()); + + Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + } + + @Test + public void testChangedRowPolicyInvalidatesSnapshot() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, + ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1"))); + ConnectContext connectContext = contextWithPolicies( + ImmutableList.of(RowFilterSpec.restrictive("row:2", "tenant_id = 2")), ImmutableMap.of()); + + Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + } + + @Test + public void testAddedDataMaskInvalidatesNegativeSnapshot() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.empty()); + ConnectContext connectContext = contextWithPolicies(ImmutableList.of(), + ImmutableMap.of(COLUMN, new DataMaskSpec("mask:1", "null"))); + + Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + } + + @Test + public void testDifferentExecutingIdentityInvalidatesSnapshot() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + ConnectContext connectContext = Mockito.mock(ConnectContext.class); + Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn( + UserIdentity.createAnalyzedUserIdentWithIp("other", "%")); + + Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + } + + @Test + public void testMissingPlanningIdentityFailsClosed() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(null); + ConnectContext connectContext = Mockito.mock(ConnectContext.class); + Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(UserIdentity.ROOT); + + Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + } + + @Test + public void testMissingPrivilegeRecordingDisablesShortCircuitReuse() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + ConnectContext connectContext = Mockito.mock(ConnectContext.class); + Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); + + Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(connectContext)); + } + + @Test + @SuppressWarnings({"rawtypes", "unchecked"}) + public void testPrivilegeRevocationOrAuthorizationFailureInvalidatesSnapshot() throws Exception { + CatalogIf catalog = Mockito.mock(CatalogIf.class); + DatabaseIf database = Mockito.mock(DatabaseIf.class); + TableIf table = Mockito.mock(TableIf.class); + CatalogMgr catalogMgr = Mockito.mock(CatalogMgr.class); + AccessControllerManager accessManager = Mockito.mock(AccessControllerManager.class); + Env env = Mockito.mock(Env.class); + ConnectContext connectContext = Mockito.mock(ConnectContext.class); + + Mockito.when(catalog.getName()).thenReturn(CATALOG); + Mockito.when(catalog.getDb(DATABASE)).thenReturn(Optional.of(database)); + Mockito.when(database.getCatalog()).thenReturn(catalog); + Mockito.when(database.getFullName()).thenReturn(DATABASE); + Mockito.when(database.getTable(TABLE)).thenReturn(Optional.of(table)); + Mockito.when(table.getDatabase()).thenReturn(database); + Mockito.when(table.getName()).thenReturn(TABLE); + Mockito.when(catalogMgr.getCatalog(CATALOG)).thenReturn(catalog); + Mockito.when(env.getCatalogMgr()).thenReturn(catalogMgr); + Mockito.when(env.getAccessManager()).thenReturn(accessManager); + Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); + Mockito.when(connectContext.getEnv()).thenReturn(env); + Mockito.when(connectContext.getSessionVariable()).thenReturn(new SessionVariable()); + + SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + dependencies.addCheckedPrivilege(table, ImmutableSet.of(COLUMN)); + SecurityDependencyContext snapshot = dependencies.snapshot(); + Assertions.assertTrue(snapshot.isValid(connectContext)); + + Mockito.doThrow(new UserException("SELECT was revoked or authorization is unavailable")) + .when(accessManager).checkColumnsPriv( + connectContext, CATALOG, DATABASE, TABLE, ImmutableSet.of(COLUMN), PrivPredicate.SELECT); + Assertions.assertFalse(snapshot.isValid(connectContext)); + } + + private ConnectContext contextWithPolicies( + ImmutableList rowFilters, ImmutableMap dataMasks) { + AccessControllerManager accessManager = Mockito.mock(AccessControllerManager.class); + Mockito.when(accessManager.evalRowFilterPolicies(USER, CATALOG, DATABASE, TABLE)).thenReturn(rowFilters); + Mockito.when(accessManager.evalDataMaskPolicies( + ArgumentMatchers.eq(USER), ArgumentMatchers.eq(CATALOG), ArgumentMatchers.eq(DATABASE), + ArgumentMatchers.eq(TABLE), ArgumentMatchers.anySet())).thenReturn(dataMasks); + Env env = Mockito.mock(Env.class); + Mockito.when(env.getAccessManager()).thenReturn(accessManager); + ConnectContext connectContext = Mockito.mock(ConnectContext.class); + Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); + Mockito.when(connectContext.getEnv()).thenReturn(env); + return connectContext; + } +} diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java index 1fde66587cde6d..892815b5a89280 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java @@ -347,6 +347,38 @@ public void testUnsafePointKeyFallsBackToNormalPreparedExecution() throws Except Assertions.assertNull(preparedStatement.getStatementContext().getPointQueryExecutionContext()); } + @Test + public void testInvalidSecurityDependenciesRefreshInsteadOfDirectReuse() throws Exception { + String sql = "select * from tbl"; + LogicalPlan logicalPlan = new NereidsParser().parseSingle(sql); + + ConnectContext connectContext = Mockito.mock(ConnectContext.class); + StatementContext statementContext = new StatementContext(); + statementContext.setShortCircuitQuery(true); + PrepareCommand prepareCommand = new PrepareCommand( + "stmt", logicalPlan, Collections.emptyList(), new OriginStatement(sql, 0)); + PreparedStatementContext preparedStatement = new PreparedStatementContext( + prepareCommand, connectContext, statementContext, "stmt"); + ShortCircuitQueryContext cachedPlan = Mockito.mock(ShortCircuitQueryContext.class); + Mockito.when(cachedPlan.isReusable(connectContext)).thenReturn(false); + preparedStatement.shortCircuitQueryContext = Optional.of(cachedPlan); + + StmtExecutor executor = Mockito.mock(StmtExecutor.class); + Mockito.when(connectContext.getPreparedStementContext("stmt")).thenReturn(preparedStatement); + Mockito.when(connectContext.getSessionVariable()).thenReturn(new SessionVariable()); + Mockito.when(connectContext.getStatementContext()).thenReturn(statementContext); + Mockito.when(executor.getContext()).thenReturn(connectContext); + + new ExecuteCommand("stmt", prepareCommand, statementContext).run(connectContext, executor); + + Mockito.verify(cachedPlan).isReusable(connectContext); + Mockito.verify(executor).execute(); + Mockito.verify(executor, Mockito.never()).executeAndSendResult(Mockito.anyBoolean(), Mockito.anyBoolean(), + Mockito.any(), Mockito.any(), Mockito.any(), Mockito.any()); + Assertions.assertNotSame(prepareCommand, preparedStatement.command, + "rejecting a stale fast-path plan must rebuild the retained prepared command"); + } + private String resolveNextSnapshot(TableScanParams scanParams, AtomicInteger snapshotId) { return scanParams.getOrResolveMapParams(ignored -> ImmutableMap.of( "scan.snapshot-id", String.valueOf(snapshotId.incrementAndGet()))) diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java index ddc0c3d4c92671..4f31a52c5b4e6d 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java @@ -19,6 +19,7 @@ import org.apache.doris.analysis.DescriptorTable; import org.apache.doris.analysis.Queriable; +import org.apache.doris.analysis.UserIdentity; import org.apache.doris.catalog.Column; import org.apache.doris.catalog.KeysType; import org.apache.doris.catalog.MaterializedIndex; @@ -27,6 +28,7 @@ import org.apache.doris.catalog.PrimitiveType; import org.apache.doris.catalog.RandomDistributionInfo; import org.apache.doris.catalog.SinglePartitionInfo; +import org.apache.doris.nereids.SecurityDependencyContext; import org.apache.doris.nereids.StatementContext; import org.apache.doris.nereids.trees.expressions.Placeholder; import org.apache.doris.nereids.trees.expressions.SlotReference; @@ -110,6 +112,18 @@ public void testReusableRequiresSamePartitionTopologyVersion() { Assertions.assertFalse(context.isReusable(connectContext(-1))); } + @Test + public void testReusableRequiresCurrentSecurityDependencies() { + ConnectContext connectContext = connectContext(-1); + SecurityDependencyContext securityDependencyContext = Mockito.mock(SecurityDependencyContext.class); + Mockito.when(securityDependencyContext.isValid(connectContext)).thenReturn(false); + ShortCircuitQueryContext context = new ShortCircuitQueryContext( + table("tbl", 10), "tbl", 10, -1, securityDependencyContext); + + Assertions.assertFalse(context.isReusable(connectContext)); + Mockito.verify(securityDependencyContext).isValid(connectContext); + } + @Test public void testSerializedQueryOptionsKeepBitmapOpCountVersion() throws Exception { TQueryOptions queryOptions = new SessionVariable().toThrift(); @@ -127,7 +141,8 @@ public void testSerializedQueryOptionsKeepBitmapOpCountVersion() throws Exceptio Mockito.when(planner.getScanNodes()).thenReturn(Collections.singletonList(scanNode)); ShortCircuitQueryContext context = - new ShortCircuitQueryContext(planner, Mockito.mock(Queriable.class)); + new ShortCircuitQueryContext(planner, Mockito.mock(Queriable.class), + new SecurityDependencyContext(UserIdentity.ROOT)); TQueryOptions serializedQueryOptions = new TQueryOptions(); new TDeserializer().deserialize(serializedQueryOptions, context.serializedQueryOptions.toByteArray()); diff --git a/regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out b/regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out new file mode 100644 index 00000000000000..0b85864c0774b3 --- /dev/null +++ b/regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out @@ -0,0 +1,15 @@ +-- This file is automatically generated. You should know what you did if you want to edit this +-- !before_policy -- +2 20 restricted + +-- !cached_before_policy -- +2 20 restricted + +-- !after_policy_added -- + +-- !after_policy_dropped -- +2 20 restricted + +-- !after_select_granted -- +2 20 restricted + diff --git a/regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy b/regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy new file mode 100644 index 00000000000000..86ccf8ccdabbbf --- /dev/null +++ b/regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy @@ -0,0 +1,131 @@ +// 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 java.sql.DriverManager +import java.sql.SQLException + +suite("prepared_short_circuit_security_refresh", "nonConcurrent") { + def dbName = context.config.getDbNameByFile(context.file) + def policyName = "prepared_short_circuit_security_refresh_policy" + def testUser = "prepared_short_circuit_security_refresh_user" + def testPassword = "PreparedSecurity@123" + def adminUser = context.config.jdbcUser + def adminPassword = context.config.jdbcPassword + String serverPrepareUrl = getServerPrepareJdbcUrl(context.config.jdbcUrl, dbName) + + sql "DROP TABLE IF EXISTS prepared_short_circuit_security_refresh_tbl" + sql "DROP USER IF EXISTS ${testUser}" + sql "CREATE USER ${testUser} IDENTIFIED BY '${testPassword}'" + sql """ + CREATE TABLE prepared_short_circuit_security_refresh_tbl ( + k INT NOT NULL, + tenant_id INT NOT NULL, + payload VARCHAR(32) NULL + ) ENGINE=OLAP + UNIQUE KEY(k) + DISTRIBUTED BY HASH(k) BUCKETS 1 + PROPERTIES ( + "replication_num" = "1", + "enable_unique_key_merge_on_write" = "true", + "store_row_column" = "true" + ) + """ + sql """DROP ROW POLICY IF EXISTS ${policyName} + ON ${dbName}.prepared_short_circuit_security_refresh_tbl FOR ${testUser}""" + sql """INSERT INTO prepared_short_circuit_security_refresh_tbl + VALUES (1, 10, 'allowed'), (2, 20, 'restricted')""" + sql "GRANT SELECT_PRIV ON ${dbName}.prepared_short_circuit_security_refresh_tbl TO ${testUser}" + sql "SET GLOBAL enable_server_side_prepared_statement = true" + sql "SYNC" + + if (isCloudMode()) { + def clusters = sql "SHOW CLUSTERS" + assertTrue(!clusters.isEmpty()) + sql "GRANT USAGE_PRIV ON CLUSTER `${clusters[0][0]}` TO ${testUser}" + } + + def adminConnection = DriverManager.getConnection(context.config.jdbcUrl, adminUser, adminPassword) + def adminExecute = { String statement -> + adminConnection.createStatement().withCloseable { adminStatement -> + adminStatement.execute(statement) + } + } + + try { + explain { + sql """ + SELECT /*+ SET_VAR(enable_nereids_planner=true, + enable_fallback_to_original_planner=false, + enable_short_circuit_query=true) */ + k, tenant_id, payload + FROM prepared_short_circuit_security_refresh_tbl + WHERE k = 2 + """ + contains "SHORT-CIRCUIT" + } + connect(testUser, testPassword, serverPrepareUrl) { + sql "SET enable_fallback_to_original_planner = false" + def prepared = prepareStatement( + """SELECT /*+ SET_VAR(enable_nereids_planner=true, + enable_fallback_to_original_planner=false, + enable_short_circuit_query=true) */ + k, tenant_id, payload + FROM prepared_short_circuit_security_refresh_tbl + WHERE k = ?""") + assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) + prepared.setInt(1, 2) + + // The second execution must use the cached point-query plan. + qe_before_policy prepared + qe_cached_before_policy prepared + + adminExecute(""" + CREATE ROW POLICY ${policyName} + ON ${dbName}.prepared_short_circuit_security_refresh_tbl + AS RESTRICTIVE TO ${testUser} USING (tenant_id = 10) + """) + qe_after_policy_added prepared + + adminExecute("""DROP ROW POLICY ${policyName} + ON ${dbName}.prepared_short_circuit_security_refresh_tbl FOR ${testUser}""") + qe_after_policy_dropped prepared + + adminExecute("""REVOKE SELECT_PRIV + ON ${dbName}.prepared_short_circuit_security_refresh_tbl FROM ${testUser}""") + boolean denied = false + try { + prepared.executeQuery().close() + } catch (SQLException e) { + denied = true + logger.info("prepared execution was denied after SELECT revoke: ${e.message}") + } + assertTrue(denied, "the cached point-query plan must not survive SELECT revocation") + + adminExecute("""GRANT SELECT_PRIV + ON ${dbName}.prepared_short_circuit_security_refresh_tbl TO ${testUser}""") + qe_after_select_granted prepared + prepared.close() + } + } finally { + adminExecute("""DROP ROW POLICY IF EXISTS ${policyName} + ON ${dbName}.prepared_short_circuit_security_refresh_tbl FOR ${testUser}""") + adminExecute("""GRANT SELECT_PRIV + ON ${dbName}.prepared_short_circuit_security_refresh_tbl TO ${testUser}""") + adminConnection.close() + sql "DROP USER IF EXISTS ${testUser}" + } +} From 6aa418cbc18c18970175df13f6503c5ad8044577 Mon Sep 17 00:00:00 2001 From: morrySnow Date: Mon, 14 Sep 2026 18:59:02 +0800 Subject: [PATCH 3/6] [fix](point query) Version cached built-in security decisions Avoid catalog, privilege, row-policy, and mask lookups on every prepared point-query execution when Doris built-in authorization is unchanged. External authorization and LDAP continue to use fail-closed full validation. --- .../apache/doris/mysql/privilege/Auth.java | 30 ++++ .../nereids/SecurityDependencyContext.java | 145 +++++++++++++++--- .../doris/nereids/StatementContext.java | 3 +- .../org/apache/doris/policy/PolicyMgr.java | 13 +- .../doris/mysql/privilege/AuthTest.java | 13 ++ .../SecurityDependencyContextTest.java | 76 +++++++-- .../org/apache/doris/policy/PolicyTest.java | 14 ++ .../qe/ShortCircuitQueryContextTest.java | 3 +- 8 files changed, 254 insertions(+), 43 deletions(-) diff --git a/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java b/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java index e2a1182bd228ac..c1c353e499c659 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java +++ b/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java @@ -118,6 +118,10 @@ public class Auth implements Writable { private PasswordPolicyManager passwdPolicyManager = new PasswordPolicyManager(); + // Prepared point-query plans use this process-local epoch to avoid repeating built-in + // privilege checks while no authorization state has changed. + private transient volatile long authorizationVersion; + private ReentrantReadWriteLock lock = new ReentrantReadWriteLock(); private void readLock() { @@ -136,6 +140,19 @@ private void writeUnlock() { lock.writeLock().unlock(); } + private void markAuthorizationChanged() { + authorizationVersion++; + } + + public long getAuthorizationVersion() { + return authorizationVersion; + } + + /** Whether this version covers every role which can affect authorization decisions. */ + public boolean isAuthorizationVersionReliable() { + return !isLdapAuthEnabled(); + } + public enum PrivLevel { GLOBAL, CATALOG, DATABASE, TABLE, RESOURCE, WORKLOAD_GROUP, CLUSTER, STAGE, STORAGE_VAULT } @@ -587,6 +604,7 @@ private void createUserInternal(UserIdentity userIdent, String roleName, byte[] if (role != null) { userRoleManager.addUserRole(userIdent, roleName); } + markAuthorizationChanged(); // other user properties propertyMgr.addUserResource(userIdent.getQualifiedUser()); MetricRepo.updateUserConnectionMaxMetric(this, userIdent.getQualifiedUser(), @@ -641,6 +659,7 @@ private void dropUserInternal(UserIdentity userIdent, boolean ignoreIfNonExists, roleManager.removeDefaultRole(userIdent); // drop user role userRoleManager.dropUser(userIdent); + markAuthorizationChanged(); passwdPolicyManager.dropUser(userIdent); userManager.removeUser(userIdent); if (CollectionUtils.isEmpty(userManager.getUserByName(userIdent.getQualifiedUser()))) { @@ -755,6 +774,7 @@ private void grantInternal(UserIdentity userIdent, String role, TablePattern tbl } Role newRole = new Role(role, tblPattern, privs, colPrivileges); roleManager.addOrMergeRole(newRole, false /* err on exist */); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, tblPattern, privs, null, role, colPrivileges); Env.getCurrentEnv().getEditLog().logGrantPriv(info); @@ -806,6 +826,7 @@ private void grantInternal(UserIdentity userIdent, String role, ResourcePattern Role newRole = new Role(role, resourcePattern, privs); roleManager.addOrMergeRole(newRole, false /* err on exist */); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, resourcePattern, privs, null, role); @@ -837,6 +858,7 @@ private void grantInternal(UserIdentity userIdent, String role, WorkloadGroupPat Role newRole = new Role(role, workloadGroupPattern, privs); roleManager.addOrMergeRole(newRole, false /* err on exist */); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, workloadGroupPattern, privs, null, role); @@ -862,6 +884,7 @@ private void grantInternal(UserIdentity userIdent, List roles, boolean i } } userRoleManager.addUserRoles(userIdent, roles); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, roles); Env.getCurrentEnv().getEditLog().logGrantPriv(info); @@ -952,6 +975,7 @@ private void revokeInternal(UserIdentity userIdent, String role, TablePattern tb } // revoke privs from role roleManager.revokePrivs(role, tblPattern, privs, colPrivileges, errOnNonExist); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, tblPattern, privs, null, role, colPrivileges); @@ -973,6 +997,7 @@ private void revokeInternal(UserIdentity userIdent, String role, ResourcePattern // revoke privs from role roleManager.revokePrivs(role, resourcePattern, privs, errOnNonExist); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, resourcePattern, privs, null, role); @@ -994,6 +1019,7 @@ private void revokeInternal(UserIdentity userIdent, String role, WorkloadGroupPa // revoke privs from role roleManager.revokePrivs(role, workloadGroupPattern, privs, errOnNonExist); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, workloadGroupPattern, privs, null, role); @@ -1019,6 +1045,7 @@ private void revokeInternal(UserIdentity userIdent, List roles, boolean } } userRoleManager.removeUserRoles(userIdent, roles); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, roles); Env.getCurrentEnv().getEditLog().logRevokePriv(info); @@ -1134,6 +1161,7 @@ private void createRoleInternal(String role, boolean ignoreIfExists, String comm } roleManager.addOrMergeRole(emptyPrivsRole, true /* err on exist */); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(role, comment); @@ -1167,6 +1195,7 @@ private void dropRoleInternal(String role, boolean ignoreIfNonExists, boolean is roleManager.dropRole(role, true /* err on non exist */); userRoleManager.dropRole(role); + markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(null, null, null, role, null, null, ""); Env.getCurrentEnv().getEditLog().logDropRole(info); @@ -1996,6 +2025,7 @@ private void setRoleToUser(UserIdentity userIdent, String role) throws DdlExcept userRoleManager.dropUser(userIdent); userRoleManager.addUserRole(userIdent, role); userRoleManager.addUserRole(userIdent, roleManager.getUserDefaultRoleName(userIdent)); + markAuthorizationChanged(); } private void updateUserTlsRequirements(UserIdentity userIdent) throws DdlException { diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java index 5ec32baa64f13d..52cc9f23b35dba 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java @@ -25,17 +25,21 @@ import org.apache.doris.catalog.TableIf; import org.apache.doris.common.UserException; import org.apache.doris.datasource.CatalogIf; +import org.apache.doris.datasource.InternalCatalog; +import org.apache.doris.mysql.privilege.Auth; +import org.apache.doris.mysql.privilege.InternalAuthorizationPlugin; import org.apache.doris.nereids.SqlCacheContext.FullColumnName; import org.apache.doris.nereids.SqlCacheContext.FullTableName; import org.apache.doris.nereids.rules.analysis.UserAuthentication; +import org.apache.doris.policy.PolicyMgr; import org.apache.doris.qe.ConnectContext; +import org.apache.doris.qe.SessionVariable; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableSet; import com.google.common.collect.Maps; import org.apache.commons.collections4.CollectionUtils; -import java.util.LinkedHashMap; import java.util.LinkedHashSet; import java.util.List; import java.util.Locale; @@ -53,16 +57,39 @@ * built before that policy existed. */ public class SecurityDependencyContext { - private final UserIdentity userIdentity; + private static final long UNKNOWN_VERSION = -1; + + private final Env planningEnv; + private final long authorizationVersion; + private final long rowPolicyVersion; + private final boolean versionValidationEligible; private final Map> checkedPrivileges = Maps.newLinkedHashMap(); private final Map> rowPolicies = Maps.newLinkedHashMap(); private final Map> dataMaskPolicies = Maps.newLinkedHashMap(); - private boolean complete; + private final Map> dataMaskColumnsByTable = Maps.newLinkedHashMap(); + private boolean useVersionValidation; + private boolean complete = true; + + /** Create a context which always uses full security revalidation. */ + public SecurityDependencyContext() { + this(null, UNKNOWN_VERSION, UNKNOWN_VERSION, false); + } - /** SecurityDependencyContext */ - public SecurityDependencyContext(UserIdentity userIdentity) { - this.userIdentity = userIdentity; - this.complete = userIdentity != null; + /** Create a context and capture the security versions before analysis starts. */ + public SecurityDependencyContext(ConnectContext connectContext) { + this(connectContext == null ? null : connectContext.getEnv(), usesAuthorizationChecks(connectContext)); + } + + private SecurityDependencyContext(Env env, boolean versionValidationEligible) { + this(env, currentAuthorizationVersion(env), currentRowPolicyVersion(env), versionValidationEligible); + } + + private SecurityDependencyContext(Env planningEnv, long authorizationVersion, long rowPolicyVersion, + boolean versionValidationEligible) { + this.planningEnv = planningEnv; + this.authorizationVersion = authorizationVersion; + this.rowPolicyVersion = rowPolicyVersion; + this.versionValidationEligible = versionValidationEligible; } /** Record the columns whose SELECT privilege was checked while the plan was analyzed. */ @@ -90,13 +117,16 @@ public synchronized void setRowPolicies( /** Record the mask answer for a column, including the absence of a mask. */ public synchronized void addDataMask( String catalog, String database, String table, String column, Optional mask) { - dataMaskPolicies.put(new FullColumnName( - catalog, database, table, column.toLowerCase(Locale.ROOT)), mask); + String normalizedColumn = column.toLowerCase(Locale.ROOT); + FullTableName tableName = new FullTableName(catalog, database, table); + dataMaskPolicies.put(new FullColumnName(catalog, database, table, normalizedColumn), mask); + dataMaskColumnsByTable.computeIfAbsent(tableName, ignored -> new LinkedHashSet<>()).add(normalizedColumn); } /** Freeze the decisions used by a completed plan before storing them in a reusable context. */ public synchronized SecurityDependencyContext snapshot() { - SecurityDependencyContext snapshot = new SecurityDependencyContext(userIdentity); + SecurityDependencyContext snapshot = new SecurityDependencyContext( + planningEnv, authorizationVersion, rowPolicyVersion, versionValidationEligible); snapshot.complete = complete; for (Map.Entry> entry : checkedPrivileges.entrySet()) { snapshot.checkedPrivileges.put(entry.getKey(), ImmutableSet.copyOf(entry.getValue())); @@ -105,6 +135,10 @@ public synchronized SecurityDependencyContext snapshot() { snapshot.rowPolicies.put(entry.getKey(), ImmutableList.copyOf(entry.getValue())); } snapshot.dataMaskPolicies.putAll(dataMaskPolicies); + for (Map.Entry> entry : dataMaskColumnsByTable.entrySet()) { + snapshot.dataMaskColumnsByTable.put(entry.getKey(), ImmutableSet.copyOf(entry.getValue())); + } + snapshot.useVersionValidation = snapshot.canUseVersionValidation(); return snapshot; } @@ -124,15 +158,19 @@ public synchronized SecurityDependencyContext snapshotForShortCircuit() { * planning path performs the authoritative checks and returns the usual user-facing error when access was * revoked. Authorization-source failures also reject reuse, so this fast path always fails closed. */ - public synchronized boolean isValid(ConnectContext connectContext) { + public boolean isValid(ConnectContext connectContext) { if (!complete || connectContext == null) { return false; } try { - if (!Objects.equals(userIdentity, connectContext.getCurrentUserIdentity())) { + Env env = connectContext.getEnv(); + if (useVersionValidation) { + return usesAuthorizationChecks(connectContext) && versionsAreCurrent(env); + } + UserIdentity currentUser = connectContext.getCurrentUserIdentity(); + if (currentUser == null) { return false; } - Env env = connectContext.getEnv(); for (Map.Entry> entry : checkedPrivileges.entrySet()) { TableIf table = findTable(env, entry.getKey()); if (table == null) { @@ -143,27 +181,90 @@ public synchronized boolean isValid(ConnectContext connectContext) { for (Map.Entry> entry : rowPolicies.entrySet()) { FullTableName table = entry.getKey(); List current = env.getAccessManager().evalRowFilterPolicies( - userIdentity, table.catalog, table.db, table.table); + currentUser, table.catalog, table.db, table.table); if (!CollectionUtils.isEqualCollection(entry.getValue(), current)) { return false; } } - return dataMasksAreValid(env); + return dataMasksAreValid(env, currentUser); } catch (UserException | RuntimeException e) { return false; } } - private boolean dataMasksAreValid(Env env) { - Map> columnsByTable = new LinkedHashMap<>(); - for (FullColumnName column : dataMaskPolicies.keySet()) { - columnsByTable.computeIfAbsent(new FullTableName(column.catalog, column.db, column.table), - table -> new LinkedHashSet<>()).add(column.column); + private boolean canUseVersionValidation() { + if (!complete || !versionValidationEligible || checkedPrivileges.isEmpty() + || authorizationVersion == UNKNOWN_VERSION || rowPolicyVersion == UNKNOWN_VERSION) { + return false; + } + return allDependenciesUseInternalCatalog() && usesVersionedBuiltInAuthorization(planningEnv); + } + + private boolean versionsAreCurrent(Env env) { + if (env == null || env != planningEnv) { + return false; + } + Auth auth = env.getAuth(); + PolicyMgr policyMgr = env.getPolicyMgr(); + return auth != null && policyMgr != null + && auth.isAuthorizationVersionReliable() + && env.getAccessManager().getAccessControllerOrDefault(InternalCatalog.INTERNAL_CATALOG_NAME) + instanceof InternalAuthorizationPlugin + && auth.getAuthorizationVersion() == authorizationVersion + && policyMgr.getRowPolicyVersion() == rowPolicyVersion; + } + + private static boolean usesVersionedBuiltInAuthorization(Env env) { + if (env == null || env.getAuth() == null || env.getPolicyMgr() == null + || !env.getAuth().isAuthorizationVersionReliable()) { + return false; + } + return env.getAccessManager().getAccessControllerOrDefault(InternalCatalog.INTERNAL_CATALOG_NAME) + instanceof InternalAuthorizationPlugin; + } + + private boolean allDependenciesUseInternalCatalog() { + for (FullTableName table : checkedPrivileges.keySet()) { + if (!InternalCatalog.INTERNAL_CATALOG_NAME.equals(table.catalog)) { + return false; + } } - for (Map.Entry> entry : columnsByTable.entrySet()) { + for (FullTableName table : rowPolicies.keySet()) { + if (!InternalCatalog.INTERNAL_CATALOG_NAME.equals(table.catalog)) { + return false; + } + } + for (FullTableName table : dataMaskColumnsByTable.keySet()) { + if (!InternalCatalog.INTERNAL_CATALOG_NAME.equals(table.catalog)) { + return false; + } + } + return true; + } + + private static long currentAuthorizationVersion(Env env) { + Auth auth = env == null ? null : env.getAuth(); + return auth == null ? UNKNOWN_VERSION : auth.getAuthorizationVersion(); + } + + private static long currentRowPolicyVersion(Env env) { + PolicyMgr policyMgr = env == null ? null : env.getPolicyMgr(); + return policyMgr == null ? UNKNOWN_VERSION : policyMgr.getRowPolicyVersion(); + } + + private static boolean usesAuthorizationChecks(ConnectContext connectContext) { + if (connectContext == null || connectContext.isSkipAuth()) { + return false; + } + SessionVariable sessionVariable = connectContext.getSessionVariable(); + return sessionVariable != null && !sessionVariable.isPlayNereidsDump(); + } + + private boolean dataMasksAreValid(Env env, UserIdentity currentUser) { + for (Map.Entry> entry : dataMaskColumnsByTable.entrySet()) { FullTableName table = entry.getKey(); Map current = env.getAccessManager().evalDataMaskPolicies( - userIdentity, table.catalog, table.db, table.table, entry.getValue()); + currentUser, table.catalog, table.db, table.table, entry.getValue()); for (String column : entry.getValue()) { Optional currentMask = Optional.ofNullable( current.get(column.toLowerCase(Locale.ROOT))); diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java index 32a121914ba46f..7b966e1e77ec26 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java @@ -394,8 +394,7 @@ private StatementContext(ConnectContext connectContext, OriginStatement originSt this.connectContext = connectContext; this.originStatement = originStatement; exprIdGenerator = ExprId.createGenerator(initialId); - this.securityDependencyContext = new SecurityDependencyContext( - connectContext == null ? null : connectContext.getCurrentUserIdentity()); + this.securityDependencyContext = new SecurityDependencyContext(connectContext); if (connectContext != null && connectContext.getSessionVariable() != null) { if (CacheAnalyzer.canUseSqlCache(connectContext.getSessionVariable())) { // cannot set the queryId here because the queryId for the current query is set diff --git a/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java b/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java index 0ecb66853d28e2..4e4f104606a1a2 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java +++ b/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java @@ -73,6 +73,13 @@ public class PolicyMgr implements Writable { // ctlName -> dbName -> tableName -> List private Map>>> tablePolicies = Maps.newConcurrentMap(); + // Process-local epoch used to invalidate prepared point-query plans after row-policy changes. + private transient volatile long rowPolicyVersion; + + public long getRowPolicyVersion() { + return rowPolicyVersion; + } + private void writeLock() { lock.writeLock().lock(); } @@ -301,6 +308,7 @@ private void unprotectedAdd(Policy policy) { typeToPolicyMap.put(policy.getType(), dbPolicies); if (PolicyTypeEnum.ROW == policy.getType()) { addTablePolicies((RowPolicy) policy); + rowPolicyVersion++; } } @@ -336,7 +344,7 @@ public void replayStoragePolicyAlter(StoragePolicy log) { private void unprotectedDrop(DropPolicyLog log) { List policies = getPoliciesByType(log.getType()); - policies.removeIf(policy -> { + boolean removed = policies.removeIf(policy -> { if (policy.matchPolicy(log)) { if (policy instanceof StoragePolicy) { ((StoragePolicy) policy).removeResourceReference(); @@ -352,6 +360,9 @@ private void unprotectedDrop(DropPolicyLog log) { return false; }); typeToPolicyMap.put(log.getType(), policies); + if (removed && log.getType() == PolicyTypeEnum.ROW) { + rowPolicyVersion++; + } } public List getUserPolicies(String ctlName, String dbName, String tableName, UserIdentity user) { diff --git a/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java b/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java index 993973a905d8f3..c5abcf7e70c229 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java @@ -87,4 +87,17 @@ public void testCheckDbPrivWithSessionMappedRoleForTempUser() throws Exception { } } + @Test + public void testAuthorizationVersionAdvancesOnMutation() throws Exception { + Auth auth = Env.getCurrentEnv().getAuth(); + long version = auth.getAuthorizationVersion(); + + addUser("authorization_version_user", true); + + Assertions.assertTrue(auth.getAuthorizationVersion() > version); + long changedVersion = auth.getAuthorizationVersion(); + auth.refreshUserPrivEntriesByResovledIPs(Collections.emptyMap()); + Assertions.assertEquals(changedVersion, auth.getAuthorizationVersion()); + } + } diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java index 42e628b5bcf0a7..20163bd89f4620 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java @@ -27,7 +27,10 @@ import org.apache.doris.datasource.CatalogIf; import org.apache.doris.datasource.CatalogMgr; import org.apache.doris.mysql.privilege.AccessControllerManager; +import org.apache.doris.mysql.privilege.Auth; +import org.apache.doris.mysql.privilege.InternalAuthorizationPlugin; import org.apache.doris.mysql.privilege.PrivPredicate; +import org.apache.doris.policy.PolicyMgr; import org.apache.doris.qe.ConnectContext; import org.apache.doris.qe.SessionVariable; @@ -52,7 +55,7 @@ public class SecurityDependencyContextTest { public void testUnchangedPoliciesAreValid() { RowFilterSpec rowFilter = RowFilterSpec.restrictive("row:1", "tenant_id = 1"); DataMaskSpec dataMask = new DataMaskSpec("mask:1", "null"); - SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + SecurityDependencyContext dependencies = new SecurityDependencyContext(); dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of(rowFilter)); dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.of(dataMask)); @@ -65,7 +68,7 @@ public void testUnchangedPoliciesAreValid() { @Test public void testAddedRowPolicyInvalidatesNegativeSnapshot() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + SecurityDependencyContext dependencies = new SecurityDependencyContext(); dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of()); ConnectContext connectContext = contextWithPolicies( ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1")), ImmutableMap.of()); @@ -75,7 +78,7 @@ public void testAddedRowPolicyInvalidatesNegativeSnapshot() { @Test public void testChangedRowPolicyInvalidatesSnapshot() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + SecurityDependencyContext dependencies = new SecurityDependencyContext(); dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1"))); ConnectContext connectContext = contextWithPolicies( @@ -86,7 +89,7 @@ public void testChangedRowPolicyInvalidatesSnapshot() { @Test public void testAddedDataMaskInvalidatesNegativeSnapshot() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + SecurityDependencyContext dependencies = new SecurityDependencyContext(); dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.empty()); ConnectContext connectContext = contextWithPolicies(ImmutableList.of(), ImmutableMap.of(COLUMN, new DataMaskSpec("mask:1", "null"))); @@ -95,31 +98,72 @@ public void testAddedDataMaskInvalidatesNegativeSnapshot() { } @Test - public void testDifferentExecutingIdentityInvalidatesSnapshot() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + public void testMissingCurrentIdentityFailsClosed() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(); ConnectContext connectContext = Mockito.mock(ConnectContext.class); - Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn( - UserIdentity.createAnalyzedUserIdentWithIp("other", "%")); Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); } @Test - public void testMissingPlanningIdentityFailsClosed() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(null); + public void testMissingPrivilegeRecordingDisablesShortCircuitReuse() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(); ConnectContext connectContext = Mockito.mock(ConnectContext.class); - Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(UserIdentity.ROOT); + Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); - Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(connectContext)); } @Test - public void testMissingPrivilegeRecordingDisablesShortCircuitReuse() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + @SuppressWarnings({"rawtypes", "unchecked"}) + public void testBuiltInVersionsSkipDependencyRevalidation() throws Exception { + CatalogIf catalog = Mockito.mock(CatalogIf.class); + DatabaseIf database = Mockito.mock(DatabaseIf.class); + TableIf table = Mockito.mock(TableIf.class); + Auth auth = Mockito.mock(Auth.class); + PolicyMgr policyMgr = Mockito.mock(PolicyMgr.class); + AccessControllerManager accessManager = Mockito.mock(AccessControllerManager.class); + Env env = Mockito.mock(Env.class); ConnectContext connectContext = Mockito.mock(ConnectContext.class); + + Mockito.when(catalog.getName()).thenReturn(CATALOG); + Mockito.when(database.getCatalog()).thenReturn(catalog); + Mockito.when(database.getFullName()).thenReturn(DATABASE); + Mockito.when(table.getDatabase()).thenReturn(database); + Mockito.when(table.getName()).thenReturn(TABLE); + Mockito.when(auth.getAuthorizationVersion()).thenReturn(7L); + Mockito.when(auth.isAuthorizationVersionReliable()).thenReturn(true); + Mockito.when(policyMgr.getRowPolicyVersion()).thenReturn(11L); + Mockito.when(accessManager.getAccessControllerOrDefault(CATALOG)) + .thenReturn(new InternalAuthorizationPlugin(auth)); + Mockito.when(env.getAuth()).thenReturn(auth); + Mockito.when(env.getPolicyMgr()).thenReturn(policyMgr); + Mockito.when(env.getAccessManager()).thenReturn(accessManager); Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); + Mockito.when(connectContext.getEnv()).thenReturn(env); + Mockito.when(connectContext.getSessionVariable()).thenReturn(new SessionVariable()); - Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(connectContext)); + SecurityDependencyContext dependencies = new SecurityDependencyContext(connectContext); + dependencies.addCheckedPrivilege(table, ImmutableSet.of(COLUMN)); + dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of()); + dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.empty()); + SecurityDependencyContext snapshot = dependencies.snapshotForShortCircuit(); + + Assertions.assertTrue(snapshot.isValid(connectContext)); + Mockito.verify(accessManager, Mockito.never()).checkColumnsPriv( + connectContext, CATALOG, DATABASE, TABLE, ImmutableSet.of(COLUMN), PrivPredicate.SELECT); + Mockito.verify(accessManager, Mockito.never()).evalRowFilterPolicies( + ArgumentMatchers.any(), ArgumentMatchers.anyString(), ArgumentMatchers.anyString(), + ArgumentMatchers.anyString()); + Mockito.verify(accessManager, Mockito.never()).evalDataMaskPolicies( + ArgumentMatchers.any(), ArgumentMatchers.anyString(), ArgumentMatchers.anyString(), + ArgumentMatchers.anyString(), ArgumentMatchers.anySet()); + + Mockito.when(auth.getAuthorizationVersion()).thenReturn(8L); + Assertions.assertFalse(snapshot.isValid(connectContext)); + Mockito.when(auth.getAuthorizationVersion()).thenReturn(7L); + Mockito.when(policyMgr.getRowPolicyVersion()).thenReturn(12L); + Assertions.assertFalse(snapshot.isValid(connectContext)); } @Test @@ -147,7 +191,7 @@ public void testPrivilegeRevocationOrAuthorizationFailureInvalidatesSnapshot() t Mockito.when(connectContext.getEnv()).thenReturn(env); Mockito.when(connectContext.getSessionVariable()).thenReturn(new SessionVariable()); - SecurityDependencyContext dependencies = new SecurityDependencyContext(USER); + SecurityDependencyContext dependencies = new SecurityDependencyContext(); dependencies.addCheckedPrivilege(table, ImmutableSet.of(COLUMN)); SecurityDependencyContext snapshot = dependencies.snapshot(); Assertions.assertTrue(snapshot.isValid(connectContext)); diff --git a/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java b/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java index 33759675ecf192..27fdab19e50836 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java @@ -188,6 +188,20 @@ public void testDropPolicy() throws Exception { () -> dropPolicy("DROP ROW POLICY test_row_policy1 ON test.table1")); } + @Test + public void testRowPolicyVersionAdvancesOnMutation() throws Exception { + PolicyMgr policyMgr = Env.getCurrentEnv().getPolicyMgr(); + long version = policyMgr.getRowPolicyVersion(); + + createPolicy("CREATE ROW POLICY test_row_policy_version ON test.table1 AS PERMISSIVE" + + " TO test_policy USING (k1 = 1)"); + long createdVersion = policyMgr.getRowPolicyVersion(); + Assertions.assertTrue(createdVersion > version); + + dropPolicy("DROP ROW POLICY test_row_policy_version ON test.table1"); + Assertions.assertTrue(policyMgr.getRowPolicyVersion() > createdVersion); + } + @Test public void testMergeFilter() throws Exception { createPolicy("CREATE ROW POLICY test_row_policy1 ON test.table1 AS RESTRICTIVE TO test_policy USING (k1 = 1)"); diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java index 4f31a52c5b4e6d..35ae004acbb8b1 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java @@ -19,7 +19,6 @@ import org.apache.doris.analysis.DescriptorTable; import org.apache.doris.analysis.Queriable; -import org.apache.doris.analysis.UserIdentity; import org.apache.doris.catalog.Column; import org.apache.doris.catalog.KeysType; import org.apache.doris.catalog.MaterializedIndex; @@ -142,7 +141,7 @@ public void testSerializedQueryOptionsKeepBitmapOpCountVersion() throws Exceptio ShortCircuitQueryContext context = new ShortCircuitQueryContext(planner, Mockito.mock(Queriable.class), - new SecurityDependencyContext(UserIdentity.ROOT)); + new SecurityDependencyContext()); TQueryOptions serializedQueryOptions = new TQueryOptions(); new TDeserializer().deserialize(serializedQueryOptions, context.serializedQueryOptions.toByteArray()); From d3ac58ea8a419608c4f6ca29e69b62ba739e3adc Mon Sep 17 00:00:00 2001 From: morrySnow Date: Tue, 15 Sep 2026 11:00:31 +0800 Subject: [PATCH 4/6] [fix](point query) Restrict reusable prepared plans Disable the short-circuit path for row-policy and view plans, and keep cached plans bound to the planning user, authenticated roles, security epochs, and table namespace. Validate expected SELECT denial on COM_CHANGE_USER and privilege refresh. --- .../nereids/SecurityDependencyContext.java | 233 ++++++------------ .../doris/nereids/StatementContext.java | 8 +- .../rules/analysis/ExpressionAnalyzer.java | 5 +- ...calResultSinkToShortCircuitPointQuery.java | 8 +- .../doris/qe/ShortCircuitQueryContext.java | 41 +++ .../SecurityDependencyContextTest.java | 229 ++++++----------- .../rewrite/ShortCircuitPointQueryTest.java | 9 + .../plans/commands/ExecuteCommandTest.java | 9 +- .../qe/ShortCircuitQueryContextTest.java | 38 ++- ...repared_short_circuit_security_refresh.out | 4 +- .../prepared_point_query_row_policy.groovy | 34 ++- ...ared_short_circuit_security_refresh.groovy | 45 +++- 12 files changed, 331 insertions(+), 332 deletions(-) diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java index 52cc9f23b35dba..f08d3193aa5ba6 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java @@ -22,182 +22,163 @@ import org.apache.doris.authorization.RowFilterSpec; import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.Env; +import org.apache.doris.catalog.OlapTable; import org.apache.doris.catalog.TableIf; -import org.apache.doris.common.UserException; import org.apache.doris.datasource.CatalogIf; import org.apache.doris.datasource.InternalCatalog; import org.apache.doris.mysql.privilege.Auth; import org.apache.doris.mysql.privilege.InternalAuthorizationPlugin; -import org.apache.doris.nereids.SqlCacheContext.FullColumnName; -import org.apache.doris.nereids.SqlCacheContext.FullTableName; -import org.apache.doris.nereids.rules.analysis.UserAuthentication; import org.apache.doris.policy.PolicyMgr; import org.apache.doris.qe.ConnectContext; import org.apache.doris.qe.SessionVariable; -import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableSet; -import com.google.common.collect.Maps; -import org.apache.commons.collections4.CollectionUtils; -import java.util.LinkedHashSet; import java.util.List; -import java.util.Locale; -import java.util.Map; import java.util.Objects; import java.util.Optional; import java.util.Set; /** - * Security decisions which an analyzed plan depends on. + * Security state which a reusable prepared point-query plan depends on. * - *

Unlike {@link SqlCacheContext}, this context exists independently of the SQL result-cache switch. A prepared - * short-circuit plan can otherwise outlive the privilege and data-policy decisions made while it was analyzed. - * Callers record both positive and negative policy answers so that adding a policy invalidates a plan which was - * built before that policy existed. + *

Row policies are deliberately not copied or compared here. A plan with a row policy is not eligible for + * short-circuit execution, while the process-local policy version invalidates a previously cached no-policy plan + * when a policy is later added. This keeps the common validation path to a few identity and volatile-version reads. + * Authorization sources without reliable local versions are never reused without replanning. */ public class SecurityDependencyContext { private static final long UNKNOWN_VERSION = -1; + private final UserIdentity planningUserIdentity; + private final Set planningAuthenticatedRoles; private final Env planningEnv; private final long authorizationVersion; private final long rowPolicyVersion; private final boolean versionValidationEligible; - private final Map> checkedPrivileges = Maps.newLinkedHashMap(); - private final Map> rowPolicies = Maps.newLinkedHashMap(); - private final Map> dataMaskPolicies = Maps.newLinkedHashMap(); - private final Map> dataMaskColumnsByTable = Maps.newLinkedHashMap(); + private boolean privilegeChecked; + private boolean internalCatalogOnly = true; + private boolean olapTableOnly = true; + private boolean hasRowPolicy; + private boolean hasDataMask; private boolean useVersionValidation; private boolean complete = true; - /** Create a context which always uses full security revalidation. */ + /** Create an incomplete context for tests and callers without a connection. */ public SecurityDependencyContext() { - this(null, UNKNOWN_VERSION, UNKNOWN_VERSION, false); + this(null, ImmutableSet.of(), null, UNKNOWN_VERSION, UNKNOWN_VERSION, false); } - /** Create a context and capture the security versions before analysis starts. */ + /** Capture the effective authorization subject and security versions before analysis starts. */ public SecurityDependencyContext(ConnectContext connectContext) { - this(connectContext == null ? null : connectContext.getEnv(), usesAuthorizationChecks(connectContext)); + this(connectContext == null ? null : connectContext.getCurrentUserIdentity(), + authenticatedRoles(connectContext), + connectContext == null ? null : connectContext.getEnv(), + usesAuthorizationChecks(connectContext)); } - private SecurityDependencyContext(Env env, boolean versionValidationEligible) { - this(env, currentAuthorizationVersion(env), currentRowPolicyVersion(env), versionValidationEligible); + private SecurityDependencyContext(UserIdentity planningUserIdentity, Set planningAuthenticatedRoles, + Env planningEnv, boolean versionValidationEligible) { + this(planningUserIdentity, planningAuthenticatedRoles, planningEnv, + currentAuthorizationVersion(planningEnv), currentRowPolicyVersion(planningEnv), + versionValidationEligible); } - private SecurityDependencyContext(Env planningEnv, long authorizationVersion, long rowPolicyVersion, + private SecurityDependencyContext(UserIdentity planningUserIdentity, Set planningAuthenticatedRoles, + Env planningEnv, long authorizationVersion, long rowPolicyVersion, boolean versionValidationEligible) { + this.planningUserIdentity = planningUserIdentity; + this.planningAuthenticatedRoles = planningAuthenticatedRoles; this.planningEnv = planningEnv; this.authorizationVersion = authorizationVersion; this.rowPolicyVersion = rowPolicyVersion; this.versionValidationEligible = versionValidationEligible; } - /** Record the columns whose SELECT privilege was checked while the plan was analyzed. */ + /** Record that SELECT privileges were checked and whether the relation supports version-only validation. */ public synchronized void addCheckedPrivilege(TableIf table, Set usedColumns) { - Optional tableName = qualifiedName(table); - if (!tableName.isPresent()) { + if (table == null) { complete = false; return; } - Set existing = checkedPrivileges.get(tableName.get()); - if (existing == null) { - checkedPrivileges.put(tableName.get(), ImmutableSet.copyOf(usedColumns)); - } else { - checkedPrivileges.put(tableName.get(), ImmutableSet.builder() - .addAll(existing).addAll(usedColumns).build()); + DatabaseIf database = table.getDatabase(); + CatalogIf catalog = database == null ? null : database.getCatalog(); + if (catalog == null) { + complete = false; + return; } + privilegeChecked = true; + internalCatalogOnly &= InternalCatalog.INTERNAL_CATALOG_NAME.equals(catalog.getName()); + olapTableOnly &= table instanceof OlapTable; } - /** Record the complete row-filter answer, including an empty answer. */ + /** Record only whether a row policy exists; policy objects never enter the point-query cache. */ public synchronized void setRowPolicies( String catalog, String database, String table, List policies) { - rowPolicies.put(new FullTableName(catalog, database, table), ImmutableList.copyOf(policies)); + hasRowPolicy |= policies != null && !policies.isEmpty(); } - /** Record the mask answer for a column, including the absence of a mask. */ + /** A row policy makes the statement ineligible for short-circuit execution. */ + public synchronized boolean hasRowPolicy() { + return hasRowPolicy; + } + + /** Record mask presence so a masked plan is never accepted by the version-only cache path. */ public synchronized void addDataMask( String catalog, String database, String table, String column, Optional mask) { - String normalizedColumn = column.toLowerCase(Locale.ROOT); - FullTableName tableName = new FullTableName(catalog, database, table); - dataMaskPolicies.put(new FullColumnName(catalog, database, table, normalizedColumn), mask); - dataMaskColumnsByTable.computeIfAbsent(tableName, ignored -> new LinkedHashSet<>()).add(normalizedColumn); + hasDataMask |= mask.isPresent(); } /** Freeze the decisions used by a completed plan before storing them in a reusable context. */ public synchronized SecurityDependencyContext snapshot() { SecurityDependencyContext snapshot = new SecurityDependencyContext( - planningEnv, authorizationVersion, rowPolicyVersion, versionValidationEligible); + planningUserIdentity, planningAuthenticatedRoles, planningEnv, + authorizationVersion, rowPolicyVersion, versionValidationEligible); + snapshot.privilegeChecked = privilegeChecked; + snapshot.internalCatalogOnly = internalCatalogOnly; + snapshot.olapTableOnly = olapTableOnly; + snapshot.hasRowPolicy = hasRowPolicy; + snapshot.hasDataMask = hasDataMask; snapshot.complete = complete; - for (Map.Entry> entry : checkedPrivileges.entrySet()) { - snapshot.checkedPrivileges.put(entry.getKey(), ImmutableSet.copyOf(entry.getValue())); - } - for (Map.Entry> entry : rowPolicies.entrySet()) { - snapshot.rowPolicies.put(entry.getKey(), ImmutableList.copyOf(entry.getValue())); - } - snapshot.dataMaskPolicies.putAll(dataMaskPolicies); - for (Map.Entry> entry : dataMaskColumnsByTable.entrySet()) { - snapshot.dataMaskColumnsByTable.put(entry.getKey(), ImmutableSet.copyOf(entry.getValue())); - } snapshot.useVersionValidation = snapshot.canUseVersionValidation(); return snapshot; } - /** Freeze the decisions for a prepared short-circuit plan, failing closed if authorization was not recorded. */ + /** Freeze a prepared short-circuit dependency set, failing closed if its proof is incomplete. */ public synchronized SecurityDependencyContext snapshotForShortCircuit() { SecurityDependencyContext snapshot = snapshot(); - if (checkedPrivileges.isEmpty()) { + if (!privilegeChecked || hasRowPolicy) { snapshot.complete = false; + snapshot.useVersionValidation = false; } return snapshot; } /** - * Revalidate every security decision before a cached plan bypasses analysis. + * Check whether an analyzed point-query plan can bypass planning again. * - *

A false result does not deny the statement itself. It rejects only the cached plan, after which the normal - * planning path performs the authoritative checks and returns the usual user-facing error when access was - * revoked. Authorization-source failures also reject reuse, so this fast path always fails closed. + *

A false result rejects only cached reuse. The prepared statement is reparsed and analyzed normally, so + * authorization failures retain their standard user-facing error. The authorization subject is checked before + * the version shortcut because COM_CHANGE_USER keeps the connection's prepared statements alive. */ public boolean isValid(ConnectContext connectContext) { - if (!complete || connectContext == null) { + if (!complete || !useVersionValidation || connectContext == null + || !Objects.equals(planningUserIdentity, connectContext.getCurrentUserIdentity()) + || !planningAuthenticatedRoles.equals(authenticatedRoles(connectContext))) { return false; } try { - Env env = connectContext.getEnv(); - if (useVersionValidation) { - return usesAuthorizationChecks(connectContext) && versionsAreCurrent(env); - } - UserIdentity currentUser = connectContext.getCurrentUserIdentity(); - if (currentUser == null) { - return false; - } - for (Map.Entry> entry : checkedPrivileges.entrySet()) { - TableIf table = findTable(env, entry.getKey()); - if (table == null) { - return false; - } - UserAuthentication.checkPermission(table, connectContext, entry.getValue()); - } - for (Map.Entry> entry : rowPolicies.entrySet()) { - FullTableName table = entry.getKey(); - List current = env.getAccessManager().evalRowFilterPolicies( - currentUser, table.catalog, table.db, table.table); - if (!CollectionUtils.isEqualCollection(entry.getValue(), current)) { - return false; - } - } - return dataMasksAreValid(env, currentUser); - } catch (UserException | RuntimeException e) { + return usesAuthorizationChecks(connectContext) && versionsAreCurrent(connectContext.getEnv()); + } catch (RuntimeException e) { return false; } } private boolean canUseVersionValidation() { - if (!complete || !versionValidationEligible || checkedPrivileges.isEmpty() - || authorizationVersion == UNKNOWN_VERSION || rowPolicyVersion == UNKNOWN_VERSION) { - return false; - } - return allDependenciesUseInternalCatalog() && usesVersionedBuiltInAuthorization(planningEnv); + return complete && versionValidationEligible && privilegeChecked && internalCatalogOnly && olapTableOnly + && !hasRowPolicy && !hasDataMask + && authorizationVersion != UNKNOWN_VERSION && rowPolicyVersion != UNKNOWN_VERSION + && usesVersionedBuiltInAuthorization(planningEnv); } private boolean versionsAreCurrent(Env env) { @@ -215,31 +196,10 @@ private boolean versionsAreCurrent(Env env) { } private static boolean usesVersionedBuiltInAuthorization(Env env) { - if (env == null || env.getAuth() == null || env.getPolicyMgr() == null - || !env.getAuth().isAuthorizationVersionReliable()) { - return false; - } - return env.getAccessManager().getAccessControllerOrDefault(InternalCatalog.INTERNAL_CATALOG_NAME) - instanceof InternalAuthorizationPlugin; - } - - private boolean allDependenciesUseInternalCatalog() { - for (FullTableName table : checkedPrivileges.keySet()) { - if (!InternalCatalog.INTERNAL_CATALOG_NAME.equals(table.catalog)) { - return false; - } - } - for (FullTableName table : rowPolicies.keySet()) { - if (!InternalCatalog.INTERNAL_CATALOG_NAME.equals(table.catalog)) { - return false; - } - } - for (FullTableName table : dataMaskColumnsByTable.keySet()) { - if (!InternalCatalog.INTERNAL_CATALOG_NAME.equals(table.catalog)) { - return false; - } - } - return true; + return env != null && env.getAuth() != null && env.getPolicyMgr() != null + && env.getAuth().isAuthorizationVersionReliable() + && env.getAccessManager().getAccessControllerOrDefault(InternalCatalog.INTERNAL_CATALOG_NAME) + instanceof InternalAuthorizationPlugin; } private static long currentAuthorizationVersion(Env env) { @@ -260,44 +220,11 @@ private static boolean usesAuthorizationChecks(ConnectContext connectContext) { return sessionVariable != null && !sessionVariable.isPlayNereidsDump(); } - private boolean dataMasksAreValid(Env env, UserIdentity currentUser) { - for (Map.Entry> entry : dataMaskColumnsByTable.entrySet()) { - FullTableName table = entry.getKey(); - Map current = env.getAccessManager().evalDataMaskPolicies( - currentUser, table.catalog, table.db, table.table, entry.getValue()); - for (String column : entry.getValue()) { - Optional currentMask = Optional.ofNullable( - current.get(column.toLowerCase(Locale.ROOT))); - if (!Objects.equals(dataMaskPolicies.get( - new FullColumnName(table.catalog, table.db, table.table, column)), currentMask)) { - return false; - } - } - } - return true; - } - - private Optional qualifiedName(TableIf table) { - if (table == null) { - return Optional.empty(); - } - DatabaseIf database = table.getDatabase(); - if (database == null || database.getCatalog() == null) { - return Optional.empty(); - } - return Optional.of(new FullTableName( - database.getCatalog().getName(), database.getFullName(), table.getName())); - } - - private TableIf findTable(Env env, FullTableName fullTableName) { - CatalogIf> catalog = env.getCatalogMgr().getCatalog(fullTableName.catalog); - if (catalog == null) { - return null; - } - Optional> database = catalog.getDb(fullTableName.db); - if (!database.isPresent()) { - return null; + private static Set authenticatedRoles(ConnectContext connectContext) { + if (connectContext == null) { + return ImmutableSet.of(); } - return database.get().getTable(fullTableName.table).orElse(null); + Set roles = connectContext.getAuthenticatedRoles(); + return roles == null || roles.isEmpty() ? ImmutableSet.of() : ImmutableSet.copyOf(roles); } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java index 7b966e1e77ec26..7bf7ab6f3e0c1f 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java @@ -194,10 +194,10 @@ public enum TableFrom { // Map placeholder id to the physical key slot used by the immutable point-query template. private final Map idToComparisonSlot = new TreeMap<>(); - // Equality literals that were written as constants in the statement or injected by a - // security policy. They are deliberately separate from placeholder bindings: a prepared - // point query must never replace a fixed predicate merely because it references the same - // column as a placeholder. + // Equality literals written as constants in the statement. They are deliberately separate + // from placeholder bindings: a prepared point query must never replace a fixed predicate + // merely because it references the same column as a placeholder. Plans with row policies + // are not eligible for the point-query shortcut. private final List pointQueryFixedKeyConstraints = new ArrayList<>(); private boolean pointQueryFixedKeyConstraintsComplete = true; 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 f8c9950eefa81d..0e6891933b47cb 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 @@ -953,8 +953,9 @@ public Expression visitComparisonPredicate(ComparisonPredicate cp, ExpressionRew /** * Keep fixed equality values distinct from prepared-statement placeholders. The point-query * executor used to rediscover both from the translated scan conjuncts and then update every - * predicate sharing a column name. Once a row policy adds {@code key = constant}, that loses - * provenance and turns the policy constant into caller-controlled state. + * predicate sharing a column name. A statement containing both {@code key = ?} and + * {@code key = constant} would therefore lose provenance and turn the fixed value into + * caller-controlled state. */ private void registerPointQueryFixedKeyConstraint(ComparisonPredicate original, Expression analyzed, ExpressionRewriteContext context) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java index 3a651ad2001b53..ff92869fadecb7 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java @@ -103,7 +103,13 @@ boolean scanMatchShortCircuitCondition(LogicalOlapScan olapScan) { // set short circuit flag and return the original plan private Plan shortCircuit(Plan root, OlapTable olapTable, Set conjuncts, StatementContext statementContext) { - if (!statementContext.arePointQueryFixedKeyConstraintsComplete()) { + // Row filters are injected into the analyzed plan and views are inlined. Neither shape has a + // cheap, stable dependency fence suitable for a reusable direct plan, so keep both on the + // normal execution path. A global row-policy epoch still invalidates a no-policy plan if a + // policy is added after it was cached. + if (statementContext.getSecurityDependencyContext().hasRowPolicy() + || !statementContext.getViewDdlSqls().isEmpty() + || !statementContext.arePointQueryFixedKeyConstraintsComplete()) { return root; } // All key columns in conjuncts diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java index e991aaa21c9ed9..755f72dd5eca24 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java @@ -24,8 +24,10 @@ import org.apache.doris.analysis.LiteralExprUtils; import org.apache.doris.analysis.Queriable; import org.apache.doris.catalog.Column; +import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.OlapTable; import org.apache.doris.catalog.Type; +import org.apache.doris.datasource.CatalogIf; import org.apache.doris.nereids.NereidsPlanner; import org.apache.doris.nereids.SecurityDependencyContext; import org.apache.doris.nereids.StatementContext; @@ -81,6 +83,7 @@ public class ShortCircuitQueryContext { public final String tableName; private final long fileCacheQueryLimitBytes; private final long partitionTopologyVersion; + private final TableNamespaceSnapshot tableNamespaceSnapshot; private final SecurityDependencyContext securityDependencyContext; public final OlapScanNode scanNode; @@ -157,6 +160,7 @@ private ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.tableName = this.scanNode.getTableNameInPlan(); this.schemaVersion = this.tbl.getBaseSchemaVersion(); this.partitionTopologyVersion = this.tbl.getPartitionTopologyVersion(); + this.tableNamespaceSnapshot = TableNamespaceSnapshot.from(this.tbl); this.analzyedQuery = analzyedQuery; this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(this.scanNode, statementContext); this.securityDependencyContext = securityDependencyContext == null @@ -182,6 +186,7 @@ private ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.schemaVersion = schemaVersion; this.fileCacheQueryLimitBytes = fileCacheQueryLimitBytes; this.partitionTopologyVersion = tbl.getPartitionTopologyVersion(); + this.tableNamespaceSnapshot = TableNamespaceSnapshot.from(tbl); this.scanNode = null; this.analzyedQuery = null; this.pointQueryKeyTemplate = PointQueryKeyTemplate.unsupported(); @@ -201,6 +206,7 @@ private ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.schemaVersion = tbl.getBaseSchemaVersion(); this.fileCacheQueryLimitBytes = -1; this.partitionTopologyVersion = tbl.getPartitionTopologyVersion(); + this.tableNamespaceSnapshot = TableNamespaceSnapshot.from(tbl); this.analzyedQuery = null; this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(scanNode, statementContext); this.securityDependencyContext = null; @@ -212,9 +218,44 @@ public boolean isReusable(ConnectContext ctx) { && Objects.equals(this.tableName, this.tbl.getName()) && this.fileCacheQueryLimitBytes == ctx.getSessionVariable().fileCacheQueryLimitBytes && this.tbl.getPartitionTopologyVersion() == this.partitionTopologyVersion + && this.tableNamespaceSnapshot.matches(this.tbl) && (securityDependencyContext == null || securityDependencyContext.isValid(ctx)); } + /** Fence name-scoped grants when a catalog or database is renamed or replaced. */ + private static class TableNamespaceSnapshot { + private final DatabaseIf database; + private final CatalogIf catalog; + private final String databaseName; + private final String catalogName; + + private TableNamespaceSnapshot(DatabaseIf database, CatalogIf catalog, + String databaseName, String catalogName) { + this.database = database; + this.catalog = catalog; + this.databaseName = databaseName; + this.catalogName = catalogName; + } + + private static TableNamespaceSnapshot from(OlapTable table) { + DatabaseIf database = table.getDatabase(); + CatalogIf catalog = database == null ? null : database.getCatalog(); + return new TableNamespaceSnapshot(database, catalog, + database == null ? null : database.getFullName(), + catalog == null ? null : catalog.getName()); + } + + private boolean matches(OlapTable table) { + DatabaseIf currentDatabase = table.getDatabase(); + CatalogIf currentCatalog = currentDatabase == null ? null : currentDatabase.getCatalog(); + return currentDatabase == database + && currentCatalog == catalog + && Objects.equals(databaseName, + currentDatabase == null ? null : currentDatabase.getFullName()) + && Objects.equals(catalogName, currentCatalog == null ? null : currentCatalog.getName()); + } + } + public void sanitize() { Preconditions.checkNotNull(serializedDescTable); Preconditions.checkNotNull(serializedOutputExpr); diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java index 20163bd89f4620..2c1b65778667e2 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java @@ -22,198 +22,133 @@ import org.apache.doris.authorization.RowFilterSpec; import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.Env; -import org.apache.doris.catalog.TableIf; -import org.apache.doris.common.UserException; +import org.apache.doris.catalog.OlapTable; import org.apache.doris.datasource.CatalogIf; -import org.apache.doris.datasource.CatalogMgr; import org.apache.doris.mysql.privilege.AccessControllerManager; import org.apache.doris.mysql.privilege.Auth; import org.apache.doris.mysql.privilege.InternalAuthorizationPlugin; -import org.apache.doris.mysql.privilege.PrivPredicate; import org.apache.doris.policy.PolicyMgr; import org.apache.doris.qe.ConnectContext; import org.apache.doris.qe.SessionVariable; import com.google.common.collect.ImmutableList; -import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; -import org.mockito.ArgumentMatchers; import org.mockito.Mockito; import java.util.Optional; public class SecurityDependencyContextTest { private static final UserIdentity USER = UserIdentity.createAnalyzedUserIdentWithIp("reader", "%"); + private static final UserIdentity OTHER_USER = UserIdentity.createAnalyzedUserIdentWithIp("other", "%"); private static final String CATALOG = "internal"; private static final String DATABASE = "db"; private static final String TABLE = "tbl"; private static final String COLUMN = "value"; @Test - public void testUnchangedPoliciesAreValid() { - RowFilterSpec rowFilter = RowFilterSpec.restrictive("row:1", "tenant_id = 1"); - DataMaskSpec dataMask = new DataMaskSpec("mask:1", "null"); - SecurityDependencyContext dependencies = new SecurityDependencyContext(); - dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of(rowFilter)); - dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.of(dataMask)); - - ConnectContext connectContext = contextWithPolicies( - ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1")), - ImmutableMap.of(COLUMN, new DataMaskSpec("mask:1", "null"))); - - Assertions.assertTrue(dependencies.snapshot().isValid(connectContext)); + public void testBuiltInVersionsAllowConstantTimeReuse() { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext snapshot = fixture.completeDependencies().snapshotForShortCircuit(); + + Assertions.assertTrue(snapshot.isValid(fixture.connectContext)); + Mockito.verify(fixture.accessManager, Mockito.never()).evalRowFilterPolicies( + Mockito.any(), Mockito.anyString(), Mockito.anyString(), Mockito.anyString()); + Mockito.verify(fixture.accessManager, Mockito.never()).evalDataMaskPolicies( + Mockito.any(), Mockito.anyString(), Mockito.anyString(), Mockito.anyString(), Mockito.anySet()); + + Mockito.when(fixture.auth.getAuthorizationVersion()).thenReturn(8L); + Assertions.assertFalse(snapshot.isValid(fixture.connectContext)); + Mockito.when(fixture.auth.getAuthorizationVersion()).thenReturn(7L); + Mockito.when(fixture.policyMgr.getRowPolicyVersion()).thenReturn(12L); + Assertions.assertFalse(snapshot.isValid(fixture.connectContext)); } @Test - public void testAddedRowPolicyInvalidatesNegativeSnapshot() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(); - dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of()); - ConnectContext connectContext = contextWithPolicies( - ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1")), ImmutableMap.of()); + public void testDifferentPlanningUserInvalidatesReuse() { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext snapshot = fixture.completeDependencies().snapshotForShortCircuit(); - Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + Mockito.when(fixture.connectContext.getCurrentUserIdentity()).thenReturn(OTHER_USER); + Assertions.assertFalse(snapshot.isValid(fixture.connectContext)); } @Test - public void testChangedRowPolicyInvalidatesSnapshot() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(); - dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, - ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1"))); - ConnectContext connectContext = contextWithPolicies( - ImmutableList.of(RowFilterSpec.restrictive("row:2", "tenant_id = 2")), ImmutableMap.of()); - - Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); - } - - @Test - public void testAddedDataMaskInvalidatesNegativeSnapshot() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(); - dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.empty()); - ConnectContext connectContext = contextWithPolicies(ImmutableList.of(), - ImmutableMap.of(COLUMN, new DataMaskSpec("mask:1", "null"))); + public void testDifferentAuthenticatedRolesInvalidateReuse() { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext snapshot = fixture.completeDependencies().snapshotForShortCircuit(); - Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + Mockito.when(fixture.connectContext.getAuthenticatedRoles()).thenReturn(ImmutableSet.of("auditor")); + Assertions.assertFalse(snapshot.isValid(fixture.connectContext)); } @Test - public void testMissingCurrentIdentityFailsClosed() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(); - ConnectContext connectContext = Mockito.mock(ConnectContext.class); + public void testRowPolicyDisablesShortCircuitReuse() { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext dependencies = fixture.completeDependencies(); + dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, + ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1"))); - Assertions.assertFalse(dependencies.snapshot().isValid(connectContext)); + Assertions.assertTrue(dependencies.hasRowPolicy()); + Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(fixture.connectContext)); } @Test - public void testMissingPrivilegeRecordingDisablesShortCircuitReuse() { - SecurityDependencyContext dependencies = new SecurityDependencyContext(); - ConnectContext connectContext = Mockito.mock(ConnectContext.class); - Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); + public void testDataMaskDisablesVersionOnlyReuse() { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext dependencies = fixture.completeDependencies(); + dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, + Optional.of(new DataMaskSpec("mask:1", "null"))); - Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(connectContext)); + Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(fixture.connectContext)); } @Test - @SuppressWarnings({"rawtypes", "unchecked"}) - public void testBuiltInVersionsSkipDependencyRevalidation() throws Exception { - CatalogIf catalog = Mockito.mock(CatalogIf.class); - DatabaseIf database = Mockito.mock(DatabaseIf.class); - TableIf table = Mockito.mock(TableIf.class); - Auth auth = Mockito.mock(Auth.class); - PolicyMgr policyMgr = Mockito.mock(PolicyMgr.class); - AccessControllerManager accessManager = Mockito.mock(AccessControllerManager.class); - Env env = Mockito.mock(Env.class); - ConnectContext connectContext = Mockito.mock(ConnectContext.class); - - Mockito.when(catalog.getName()).thenReturn(CATALOG); - Mockito.when(database.getCatalog()).thenReturn(catalog); - Mockito.when(database.getFullName()).thenReturn(DATABASE); - Mockito.when(table.getDatabase()).thenReturn(database); - Mockito.when(table.getName()).thenReturn(TABLE); - Mockito.when(auth.getAuthorizationVersion()).thenReturn(7L); - Mockito.when(auth.isAuthorizationVersionReliable()).thenReturn(true); - Mockito.when(policyMgr.getRowPolicyVersion()).thenReturn(11L); - Mockito.when(accessManager.getAccessControllerOrDefault(CATALOG)) - .thenReturn(new InternalAuthorizationPlugin(auth)); - Mockito.when(env.getAuth()).thenReturn(auth); - Mockito.when(env.getPolicyMgr()).thenReturn(policyMgr); - Mockito.when(env.getAccessManager()).thenReturn(accessManager); - Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); - Mockito.when(connectContext.getEnv()).thenReturn(env); - Mockito.when(connectContext.getSessionVariable()).thenReturn(new SessionVariable()); - - SecurityDependencyContext dependencies = new SecurityDependencyContext(connectContext); - dependencies.addCheckedPrivilege(table, ImmutableSet.of(COLUMN)); - dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of()); - dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.empty()); - SecurityDependencyContext snapshot = dependencies.snapshotForShortCircuit(); - - Assertions.assertTrue(snapshot.isValid(connectContext)); - Mockito.verify(accessManager, Mockito.never()).checkColumnsPriv( - connectContext, CATALOG, DATABASE, TABLE, ImmutableSet.of(COLUMN), PrivPredicate.SELECT); - Mockito.verify(accessManager, Mockito.never()).evalRowFilterPolicies( - ArgumentMatchers.any(), ArgumentMatchers.anyString(), ArgumentMatchers.anyString(), - ArgumentMatchers.anyString()); - Mockito.verify(accessManager, Mockito.never()).evalDataMaskPolicies( - ArgumentMatchers.any(), ArgumentMatchers.anyString(), ArgumentMatchers.anyString(), - ArgumentMatchers.anyString(), ArgumentMatchers.anySet()); - - Mockito.when(auth.getAuthorizationVersion()).thenReturn(8L); - Assertions.assertFalse(snapshot.isValid(connectContext)); - Mockito.when(auth.getAuthorizationVersion()).thenReturn(7L); - Mockito.when(policyMgr.getRowPolicyVersion()).thenReturn(12L); - Assertions.assertFalse(snapshot.isValid(connectContext)); - } + public void testMissingPrivilegeRecordingFailsClosed() { + BuiltInFixture fixture = new BuiltInFixture(); - @Test - @SuppressWarnings({"rawtypes", "unchecked"}) - public void testPrivilegeRevocationOrAuthorizationFailureInvalidatesSnapshot() throws Exception { - CatalogIf catalog = Mockito.mock(CatalogIf.class); - DatabaseIf database = Mockito.mock(DatabaseIf.class); - TableIf table = Mockito.mock(TableIf.class); - CatalogMgr catalogMgr = Mockito.mock(CatalogMgr.class); - AccessControllerManager accessManager = Mockito.mock(AccessControllerManager.class); - Env env = Mockito.mock(Env.class); - ConnectContext connectContext = Mockito.mock(ConnectContext.class); - - Mockito.when(catalog.getName()).thenReturn(CATALOG); - Mockito.when(catalog.getDb(DATABASE)).thenReturn(Optional.of(database)); - Mockito.when(database.getCatalog()).thenReturn(catalog); - Mockito.when(database.getFullName()).thenReturn(DATABASE); - Mockito.when(database.getTable(TABLE)).thenReturn(Optional.of(table)); - Mockito.when(table.getDatabase()).thenReturn(database); - Mockito.when(table.getName()).thenReturn(TABLE); - Mockito.when(catalogMgr.getCatalog(CATALOG)).thenReturn(catalog); - Mockito.when(env.getCatalogMgr()).thenReturn(catalogMgr); - Mockito.when(env.getAccessManager()).thenReturn(accessManager); - Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); - Mockito.when(connectContext.getEnv()).thenReturn(env); - Mockito.when(connectContext.getSessionVariable()).thenReturn(new SessionVariable()); - - SecurityDependencyContext dependencies = new SecurityDependencyContext(); - dependencies.addCheckedPrivilege(table, ImmutableSet.of(COLUMN)); - SecurityDependencyContext snapshot = dependencies.snapshot(); - Assertions.assertTrue(snapshot.isValid(connectContext)); - - Mockito.doThrow(new UserException("SELECT was revoked or authorization is unavailable")) - .when(accessManager).checkColumnsPriv( - connectContext, CATALOG, DATABASE, TABLE, ImmutableSet.of(COLUMN), PrivPredicate.SELECT); - Assertions.assertFalse(snapshot.isValid(connectContext)); + Assertions.assertFalse(new SecurityDependencyContext(fixture.connectContext) + .snapshotForShortCircuit().isValid(fixture.connectContext)); } - private ConnectContext contextWithPolicies( - ImmutableList rowFilters, ImmutableMap dataMasks) { - AccessControllerManager accessManager = Mockito.mock(AccessControllerManager.class); - Mockito.when(accessManager.evalRowFilterPolicies(USER, CATALOG, DATABASE, TABLE)).thenReturn(rowFilters); - Mockito.when(accessManager.evalDataMaskPolicies( - ArgumentMatchers.eq(USER), ArgumentMatchers.eq(CATALOG), ArgumentMatchers.eq(DATABASE), - ArgumentMatchers.eq(TABLE), ArgumentMatchers.anySet())).thenReturn(dataMasks); - Env env = Mockito.mock(Env.class); - Mockito.when(env.getAccessManager()).thenReturn(accessManager); - ConnectContext connectContext = Mockito.mock(ConnectContext.class); - Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); - Mockito.when(connectContext.getEnv()).thenReturn(env); - return connectContext; + private static class BuiltInFixture { + private final CatalogIf catalog = Mockito.mock(CatalogIf.class); + private final DatabaseIf database = Mockito.mock(DatabaseIf.class); + private final OlapTable table = Mockito.mock(OlapTable.class); + private final Auth auth = Mockito.mock(Auth.class); + private final PolicyMgr policyMgr = Mockito.mock(PolicyMgr.class); + private final AccessControllerManager accessManager = Mockito.mock(AccessControllerManager.class); + private final Env env = Mockito.mock(Env.class); + private final ConnectContext connectContext = Mockito.mock(ConnectContext.class); + + @SuppressWarnings({"rawtypes", "unchecked"}) + private BuiltInFixture() { + Mockito.when(catalog.getName()).thenReturn(CATALOG); + Mockito.when(database.getCatalog()).thenReturn((CatalogIf) catalog); + Mockito.when(database.getFullName()).thenReturn(DATABASE); + Mockito.when(table.getDatabase()).thenReturn((DatabaseIf) database); + Mockito.when(table.getName()).thenReturn(TABLE); + Mockito.when(auth.getAuthorizationVersion()).thenReturn(7L); + Mockito.when(auth.isAuthorizationVersionReliable()).thenReturn(true); + Mockito.when(policyMgr.getRowPolicyVersion()).thenReturn(11L); + Mockito.when(accessManager.getAccessControllerOrDefault(CATALOG)) + .thenReturn(new InternalAuthorizationPlugin(auth)); + Mockito.when(env.getAuth()).thenReturn(auth); + Mockito.when(env.getPolicyMgr()).thenReturn(policyMgr); + Mockito.when(env.getAccessManager()).thenReturn(accessManager); + Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); + Mockito.when(connectContext.getAuthenticatedRoles()).thenReturn(ImmutableSet.of("reader_role")); + Mockito.when(connectContext.getEnv()).thenReturn(env); + Mockito.when(connectContext.getSessionVariable()).thenReturn(new SessionVariable()); + } + + private SecurityDependencyContext completeDependencies() { + SecurityDependencyContext dependencies = new SecurityDependencyContext(connectContext); + dependencies.addCheckedPrivilege(table, ImmutableSet.of(COLUMN)); + dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of()); + dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.empty()); + return dependencies; + } } } diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java index 9511e8d4bb9473..cb5e839c56405f 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java @@ -78,6 +78,7 @@ protected void runBeforeAll() throws Exception { + " \"light_schema_change\" = \"true\",\n" + " \"store_row_column\" = \"true\"\n" + ");"); + createView("CREATE VIEW `view_point_query` AS SELECT `key`, `v1` FROM `tbl_point_query`"); } @Test @@ -141,6 +142,14 @@ void testPointQueryWithManualTabletDoesNotUseShortCircuit() throws Exception { Assertions.assertFalse(connectContext.getStatementContext().isShortCircuitQuery()); } + @Test + void testViewDoesNotUseShortCircuit() { + rewrite("select * from view_point_query where `key` = 1"); + + Assertions.assertFalse(connectContext.getStatementContext().isShortCircuitQuery()); + Assertions.assertFalse(connectContext.getStatementContext().getViewDdlSqls().isEmpty()); + } + @Test void testRemoteOlapTableDoesNotUseShortCircuit() throws Exception { Database database = Env.getCurrentInternalCatalog().getDbOrMetaException("test"); diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java index 892815b5a89280..07a90354ec21ef 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java @@ -266,7 +266,9 @@ public void testFastPathInstallsCachedShortCircuitContextAcrossExecutions() thro PreparedStatementContext preparedStatement = new PreparedStatementContext( prepareCommand, connectContext, statementContext, "stmt"); - // A real ShortCircuitQueryContext (built from a mocked planner) that passes isReusable(). + // Keep the real point-key binding path, but explicitly model a cache whose security + // dependencies have already been validated. A bare StatementContext intentionally + // cannot produce a reusable security snapshot because production validation is fail-closed. Planner planner = Mockito.mock(Planner.class); Mockito.when(planner.getQueryOptions()).thenReturn(new TQueryOptions()); DescriptorTable descriptorTable = new DescriptorTable(); @@ -282,8 +284,9 @@ public void testFastPathInstallsCachedShortCircuitContextAcrossExecutions() thro Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); Mockito.when(scanNode.getConjuncts()).thenReturn(Collections.emptyList()); Mockito.when(planner.getScanNodes()).thenReturn(Collections.singletonList(scanNode)); - ShortCircuitQueryContext cachedPlan = new ShortCircuitQueryContext( - planner, Mockito.mock(Queriable.class), statementContext); + ShortCircuitQueryContext cachedPlan = Mockito.spy(new ShortCircuitQueryContext( + planner, Mockito.mock(Queriable.class), statementContext)); + Mockito.doReturn(true).when(cachedPlan).isReusable(connectContext); preparedStatement.shortCircuitQueryContext = Optional.of(cachedPlan); StmtExecutor executor = Mockito.mock(StmtExecutor.class); diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java index 35ae004acbb8b1..491558a4ff13de 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java @@ -20,6 +20,7 @@ import org.apache.doris.analysis.DescriptorTable; import org.apache.doris.analysis.Queriable; import org.apache.doris.catalog.Column; +import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.KeysType; import org.apache.doris.catalog.MaterializedIndex; import org.apache.doris.catalog.OlapTable; @@ -27,6 +28,7 @@ import org.apache.doris.catalog.PrimitiveType; import org.apache.doris.catalog.RandomDistributionInfo; import org.apache.doris.catalog.SinglePartitionInfo; +import org.apache.doris.datasource.CatalogIf; import org.apache.doris.nereids.SecurityDependencyContext; import org.apache.doris.nereids.StatementContext; import org.apache.doris.nereids.trees.expressions.Placeholder; @@ -49,6 +51,7 @@ import java.math.BigDecimal; import java.util.Collections; import java.util.List; +import java.util.concurrent.atomic.AtomicReference; public class ShortCircuitQueryContextTest { private OlapTable table(String name, int schemaVersion) { @@ -91,6 +94,24 @@ public void testReusableStillChecksTableMetadata() { Assertions.assertFalse(context.isReusable(connectContext(0))); } + @Test + @SuppressWarnings({"rawtypes", "unchecked"}) + public void testReusableRequiresSameDatabaseNamespace() { + CatalogIf catalog = Mockito.mock(CatalogIf.class); + DatabaseIf database = Mockito.mock(DatabaseIf.class); + AtomicReference databaseName = new AtomicReference<>("old_db"); + Mockito.when(catalog.getName()).thenReturn("internal"); + Mockito.when(database.getCatalog()).thenReturn(catalog); + Mockito.when(database.getFullName()).thenAnswer(ignored -> databaseName.get()); + OlapTable table = table("tbl", 10); + Mockito.doReturn(database).when(table).getDatabase(); + ShortCircuitQueryContext context = new ShortCircuitQueryContext(table, "tbl", 10, -1); + + Assertions.assertTrue(context.isReusable(connectContext(-1))); + databaseName.set("new_db"); + Assertions.assertFalse(context.isReusable(connectContext(-1))); + } + @Test public void testReusableRequiresSamePartitionTopologyVersion() { long baseIndexId = 2L; @@ -153,23 +174,22 @@ public void testSerializedQueryOptionsKeepBitmapOpCountVersion() throws Exceptio public void testPreparedKeyTemplateKeepsFixedConstraintsAcrossExecutions() { Column parameterKey = new Column("parameter_key", PrimitiveType.INT); parameterKey.setIsKey(true); - Column policyKey = new Column("policy_key", PrimitiveType.INT); - policyKey.setIsKey(true); - List schema = List.of(parameterKey, policyKey); + Column fixedKey = new Column("fixed_key", PrimitiveType.INT); + fixedKey.setIsKey(true); + List schema = List.of(parameterKey, fixedKey); OlapTable table = pointQueryTable(schema); SlotReference parameterSlot = SlotReference.fromColumn( StatementScopeIdGenerator.newExprId(), table, parameterKey, Collections.emptyList()); - SlotReference policySlot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, policyKey, Collections.emptyList()); + SlotReference fixedSlot = SlotReference.fromColumn( + StatementScopeIdGenerator.newExprId(), table, fixedKey, Collections.emptyList()); PlaceholderId placeholderId = new PlaceholderId(0); StatementContext templateContext = new StatementContext(); templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); templateContext.getIdToComparisonSlot().put(placeholderId, parameterSlot); - // This models a restrictive policy on the same key as the placeholder, plus a - // policy-fixed column in a composite key. + // Fixed statement predicates remain distinct from caller-controlled placeholders. templateContext.addPointQueryFixedKeyConstraint(parameterSlot, new IntegerLiteral(1)); - templateContext.addPointQueryFixedKeyConstraint(policySlot, new IntegerLiteral(9)); + templateContext.addPointQueryFixedKeyConstraint(fixedSlot, new IntegerLiteral(9)); OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); Mockito.when(scanNode.getOlapTable()).thenReturn(table); @@ -182,7 +202,7 @@ public void testPreparedKeyTemplateKeepsFixedConstraintsAcrossExecutions() { Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, firstExecution.getDecision()); Assertions.assertEquals("1", firstExecution.getKeyValues().get("parameter_key").getStringValue()); - Assertions.assertEquals("9", firstExecution.getKeyValues().get("policy_key").getStringValue()); + Assertions.assertEquals("9", firstExecution.getKeyValues().get("fixed_key").getStringValue()); StatementContext second = execution(placeholderId, new IntegerLiteral(2)); Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.EMPTY, diff --git a/regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out b/regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out index 0b85864c0774b3..76fc30b7f8ffe8 100644 --- a/regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out +++ b/regression-test/data/prepared_stmt_p0/prepared_short_circuit_security_refresh.out @@ -5,6 +5,9 @@ -- !cached_before_policy -- 2 20 restricted +-- !after_change_user_restored -- +2 20 restricted + -- !after_policy_added -- -- !after_policy_dropped -- @@ -12,4 +15,3 @@ -- !after_select_granted -- 2 20 restricted - diff --git a/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy index 3b7b4d7ad88e0c..edb024979cee27 100644 --- a/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy +++ b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy @@ -65,9 +65,10 @@ suite("prepared_point_query_row_policy", "p0") { EXPLAIN SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value FROM prepared_point_query_row_policy WHERE tenant_id = 1 AND item_id = 10 """ - assertTrue(explainRows.toString().contains("SHORT-CIRCUIT")) + assertFalse(explainRows.toString().contains("SHORT-CIRCUIT")) } + // A statement governed by a row policy must stay on the normal planning path. connect(user, password, url) { def prepared = prepareStatement """ SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value @@ -75,6 +76,33 @@ suite("prepared_point_query_row_policy", "p0") { WHERE tenant_id = ? AND item_id = ? """ assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) + prepared.setInt(1, 2) + prepared.setInt(2, 10) + prepared.executeQuery().withCloseable { result -> + assertFalse(result.next()) + } + prepared.close() + } + + sql "DROP ROW POLICY IF EXISTS ${policyName} ON ${dbName}.prepared_point_query_row_policy FOR ${user}" + sql "SYNC" + connect(user, password, explainUrl) { + def explainRows = sql """ + EXPLAIN SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + FROM prepared_point_query_row_policy WHERE tenant_id = 1 AND item_id = 10 + """ + assertTrue(explainRows.toString().contains("SHORT-CIRCUIT")) + } + + // Without a policy, a fixed statement predicate can share a key with a placeholder. The + // immutable key template must preserve that fixed value across every execution. + connect(user, password, url) { + def prepared = prepareStatement """ + SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + FROM prepared_point_query_row_policy + WHERE tenant_id = ? AND tenant_id = 1 AND item_id = ? + """ + assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) def readRows = { Integer tenant, int item -> if (tenant == null) { @@ -115,9 +143,7 @@ suite("prepared_point_query_row_policy", "p0") { prepared.close() } - // A fixed predicate on a cast key is not an exact physical-key constraint. Keep it on - // the normal path unless a future proof can establish that the cast is lossless/injective. - sql "DROP ROW POLICY IF EXISTS ${policyName} ON ${dbName}.prepared_point_query_row_policy FOR ${user}" + // A cast row policy also keeps the statement on the normal path. sql """ CREATE ROW POLICY ${policyName} ON ${dbName}.prepared_point_query_row_policy AS RESTRICTIVE TO ${user} USING (CAST(tenant_id AS CHAR(1)) = '1') diff --git a/regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy b/regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy index 86ccf8ccdabbbf..834191248403a4 100644 --- a/regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy +++ b/regression-test/suites/prepared_stmt_p0/prepared_short_circuit_security_refresh.groovy @@ -17,19 +17,24 @@ import java.sql.DriverManager import java.sql.SQLException +import java.util.Locale suite("prepared_short_circuit_security_refresh", "nonConcurrent") { def dbName = context.config.getDbNameByFile(context.file) def policyName = "prepared_short_circuit_security_refresh_policy" def testUser = "prepared_short_circuit_security_refresh_user" def testPassword = "PreparedSecurity@123" + def switchedUser = "prepared_short_circuit_security_switched_user" + def switchedPassword = "PreparedSecurity@456" def adminUser = context.config.jdbcUser def adminPassword = context.config.jdbcPassword String serverPrepareUrl = getServerPrepareJdbcUrl(context.config.jdbcUrl, dbName) sql "DROP TABLE IF EXISTS prepared_short_circuit_security_refresh_tbl" sql "DROP USER IF EXISTS ${testUser}" + sql "DROP USER IF EXISTS ${switchedUser}" sql "CREATE USER ${testUser} IDENTIFIED BY '${testPassword}'" + sql "CREATE USER ${switchedUser} IDENTIFIED BY '${switchedPassword}'" sql """ CREATE TABLE prepared_short_circuit_security_refresh_tbl ( k INT NOT NULL, @@ -49,6 +54,9 @@ suite("prepared_short_circuit_security_refresh", "nonConcurrent") { sql """INSERT INTO prepared_short_circuit_security_refresh_tbl VALUES (1, 10, 'allowed'), (2, 20, 'restricted')""" sql "GRANT SELECT_PRIV ON ${dbName}.prepared_short_circuit_security_refresh_tbl TO ${testUser}" + // Connector/J includes the current database in COM_CHANGE_USER. Give the switched user + // enough metadata access to enter it, but deliberately no SELECT privilege on the table. + sql "GRANT SHOW_VIEW_PRIV ON ${dbName}.prepared_short_circuit_security_refresh_tbl TO ${switchedUser}" sql "SET GLOBAL enable_server_side_prepared_statement = true" sql "SYNC" @@ -56,6 +64,7 @@ suite("prepared_short_circuit_security_refresh", "nonConcurrent") { def clusters = sql "SHOW CLUSTERS" assertTrue(!clusters.isEmpty()) sql "GRANT USAGE_PRIV ON CLUSTER `${clusters[0][0]}` TO ${testUser}" + sql "GRANT USAGE_PRIV ON CLUSTER `${clusters[0][0]}` TO ${switchedUser}" } def adminConnection = DriverManager.getConnection(context.config.jdbcUrl, adminUser, adminPassword) @@ -64,6 +73,22 @@ suite("prepared_short_circuit_security_refresh", "nonConcurrent") { adminStatement.execute(statement) } } + def assertSelectDenied = { prepared, String failureMessage -> + boolean denied = false + try { + prepared.executeQuery().close() + } catch (SQLException e) { + String denial = e.message == null ? "" : e.message.toLowerCase(Locale.ROOT) + if (!denial.contains("permission denied") + || !denial.contains("select_priv") + || !denial.contains("prepared_short_circuit_security_refresh_tbl")) { + throw e + } + denied = true + logger.info("prepared execution reached the expected SELECT denial: ${e.message}") + } + assertTrue(denied, failureMessage) + } try { explain { @@ -93,6 +118,15 @@ suite("prepared_short_circuit_security_refresh", "nonConcurrent") { qe_before_policy prepared qe_cached_before_policy prepared + // COM_CHANGE_USER retains the server-prepared handle. The cached plan must remain bound + // to the identity which planned it, not merely to the physical connection. + def jdbcConnection = prepared.getConnection().unwrap(com.mysql.cj.jdbc.JdbcConnection.class) + jdbcConnection.changeUser(switchedUser, switchedPassword) + assertSelectDenied(prepared, + "the cached point-query plan must not survive a successful COM_CHANGE_USER") + jdbcConnection.changeUser(testUser, testPassword) + qe_after_change_user_restored prepared + adminExecute(""" CREATE ROW POLICY ${policyName} ON ${dbName}.prepared_short_circuit_security_refresh_tbl @@ -106,14 +140,8 @@ suite("prepared_short_circuit_security_refresh", "nonConcurrent") { adminExecute("""REVOKE SELECT_PRIV ON ${dbName}.prepared_short_circuit_security_refresh_tbl FROM ${testUser}""") - boolean denied = false - try { - prepared.executeQuery().close() - } catch (SQLException e) { - denied = true - logger.info("prepared execution was denied after SELECT revoke: ${e.message}") - } - assertTrue(denied, "the cached point-query plan must not survive SELECT revocation") + assertSelectDenied(prepared, + "the cached point-query plan must not survive SELECT revocation") adminExecute("""GRANT SELECT_PRIV ON ${dbName}.prepared_short_circuit_security_refresh_tbl TO ${testUser}""") @@ -127,5 +155,6 @@ suite("prepared_short_circuit_security_refresh", "nonConcurrent") { ON ${dbName}.prepared_short_circuit_security_refresh_tbl TO ${testUser}""") adminConnection.close() sql "DROP USER IF EXISTS ${testUser}" + sql "DROP USER IF EXISTS ${switchedUser}" } } From e5c5fbb0ca2abb85ce051ce5994495296875208f Mon Sep 17 00:00:00 2001 From: morrySnow Date: Tue, 15 Sep 2026 12:27:14 +0800 Subject: [PATCH 5/6] [fix](point query) Simplify prepared-plan security checks --- .../apache/doris/mysql/privilege/Auth.java | 30 -- .../nereids/SecurityDependencyContext.java | 211 +++++------ .../doris/nereids/StatementContext.java | 71 +--- .../rules/analysis/ExpressionAnalyzer.java | 41 +-- ...calResultSinkToShortCircuitPointQuery.java | 51 ++- .../trees/plans/commands/ExecuteCommand.java | 27 +- .../plans/logical/LogicalCheckPolicy.java | 2 +- .../apache/doris/planner/OlapScanNode.java | 12 - .../org/apache/doris/planner/ScanNode.java | 23 -- .../org/apache/doris/policy/PolicyMgr.java | 30 +- .../apache/doris/qe/PointQueryExecutor.java | 89 +++-- .../doris/qe/ShortCircuitQueryContext.java | 338 +----------------- .../org/apache/doris/qe/StmtExecutor.java | 28 +- .../doris/mysql/privilege/AuthTest.java | 13 - .../SecurityDependencyContextTest.java | 63 +++- .../rewrite/ShortCircuitPointQueryTest.java | 27 ++ .../plans/commands/ExecuteCommandTest.java | 45 +-- .../org/apache/doris/policy/PolicyTest.java | 13 +- .../doris/qe/PointQueryExecutorTest.java | 19 - .../qe/ShortCircuitQueryContextTest.java | 228 ++---------- .../prepared_point_query_row_policy.groovy | 91 ++--- 21 files changed, 389 insertions(+), 1063 deletions(-) diff --git a/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java b/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java index c1c353e499c659..e2a1182bd228ac 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java +++ b/fe/fe-core/src/main/java/org/apache/doris/mysql/privilege/Auth.java @@ -118,10 +118,6 @@ public class Auth implements Writable { private PasswordPolicyManager passwdPolicyManager = new PasswordPolicyManager(); - // Prepared point-query plans use this process-local epoch to avoid repeating built-in - // privilege checks while no authorization state has changed. - private transient volatile long authorizationVersion; - private ReentrantReadWriteLock lock = new ReentrantReadWriteLock(); private void readLock() { @@ -140,19 +136,6 @@ private void writeUnlock() { lock.writeLock().unlock(); } - private void markAuthorizationChanged() { - authorizationVersion++; - } - - public long getAuthorizationVersion() { - return authorizationVersion; - } - - /** Whether this version covers every role which can affect authorization decisions. */ - public boolean isAuthorizationVersionReliable() { - return !isLdapAuthEnabled(); - } - public enum PrivLevel { GLOBAL, CATALOG, DATABASE, TABLE, RESOURCE, WORKLOAD_GROUP, CLUSTER, STAGE, STORAGE_VAULT } @@ -604,7 +587,6 @@ private void createUserInternal(UserIdentity userIdent, String roleName, byte[] if (role != null) { userRoleManager.addUserRole(userIdent, roleName); } - markAuthorizationChanged(); // other user properties propertyMgr.addUserResource(userIdent.getQualifiedUser()); MetricRepo.updateUserConnectionMaxMetric(this, userIdent.getQualifiedUser(), @@ -659,7 +641,6 @@ private void dropUserInternal(UserIdentity userIdent, boolean ignoreIfNonExists, roleManager.removeDefaultRole(userIdent); // drop user role userRoleManager.dropUser(userIdent); - markAuthorizationChanged(); passwdPolicyManager.dropUser(userIdent); userManager.removeUser(userIdent); if (CollectionUtils.isEmpty(userManager.getUserByName(userIdent.getQualifiedUser()))) { @@ -774,7 +755,6 @@ private void grantInternal(UserIdentity userIdent, String role, TablePattern tbl } Role newRole = new Role(role, tblPattern, privs, colPrivileges); roleManager.addOrMergeRole(newRole, false /* err on exist */); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, tblPattern, privs, null, role, colPrivileges); Env.getCurrentEnv().getEditLog().logGrantPriv(info); @@ -826,7 +806,6 @@ private void grantInternal(UserIdentity userIdent, String role, ResourcePattern Role newRole = new Role(role, resourcePattern, privs); roleManager.addOrMergeRole(newRole, false /* err on exist */); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, resourcePattern, privs, null, role); @@ -858,7 +837,6 @@ private void grantInternal(UserIdentity userIdent, String role, WorkloadGroupPat Role newRole = new Role(role, workloadGroupPattern, privs); roleManager.addOrMergeRole(newRole, false /* err on exist */); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, workloadGroupPattern, privs, null, role); @@ -884,7 +862,6 @@ private void grantInternal(UserIdentity userIdent, List roles, boolean i } } userRoleManager.addUserRoles(userIdent, roles); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, roles); Env.getCurrentEnv().getEditLog().logGrantPriv(info); @@ -975,7 +952,6 @@ private void revokeInternal(UserIdentity userIdent, String role, TablePattern tb } // revoke privs from role roleManager.revokePrivs(role, tblPattern, privs, colPrivileges, errOnNonExist); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, tblPattern, privs, null, role, colPrivileges); @@ -997,7 +973,6 @@ private void revokeInternal(UserIdentity userIdent, String role, ResourcePattern // revoke privs from role roleManager.revokePrivs(role, resourcePattern, privs, errOnNonExist); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, resourcePattern, privs, null, role); @@ -1019,7 +994,6 @@ private void revokeInternal(UserIdentity userIdent, String role, WorkloadGroupPa // revoke privs from role roleManager.revokePrivs(role, workloadGroupPattern, privs, errOnNonExist); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, workloadGroupPattern, privs, null, role); @@ -1045,7 +1019,6 @@ private void revokeInternal(UserIdentity userIdent, List roles, boolean } } userRoleManager.removeUserRoles(userIdent, roles); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(userIdent, roles); Env.getCurrentEnv().getEditLog().logRevokePriv(info); @@ -1161,7 +1134,6 @@ private void createRoleInternal(String role, boolean ignoreIfExists, String comm } roleManager.addOrMergeRole(emptyPrivsRole, true /* err on exist */); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(role, comment); @@ -1195,7 +1167,6 @@ private void dropRoleInternal(String role, boolean ignoreIfNonExists, boolean is roleManager.dropRole(role, true /* err on non exist */); userRoleManager.dropRole(role); - markAuthorizationChanged(); if (!isReplay) { PrivInfo info = new PrivInfo(null, null, null, role, null, null, ""); Env.getCurrentEnv().getEditLog().logDropRole(info); @@ -2025,7 +1996,6 @@ private void setRoleToUser(UserIdentity userIdent, String role) throws DdlExcept userRoleManager.dropUser(userIdent); userRoleManager.addUserRole(userIdent, role); userRoleManager.addUserRole(userIdent, roleManager.getUserDefaultRoleName(userIdent)); - markAuthorizationChanged(); } private void updateUserTlsRequirements(UserIdentity userIdent) throws DdlException { diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java index f08d3193aa5ba6..7c3596e4def2eb 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/SecurityDependencyContext.java @@ -26,50 +26,37 @@ import org.apache.doris.catalog.TableIf; import org.apache.doris.datasource.CatalogIf; import org.apache.doris.datasource.InternalCatalog; -import org.apache.doris.mysql.privilege.Auth; import org.apache.doris.mysql.privilege.InternalAuthorizationPlugin; +import org.apache.doris.nereids.rules.analysis.UserAuthentication; import org.apache.doris.policy.PolicyMgr; import org.apache.doris.qe.ConnectContext; import org.apache.doris.qe.SessionVariable; import com.google.common.collect.ImmutableSet; +import java.util.ArrayList; import java.util.List; import java.util.Objects; import java.util.Optional; import java.util.Set; -/** - * Security state which a reusable prepared point-query plan depends on. - * - *

Row policies are deliberately not copied or compared here. A plan with a row policy is not eligible for - * short-circuit execution, while the process-local policy version invalidates a previously cached no-policy plan - * when a policy is later added. This keeps the common validation path to a few identity and volatile-version reads. - * Authorization sources without reliable local versions are never reused without replanning. - */ +/** Security dependencies of a reusable prepared point-query plan. */ public class SecurityDependencyContext { - private static final long UNKNOWN_VERSION = -1; - private final UserIdentity planningUserIdentity; private final Set planningAuthenticatedRoles; private final Env planningEnv; - private final long authorizationVersion; - private final long rowPolicyVersion; - private final boolean versionValidationEligible; - private boolean privilegeChecked; - private boolean internalCatalogOnly = true; - private boolean olapTableOnly = true; - private boolean hasRowPolicy; + private final boolean authorizationChecksEnabled; + private final List checkedPrivileges = new ArrayList<>(); + private boolean hasEffectiveRowPolicy; private boolean hasDataMask; - private boolean useVersionValidation; - private boolean complete = true; + private boolean complete; /** Create an incomplete context for tests and callers without a connection. */ public SecurityDependencyContext() { - this(null, ImmutableSet.of(), null, UNKNOWN_VERSION, UNKNOWN_VERSION, false); + this(null, ImmutableSet.of(), null, false); } - /** Capture the effective authorization subject and security versions before analysis starts. */ + /** Capture the authorization subject before analysis starts. */ public SecurityDependencyContext(ConnectContext connectContext) { this(connectContext == null ? null : connectContext.getCurrentUserIdentity(), authenticatedRoles(connectContext), @@ -78,24 +65,15 @@ public SecurityDependencyContext(ConnectContext connectContext) { } private SecurityDependencyContext(UserIdentity planningUserIdentity, Set planningAuthenticatedRoles, - Env planningEnv, boolean versionValidationEligible) { - this(planningUserIdentity, planningAuthenticatedRoles, planningEnv, - currentAuthorizationVersion(planningEnv), currentRowPolicyVersion(planningEnv), - versionValidationEligible); - } - - private SecurityDependencyContext(UserIdentity planningUserIdentity, Set planningAuthenticatedRoles, - Env planningEnv, long authorizationVersion, long rowPolicyVersion, - boolean versionValidationEligible) { + Env planningEnv, boolean authorizationChecksEnabled) { this.planningUserIdentity = planningUserIdentity; this.planningAuthenticatedRoles = planningAuthenticatedRoles; this.planningEnv = planningEnv; - this.authorizationVersion = authorizationVersion; - this.rowPolicyVersion = rowPolicyVersion; - this.versionValidationEligible = versionValidationEligible; + this.authorizationChecksEnabled = authorizationChecksEnabled; + this.complete = authorizationChecksEnabled; } - /** Record that SELECT privileges were checked and whether the relation supports version-only validation. */ + /** Record the exact SELECT check which must be repeated before direct reuse. */ public synchronized void addCheckedPrivilege(TableIf table, Set usedColumns) { if (table == null) { complete = false; @@ -103,115 +81,93 @@ public synchronized void addCheckedPrivilege(TableIf table, Set usedColu } DatabaseIf database = table.getDatabase(); CatalogIf catalog = database == null ? null : database.getCatalog(); - if (catalog == null) { + if (catalog == null + || !(table instanceof OlapTable) + || !InternalCatalog.INTERNAL_CATALOG_NAME.equals(catalog.getName())) { complete = false; return; } - privilegeChecked = true; - internalCatalogOnly &= InternalCatalog.INTERNAL_CATALOG_NAME.equals(catalog.getName()); - olapTableOnly &= table instanceof OlapTable; + checkedPrivileges.add(new CheckedPrivilege(table, database, catalog, catalog.getName(), + database.getFullName(), table.getName(), + usedColumns == null ? ImmutableSet.of() : ImmutableSet.copyOf(usedColumns))); } - /** Record only whether a row policy exists; policy objects never enter the point-query cache. */ - public synchronized void setRowPolicies( - String catalog, String database, String table, List policies) { - hasRowPolicy |= policies != null && !policies.isEmpty(); + /** Record mask presence so a masked plan is never reused without policy analysis. */ + public synchronized void addDataMask( + String catalog, String database, String table, String column, Optional mask) { + hasDataMask |= mask.isPresent(); } - /** A row policy makes the statement ineligible for short-circuit execution. */ - public synchronized boolean hasRowPolicy() { - return hasRowPolicy; + public synchronized boolean hasDataMask() { + return hasDataMask; } - /** Record mask presence so a masked plan is never accepted by the version-only cache path. */ - public synchronized void addDataMask( - String catalog, String database, String table, String column, Optional mask) { - hasDataMask |= mask.isPresent(); + /** Record only whether external policy analysis produced a row filter; definitions are not retained. */ + public synchronized void addRowPolicies(List policies) { + hasEffectiveRowPolicy |= policies != null && !policies.isEmpty(); } - /** Freeze the decisions used by a completed plan before storing them in a reusable context. */ - public synchronized SecurityDependencyContext snapshot() { - SecurityDependencyContext snapshot = new SecurityDependencyContext( - planningUserIdentity, planningAuthenticatedRoles, planningEnv, - authorizationVersion, rowPolicyVersion, versionValidationEligible); - snapshot.privilegeChecked = privilegeChecked; - snapshot.internalCatalogOnly = internalCatalogOnly; - snapshot.olapTableOnly = olapTableOnly; - snapshot.hasRowPolicy = hasRowPolicy; - snapshot.hasDataMask = hasDataMask; - snapshot.complete = complete; - snapshot.useVersionValidation = snapshot.canUseVersionValidation(); - return snapshot; + public synchronized boolean hasEffectiveRowPolicy() { + return hasEffectiveRowPolicy; } - /** Freeze a prepared short-circuit dependency set, failing closed if its proof is incomplete. */ + /** Freeze the completed dependency set before storing it with the prepared plan. */ public synchronized SecurityDependencyContext snapshotForShortCircuit() { - SecurityDependencyContext snapshot = snapshot(); - if (!privilegeChecked || hasRowPolicy) { - snapshot.complete = false; - snapshot.useVersionValidation = false; - } + SecurityDependencyContext snapshot = new SecurityDependencyContext( + planningUserIdentity, planningAuthenticatedRoles, planningEnv, authorizationChecksEnabled); + snapshot.checkedPrivileges.addAll(checkedPrivileges); + snapshot.hasEffectiveRowPolicy = hasEffectiveRowPolicy; + snapshot.hasDataMask = hasDataMask; + snapshot.complete = complete && !checkedPrivileges.isEmpty() + && !hasEffectiveRowPolicy && !hasDataMask + && checkedPrivileges.stream().allMatch(CheckedPrivilege::matchesNamespace); return snapshot; } /** - * Check whether an analyzed point-query plan can bypass planning again. - * - *

A false result rejects only cached reuse. The prepared statement is reparsed and analyzed normally, so - * authorization failures retain their standard user-facing error. The authorization subject is checked before - * the version shortcut because COM_CHANGE_USER keeps the connection's prepared statements alive. + * Recheck the small set of facts needed to bypass planning. Returning false only rejects direct reuse; normal + * planning then performs the authoritative check and reports its standard error. */ public boolean isValid(ConnectContext connectContext) { - if (!complete || !useVersionValidation || connectContext == null + if (!complete || !authorizationChecksEnabled || connectContext == null + || planningEnv == null || planningEnv != connectContext.getEnv() || !Objects.equals(planningUserIdentity, connectContext.getCurrentUserIdentity()) - || !planningAuthenticatedRoles.equals(authenticatedRoles(connectContext))) { + || !planningAuthenticatedRoles.equals(authenticatedRoles(connectContext)) + || !usesAuthorizationChecks(connectContext) + || !usesBuiltInAuthorization(planningEnv)) { return false; } try { - return usesAuthorizationChecks(connectContext) && versionsAreCurrent(connectContext.getEnv()); - } catch (RuntimeException e) { + PolicyMgr policyMgr = planningEnv.getPolicyMgr(); + if (policyMgr == null) { + return false; + } + for (CheckedPrivilege checkedPrivilege : checkedPrivileges) { + if (!checkedPrivilege.matchesNamespace()) { + return false; + } + if (policyMgr.hasRowPolicy(checkedPrivilege.catalog, checkedPrivilege.database, + checkedPrivilege.tableName)) { + return false; + } + UserAuthentication.checkPermission( + checkedPrivilege.table, connectContext, checkedPrivilege.usedColumns); + if (!checkedPrivilege.matchesNamespace()) { + return false; + } + } + return true; + } catch (Exception e) { return false; } } - private boolean canUseVersionValidation() { - return complete && versionValidationEligible && privilegeChecked && internalCatalogOnly && olapTableOnly - && !hasRowPolicy && !hasDataMask - && authorizationVersion != UNKNOWN_VERSION && rowPolicyVersion != UNKNOWN_VERSION - && usesVersionedBuiltInAuthorization(planningEnv); - } - - private boolean versionsAreCurrent(Env env) { - if (env == null || env != planningEnv) { - return false; - } - Auth auth = env.getAuth(); - PolicyMgr policyMgr = env.getPolicyMgr(); - return auth != null && policyMgr != null - && auth.isAuthorizationVersionReliable() - && env.getAccessManager().getAccessControllerOrDefault(InternalCatalog.INTERNAL_CATALOG_NAME) - instanceof InternalAuthorizationPlugin - && auth.getAuthorizationVersion() == authorizationVersion - && policyMgr.getRowPolicyVersion() == rowPolicyVersion; - } - - private static boolean usesVersionedBuiltInAuthorization(Env env) { - return env != null && env.getAuth() != null && env.getPolicyMgr() != null - && env.getAuth().isAuthorizationVersionReliable() + private static boolean usesBuiltInAuthorization(Env env) { + return env != null && env.getAccessManager() != null && env.getAccessManager().getAccessControllerOrDefault(InternalCatalog.INTERNAL_CATALOG_NAME) instanceof InternalAuthorizationPlugin; } - private static long currentAuthorizationVersion(Env env) { - Auth auth = env == null ? null : env.getAuth(); - return auth == null ? UNKNOWN_VERSION : auth.getAuthorizationVersion(); - } - - private static long currentRowPolicyVersion(Env env) { - PolicyMgr policyMgr = env == null ? null : env.getPolicyMgr(); - return policyMgr == null ? UNKNOWN_VERSION : policyMgr.getRowPolicyVersion(); - } - private static boolean usesAuthorizationChecks(ConnectContext connectContext) { if (connectContext == null || connectContext.isSkipAuth()) { return false; @@ -227,4 +183,35 @@ private static Set authenticatedRoles(ConnectContext connectContext) { Set roles = connectContext.getAuthenticatedRoles(); return roles == null || roles.isEmpty() ? ImmutableSet.of() : ImmutableSet.copyOf(roles); } + + private static class CheckedPrivilege { + private final TableIf table; + private final DatabaseIf databaseObject; + private final CatalogIf catalogObject; + private final String catalog; + private final String database; + private final String tableName; + private final Set usedColumns; + + private CheckedPrivilege(TableIf table, DatabaseIf databaseObject, CatalogIf catalogObject, + String catalog, String database, String tableName, Set usedColumns) { + this.table = table; + this.databaseObject = databaseObject; + this.catalogObject = catalogObject; + this.catalog = catalog; + this.database = database; + this.tableName = tableName; + this.usedColumns = usedColumns; + } + + private boolean matchesNamespace() { + DatabaseIf currentDatabase = table.getDatabase(); + CatalogIf currentCatalog = currentDatabase == null ? null : currentDatabase.getCatalog(); + return currentDatabase == databaseObject + && currentCatalog == catalogObject + && Objects.equals(database, currentDatabase == null ? null : currentDatabase.getFullName()) + && Objects.equals(catalog, currentCatalog == null ? null : currentCatalog.getName()) + && Objects.equals(tableName, table.getName()); + } + } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java index 7bf7ab6f3e0c1f..2b75cf43557ed0 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java @@ -55,7 +55,6 @@ import org.apache.doris.nereids.trees.expressions.Slot; import org.apache.doris.nereids.trees.expressions.SlotReference; import org.apache.doris.nereids.trees.expressions.StatementScopeIdGenerator; -import org.apache.doris.nereids.trees.expressions.literal.Literal; import org.apache.doris.nereids.trees.plans.ObjectId; import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.nereids.trees.plans.Plan; @@ -185,22 +184,17 @@ public enum TableFrom { private final Set viewDdlSqlSet = Sets.newHashSet(); private final SqlCacheContext sqlCacheContext; private final SecurityDependencyContext securityDependencyContext; + private boolean hasNonFilterPlaceholder; // generate for next id for prepared statement's placeholders, which is // connection level private final IdGenerator placeHolderIdGenerator = PlaceholderId.createGenerator(); // relation id to placeholders for prepared statement, ordered by placeholder id private final Map idToPlaceholderRealExpr = new TreeMap<>(); - // Map placeholder id to the physical key slot used by the immutable point-query template. + // map placeholder id to comparison slot, which will used to replace conjuncts + // directly private final Map idToComparisonSlot = new TreeMap<>(); - // Equality literals written as constants in the statement. They are deliberately separate - // from placeholder bindings: a prepared point query must never replace a fixed predicate - // merely because it references the same column as a placeholder. Plans with row policies - // are not eligible for the point-query shortcut. - private final List pointQueryFixedKeyConstraints = new ArrayList<>(); - private boolean pointQueryFixedKeyConstraintsComplete = true; - // collect all hash join conditions to compute node connectivity in join graph private final List joinFilters = new ArrayList<>(); @@ -290,9 +284,6 @@ public enum TableFrom { private ShortCircuitQueryContext shortCircuitQueryContext; - // Built afresh for one EXECUTE. Never copied into the next StatementContext. - private ShortCircuitQueryContext.PointQueryExecutionContext pointQueryExecutionContext; - private FormatOptions formatOptions = FormatOptions.getDefault(); private Set plannerHooks = new HashSet<>(); @@ -447,8 +438,8 @@ public StatementContext createNextExecuteContext() { next.cteIdGenerator.resetId(cteIdGenerator.getCurrentId()); next.talbeIdGenerator.resetId(talbeIdGenerator.getCurrentId()); next.placeHolderIdGenerator.resetId(placeHolderIdGenerator.getCurrentId()); - // Copy this EXECUTE's placeholder values and the stable placeholder-to-key registry. - // Fixed constraints and bound key tuples remain local to the context that owns them. + // Placeholder bindings of this EXECUTE, and the comparison-slot registry used to replace + // conjuncts on the cached short-circuit plan without re-planning. next.idToPlaceholderRealExpr.putAll(idToPlaceholderRealExpr); next.idToComparisonSlot.putAll(idToComparisonSlot); next.placeholders = new ArrayList<>(placeholders); @@ -700,15 +691,6 @@ public void setShortCircuitQueryContext(ShortCircuitQueryContext shortCircuitQue this.shortCircuitQueryContext = shortCircuitQueryContext; } - public ShortCircuitQueryContext.PointQueryExecutionContext getPointQueryExecutionContext() { - return pointQueryExecutionContext; - } - - public void setPointQueryExecutionContext( - ShortCircuitQueryContext.PointQueryExecutionContext pointQueryExecutionContext) { - this.pointQueryExecutionContext = pointQueryExecutionContext; - } - public Optional getSqlCacheContext() { return Optional.ofNullable(sqlCacheContext); } @@ -717,6 +699,14 @@ public SecurityDependencyContext getSecurityDependencyContext() { return securityDependencyContext; } + public boolean hasNonFilterPlaceholder() { + return hasNonFilterPlaceholder; + } + + public void setHasNonFilterPlaceholder(boolean hasNonFilterPlaceholder) { + this.hasNonFilterPlaceholder = hasNonFilterPlaceholder; + } + public boolean isDpHyp() { return isDpHyp; } @@ -862,22 +852,6 @@ public Map getIdToComparisonSlot() { return idToComparisonSlot; } - public void addPointQueryFixedKeyConstraint(SlotReference slot, Literal literal) { - pointQueryFixedKeyConstraints.add(new PointQueryFixedKeyConstraint(slot, literal)); - } - - public List getPointQueryFixedKeyConstraints() { - return pointQueryFixedKeyConstraints; - } - - public void markPointQueryFixedKeyConstraintsIncomplete() { - pointQueryFixedKeyConstraintsComplete = false; - } - - public boolean arePointQueryFixedKeyConstraintsComplete() { - return pointQueryFixedKeyConstraintsComplete; - } - public Map, Group>>> getCteIdToConsumerGroup() { return cteIdToConsumerGroup; } @@ -1706,23 +1680,4 @@ public boolean isDelete() { public void setIsDelete(boolean del) { isDelete = del; } - - /** A fixed equality operand and the exact bound slot it constrains. */ - public static class PointQueryFixedKeyConstraint { - private final SlotReference slot; - private final Literal literal; - - public PointQueryFixedKeyConstraint(SlotReference slot, Literal literal) { - this.slot = Objects.requireNonNull(slot); - this.literal = Objects.requireNonNull(literal); - } - - public SlotReference getSlot() { - return slot; - } - - public Literal getLiteral() { - return literal; - } - } } 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 0e6891933b47cb..f8fd0910ae12db 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 @@ -92,7 +92,6 @@ import org.apache.doris.nereids.trees.expressions.typecoercion.ImplicitCastInputTypes; import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.nereids.trees.plans.Plan; -import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalJoin; import org.apache.doris.nereids.trees.plans.logical.LogicalPlan; import org.apache.doris.nereids.types.ArrayType; @@ -921,11 +920,13 @@ public Expression visitPlaceholder(Placeholder placeholder, ExpressionRewriteCon return visit(realExpr, context); } - // Register each prepared-statement placeholder with its point-query key slot. + // Register prepared statement placeholder id to related slot in comparison predicate. + // Used to replace expression in ShortCircuit plan private void registerPlaceholderIdToSlot(ComparisonPredicate cp, ExpressionRewriteContext context, Expression left, Expression right) { if (ConnectContext.get() != null && ConnectContext.get().getCommand() == MysqlCommand.COM_STMT_EXECUTE) { + // Used to replace expression in ShortCircuit plan if (cp.right() instanceof Placeholder && left instanceof SlotReference) { PlaceholderId id = ((Placeholder) cp.right()).getPlaceholderId(); context.cascadesContext.getStatementContext().getIdToComparisonSlot().put(id, (SlotReference) left); @@ -940,44 +941,12 @@ private void registerPlaceholderIdToSlot(ComparisonPredicate cp, public Expression visitComparisonPredicate(ComparisonPredicate cp, ExpressionRewriteContext context) { Expression left = cp.left().accept(this, context); Expression right = cp.right().accept(this, context); + // Used to replace expression in ShortCircuit plan registerPlaceholderIdToSlot(cp, context, left, right); - ComparisonPredicate original = cp; cp = (ComparisonPredicate) cp.withChildren(left, right); - Expression analyzed = isEqualityBetweenJoinChildren(cp) + return isEqualityBetweenJoinChildren(cp) ? TypeCoercionUtils.processJoinComparisonPredicate(cp) : TypeCoercionUtils.processComparisonPredicate(cp); - registerPointQueryFixedKeyConstraint(original, analyzed, context); - return analyzed; - } - - /** - * Keep fixed equality values distinct from prepared-statement placeholders. The point-query - * executor used to rediscover both from the translated scan conjuncts and then update every - * predicate sharing a column name. A statement containing both {@code key = ?} and - * {@code key = constant} would therefore lose provenance and turn the fixed value into - * caller-controlled state. - */ - private void registerPointQueryFixedKeyConstraint(ComparisonPredicate original, - Expression analyzed, ExpressionRewriteContext context) { - if (!(currentPlan instanceof LogicalFilter) - || !(original instanceof EqualTo) - || original.left() instanceof Placeholder - || original.right() instanceof Placeholder - || !(analyzed instanceof EqualTo)) { - return; - } - Expression left = analyzed.child(0); - Expression right = analyzed.child(1); - if (left instanceof SlotReference && right instanceof Literal) { - context.cascadesContext.getStatementContext().addPointQueryFixedKeyConstraint( - (SlotReference) left, (Literal) right); - } else { - // The logical short-circuit shape checker currently peels one Cast from the key. - // A cast can be lossy (for example CAST(INT AS CHAR(1))), so it is not evidence for - // an exact physical lookup key. Keep the normal plan unless the fixed predicate is - // literally Slot = Literal. - context.cascadesContext.getStatementContext().markPointQueryFixedKeyConstraintsIncomplete(); - } } private boolean isEqualityBetweenJoinChildren(ComparisonPredicate comparisonPredicate) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java index ff92869fadecb7..cfe01c1299ee1b 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java @@ -18,7 +18,9 @@ package org.apache.doris.nereids.rules.rewrite; import org.apache.doris.catalog.Column; +import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.OlapTable; +import org.apache.doris.datasource.CatalogIf; import org.apache.doris.datasource.doris.RemoteOlapTable; import org.apache.doris.nereids.StatementContext; import org.apache.doris.nereids.rules.Rule; @@ -30,6 +32,7 @@ import org.apache.doris.nereids.trees.plans.Plan; import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalOlapScan; +import org.apache.doris.policy.PolicyMgr; import org.apache.doris.qe.ConnectContext; import org.apache.doris.qe.ConnectContext.ConnectType; @@ -46,8 +49,9 @@ */ public class LogicalResultSinkToShortCircuitPointQuery implements RewriteRuleFactory { - private Expression removeCast(Expression expression) { - if (expression instanceof Cast) { + private Expression removeInjectiveCast(Expression expression) { + if (expression instanceof Cast + && expression.child(0).getDataType().isInjectiveCastTo(expression.getDataType())) { return expression.child(0); } return expression; @@ -55,14 +59,29 @@ private Expression removeCast(Expression expression) { private boolean filterMatchShortCircuitCondition(LogicalFilter filter) { return filter.getConjuncts().stream().allMatch( - // all conjuncts match with pattern `key = ?` + // all conjuncts match with pattern `key = literal` expression -> (expression instanceof EqualTo) - && (removeCast(expression.child(0)).isKeyColumnFromTable() + && (removeInjectiveCast(expression.child(0)).isKeyColumnFromTable() || (expression.child(0) instanceof SlotReference && ((SlotReference) expression.child(0)).getName().equals(Column.DELETE_SIGN))) && expression.child(1).isLiteral()); } + /** Any row policy on the table makes point-query planning ineligible, regardless of its target. */ + private boolean hasRowPolicy(OlapTable table, StatementContext statementContext) { + try { + DatabaseIf database = table.getDatabase(); + CatalogIf catalog = database == null ? null : database.getCatalog(); + ConnectContext connectContext = statementContext.getConnectContext(); + PolicyMgr policyMgr = connectContext == null || connectContext.getEnv() == null + ? null : connectContext.getEnv().getPolicyMgr(); + return database == null || catalog == null || policyMgr == null + || policyMgr.hasRowPolicy(catalog.getName(), database.getFullName(), table.getName()); + } catch (RuntimeException e) { + return true; + } + } + @VisibleForTesting boolean scanMatchShortCircuitCondition(LogicalOlapScan olapScan) { ConnectContext connectContext = ConnectContext.get(); @@ -103,25 +122,29 @@ boolean scanMatchShortCircuitCondition(LogicalOlapScan olapScan) { // set short circuit flag and return the original plan private Plan shortCircuit(Plan root, OlapTable olapTable, Set conjuncts, StatementContext statementContext) { - // Row filters are injected into the analyzed plan and views are inlined. Neither shape has a - // cheap, stable dependency fence suitable for a reusable direct plan, so keep both on the - // normal execution path. A global row-policy epoch still invalidates a no-policy plan if a - // policy is added after it was cached. - if (statementContext.getSecurityDependencyContext().hasRowPolicy() - || !statementContext.getViewDdlSqls().isEmpty() - || !statementContext.arePointQueryFixedKeyConstraintsComplete()) { + // Keep policy-bearing tables, inlined views, and placeholders outside the final filter on + // the normal path. A cached no-policy plan repeats the table-level lookup before reuse. + if (hasRowPolicy(olapTable, statementContext) + || statementContext.getSecurityDependencyContext().hasEffectiveRowPolicy() + || statementContext.getSecurityDependencyContext().hasDataMask() + || statementContext.hasNonFilterPlaceholder() + || !statementContext.getViewDdlSqls().isEmpty()) { return root; } // All key columns in conjuncts Set colNames = Sets.newHashSet(); for (Expression expr : conjuncts) { - SlotReference slot = ((SlotReference) removeCast((expr.child(0)))); + SlotReference slot = (SlotReference) removeInjectiveCast(expr.child(0)); if (slot.isKeyColumnFromTable()) { - colNames.add(slot.getName()); + // The executor updates cached conjuncts by column name. More than one predicate on + // the same key would make a fixed literal indistinguishable from a placeholder. + if (!colNames.add(slot.getName())) { + return root; + } } } // set short circuit flag and modify nothing to the plan - if (olapTable.getBaseSchemaKeyColumns().size() <= colNames.size()) { + if (olapTable.getBaseSchemaKeyColumns().size() == colNames.size()) { statementContext.setShortCircuitQuery(true); } return root; diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java index fedeab96956e59..d9f558e10f0b40 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java @@ -35,6 +35,7 @@ import org.apache.doris.nereids.trees.plans.commands.insert.InsertOverwriteTableCommand; import org.apache.doris.nereids.trees.plans.commands.insert.OlapGroupCommitInsertExecutor; import org.apache.doris.nereids.trees.plans.commands.merge.MergeIntoCommand; +import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalPlan; import org.apache.doris.nereids.trees.plans.logical.LogicalSqlCache; import org.apache.doris.nereids.trees.plans.visitor.PlanVisitor; @@ -117,6 +118,7 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { } // Commands hide their retained query trees from normal plan traversal. Reset every exposed // root so a later EXECUTE cannot reuse a relation-local snapshot from an earlier execution. + boolean hasNonFilterPlaceholder = false; for (int rootIndex = 0; rootIndex < relationRoots.size(); rootIndex++) { LogicalPlan relationRoot = relationRoots.get(rootIndex); for (UnboundRelation relation : relationRoot.collectToList( @@ -128,6 +130,10 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { } for (LogicalPlan plan : relationRoot.collectToList(node -> true)) { for (Expression expression : plan.getExpressions()) { + if (!(plan instanceof LogicalFilter) + && expression.anyMatch(Placeholder.class::isInstance)) { + hasNonFilterPlaceholder = true; + } for (SubqueryExpr subquery : expression.collectToList( SubqueryExpr.class::isInstance)) { // SubqueryExpr owns its query plan outside Plan.children(), so retained prepared @@ -137,6 +143,10 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { } } } + statementContext.setHasNonFilterPlaceholder(hasNonFilterPlaceholder); + if (hasNonFilterPlaceholder) { + statementContext.setShortCircuitQuery(false); + } if (logicalPlan instanceof LogicalSqlCache) { throw new AnalysisException("Unsupported sql cache for server prepared statement"); } @@ -162,10 +172,8 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { // statementContext.getShortCircuitQueryContext(), and the fallback (building one from a // null planner, since this path skips planning) would NPE. statementContext.setShortCircuitQueryContext(preparedStmtCtx.shortCircuitQueryContext.get()); - if (PointQueryExecutor.directExecuteShortCircuitQuery( - executor, preparedStmtCtx, statementContext)) { - return; - } + PointQueryExecutor.directExecuteShortCircuitQuery(executor, preparedStmtCtx, statementContext); + return; } if (ctx.getSessionVariable().enableGroupCommitFullPrepare) { if (preparedStmtCtx.groupCommitPlanner.isPresent()) { @@ -187,15 +195,20 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { // Drop the previously cached short-circuit context: either it was reusable and returned // early above, has just been refreshed here, or is stale and we are about to re-plan. preparedStmtCtx.shortCircuitQueryContext = Optional.empty(); + // The inherited flag is only for deciding direct reuse above. Recompute it from the + // current plan so a newly added policy or another eligibility change cannot stay cached. + statementContext.setShortCircuitQuery(false); executor.execute(); StatementContext executedStatementContext = executor.getContext().getStatementContext(); ShortCircuitQueryContext shortCircuitQueryContext = executedStatementContext.getShortCircuitQueryContext(); if (shortCircuitQueryContext != null) { Preconditions.checkState(executedStatementContext.isShortCircuitQuery()); - // Publish the exact context used by this execution so its topology generation stays - // bound to the cached partition pruner in the same planner scan node. - preparedStmtCtx.shortCircuitQueryContext = Optional.of(shortCircuitQueryContext); + // Publish the exact context used by this execution only if the security and namespace + // snapshot still matches after planning and execution. + if (shortCircuitQueryContext.isReusable(ctx)) { + preparedStmtCtx.shortCircuitQueryContext = Optional.of(shortCircuitQueryContext); + } } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java index 5e387bc482ef70..4add891c85cd47 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/logical/LogicalCheckPolicy.java @@ -229,7 +229,7 @@ public RelatedPolicy findPolicy(LogicalPlan logicalPlan, CascadesContext cascade if (sqlCacheContext.isPresent()) { sqlCacheContext.get().setRowFilterPolicy(ctlName, dbName, tableName, rowPolicies); } - securityDependencyContext.setRowPolicies(ctlName, dbName, tableName, rowPolicies); + securityDependencyContext.addRowPolicies(rowPolicies); return new RelatedPolicy( Optional.ofNullable(CollectionUtils.isEmpty(rowPolicies) diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java index f9dc65012605c8..fcbd915100261e 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/OlapScanNode.java @@ -1136,18 +1136,6 @@ public List lazyEvaluateRangeLocations() throws UserExcepti selectedIndexId = olapTable.getBaseIndexId(); // Only key columns computeColumnsFilter(olapTable.getBaseSchemaKeyColumns(), olapTable.getPartitionInfo()); - return evaluatePointQueryRangeLocations(); - } - - // Prepared point queries use execution-owned values so the cached conjunct template stays immutable. - public List lazyEvaluateRangeLocations( - Map keyValues) throws UserException { - selectedIndexId = olapTable.getBaseIndexId(); - computePointQueryColumnFilters(keyValues); - return evaluatePointQueryRangeLocations(); - } - - private List evaluatePointQueryRangeLocations() throws UserException { computePartitionInfo(); scanBackendIds.clear(); selectionHint = null; diff --git a/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java index 3be660eb17dc8a..98d9056e1af08c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/planner/ScanNode.java @@ -73,7 +73,6 @@ import org.apache.logging.log4j.Logger; import java.util.ArrayList; -import java.util.Collections; import java.util.HashSet; import java.util.LinkedHashSet; import java.util.List; @@ -202,28 +201,6 @@ public void computeColumnsFilter(List columns, PartitionInfo partitionsI } } - /** - * Build point-query pruning state from one execution's immutable key tuple. This avoids - * changing cached scan conjuncts (which are shared by every EXECUTE of a prepared handle) - * while still making partition and distribution pruning use the current parameter values. - */ - protected void computePointQueryColumnFilters(Map keyValues) { - columnFilters.clear(); - columnNameToRange.clear(); - for (Map.Entry entry : keyValues.entrySet()) { - LiteralExpr literal = entry.getValue(); - PartitionColumnFilter partitionFilter = new PartitionColumnFilter(); - partitionFilter.setLowerBound(literal, true); - partitionFilter.setUpperBound(literal, true); - columnFilters.put(entry.getKey(), partitionFilter); - - ColumnBound bound = ColumnBound.of(literal); - ColumnRange columnRange = ColumnRange.create(); - columnRange.intersect(Collections.singletonList(Range.closed(bound, bound))); - columnNameToRange.put(entry.getKey(), columnRange); - } - } - public void computeColumnsFilter() { // for load scan node, table is null // partitionsInfo maybe null for other scan node, eg: ExternalScanNode... diff --git a/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java b/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java index 4e4f104606a1a2..d8c0f02eafef27 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java +++ b/fe/fe-core/src/main/java/org/apache/doris/policy/PolicyMgr.java @@ -73,13 +73,6 @@ public class PolicyMgr implements Writable { // ctlName -> dbName -> tableName -> List private Map>>> tablePolicies = Maps.newConcurrentMap(); - // Process-local epoch used to invalidate prepared point-query plans after row-policy changes. - private transient volatile long rowPolicyVersion; - - public long getRowPolicyVersion() { - return rowPolicyVersion; - } - private void writeLock() { lock.writeLock().lock(); } @@ -308,7 +301,6 @@ private void unprotectedAdd(Policy policy) { typeToPolicyMap.put(policy.getType(), dbPolicies); if (PolicyTypeEnum.ROW == policy.getType()) { addTablePolicies((RowPolicy) policy); - rowPolicyVersion++; } } @@ -344,7 +336,7 @@ public void replayStoragePolicyAlter(StoragePolicy log) { private void unprotectedDrop(DropPolicyLog log) { List policies = getPoliciesByType(log.getType()); - boolean removed = policies.removeIf(policy -> { + policies.removeIf(policy -> { if (policy.matchPolicy(log)) { if (policy instanceof StoragePolicy) { ((StoragePolicy) policy).removeResourceReference(); @@ -360,8 +352,24 @@ private void unprotectedDrop(DropPolicyLog log) { return false; }); typeToPolicyMap.put(log.getType(), policies); - if (removed && log.getType() == PolicyTypeEnum.ROW) { - rowPolicyVersion++; + } + + /** Return whether the table has any row policy, without matching users/roles or parsing policy expressions. */ + public boolean hasRowPolicy(String ctlName, String dbName, String tableName) { + readLock(); + try { + Map>> dbPolicies = tablePolicies.get(ctlName); + if (dbPolicies == null) { + return false; + } + Map> tablePolicyMap = dbPolicies.get(dbName); + if (tablePolicyMap == null) { + return false; + } + List policies = tablePolicyMap.get(tableName); + return policies != null && !policies.isEmpty(); + } finally { + readUnlock(); } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java b/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java index 7042a776e33e6c..e60ae03020742d 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/PointQueryExecutor.java @@ -17,10 +17,14 @@ package org.apache.doris.qe; +import org.apache.doris.analysis.BinaryPredicate; import org.apache.doris.analysis.Expr; +import org.apache.doris.analysis.ExprToSqlVisitor; import org.apache.doris.analysis.ExprToThriftVisitor; import org.apache.doris.analysis.LiteralExpr; import org.apache.doris.analysis.LiteralExprUtils; +import org.apache.doris.analysis.SlotRef; +import org.apache.doris.analysis.ToSqlParams; import org.apache.doris.catalog.Column; import org.apache.doris.catalog.Env; import org.apache.doris.catalog.OlapTable; @@ -31,11 +35,13 @@ import org.apache.doris.common.UserException; import org.apache.doris.mysql.MysqlCommand; import org.apache.doris.nereids.StatementContext; +import org.apache.doris.nereids.exceptions.AnalysisException; +import org.apache.doris.nereids.trees.expressions.SlotReference; +import org.apache.doris.nereids.trees.expressions.literal.Literal; +import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.planner.OlapScanNode; import org.apache.doris.proto.InternalService; import org.apache.doris.proto.InternalService.KeyTuple; -import org.apache.doris.qe.ShortCircuitQueryContext.PointQueryExecutionContext; -import org.apache.doris.qe.ShortCircuitQueryContext.PointQueryExecutionContext.Decision; import org.apache.doris.rpc.BackendServiceProxy; import org.apache.doris.rpc.RpcException; import org.apache.doris.rpc.TCustomProtocolFactory; @@ -50,6 +56,7 @@ import com.google.common.base.Preconditions; import com.google.common.base.Strings; import com.google.common.collect.Lists; +import com.google.common.collect.Maps; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; import org.apache.thrift.TDeserializer; @@ -61,6 +68,8 @@ import java.util.HashSet; import java.util.Iterator; import java.util.List; +import java.util.Map; +import java.util.Map.Entry; import java.util.Set; import java.util.concurrent.ExecutionException; import java.util.concurrent.Future; @@ -80,13 +89,10 @@ public class PointQueryExecutor implements CoordInterface { private List snapshotVisibleVersions; private final ShortCircuitQueryContext shortCircuitQueryContext; - private final PointQueryExecutionContext executionContext; - public PointQueryExecutor(ShortCircuitQueryContext ctx, - PointQueryExecutionContext executionContext, int maxMessageSize) { + public PointQueryExecutor(ShortCircuitQueryContext ctx, int maxMessageSize) { ctx.sanitize(); this.shortCircuitQueryContext = ctx; - this.executionContext = executionContext; this.maxMsgSizeOfResultReceiver = maxMessageSize; } @@ -110,7 +116,7 @@ private void updateCloudPartitionVersions() throws RpcException { void setScanRangeLocations() throws Exception { OlapScanNode scanNode = shortCircuitQueryContext.scanNode; // compute scan range - List locations = scanNode.lazyEvaluateRangeLocations(executionContext.getKeyValues()); + List locations = scanNode.lazyEvaluateRangeLocations(); Preconditions.checkNotNull(locations); if (scanNode.getScanTabletIds().isEmpty()) { return; @@ -145,28 +151,52 @@ static boolean shouldShuffleCandidateBackends(OlapScanNode scanNode) { } // execute query without analyze & plan - public static boolean directExecuteShortCircuitQuery(StmtExecutor executor, + public static void directExecuteShortCircuitQuery(StmtExecutor executor, PreparedStatementContext preparedStmtCtx, StatementContext statementContext) throws Exception { Preconditions.checkNotNull(preparedStmtCtx.shortCircuitQueryContext); ShortCircuitQueryContext shortCircuitQueryContext = preparedStmtCtx.shortCircuitQueryContext.get(); - PointQueryExecutionContext executionContext = - shortCircuitQueryContext.createPointQueryExecutionContext(statementContext); - if (executionContext.getDecision() == Decision.FALLBACK) { - // The copied prepared StatementContext still carries the previous execution's - // short-circuit flag. Clear all fast-path state before normal planning; otherwise - // planner construction can treat an unplanned statement as a point-query plan and - // try to build a ShortCircuitQueryContext from an empty scan-node list. - statementContext.setShortCircuitQuery(false); - statementContext.setShortCircuitQueryContext(null); - statementContext.setPointQueryExecutionContext(null); - return false; + // update conjuncts + Map colNameToConjunct = Maps.newHashMap(); + for (Entry entry : statementContext.getIdToComparisonSlot().entrySet()) { + String colName = entry.getValue().getOriginalColumn().get().getName(); + Expr conjunctVal = ((Literal) statementContext.getIdToPlaceholderRealExpr() + .get(entry.getKey())).toLegacyLiteral(); + colNameToConjunct.put(colName, conjunctVal); } - statementContext.setPointQueryExecutionContext(executionContext); + if (colNameToConjunct.size() != preparedStmtCtx.command.placeholderCount()) { + throw new AnalysisException("Mismatched conjuncts values size with prepared" + + "statement parameters size, expected " + + preparedStmtCtx.command.placeholderCount() + + ", but meet " + colNameToConjunct.size()); + } + updateScanNodeConjuncts(shortCircuitQueryContext.scanNode, colNameToConjunct); // short circuit plan and execution executor.executeAndSendResult(false, false, shortCircuitQueryContext.analzyedQuery, executor.getContext().getResultSender(), null, null); - return true; + } + + private static void updateScanNodeConjuncts(OlapScanNode scanNode, + Map colNameToConjunct) { + for (Expr conjunct : scanNode.getConjuncts()) { + BinaryPredicate binaryPredicate = (BinaryPredicate) conjunct; + SlotRef slot = null; + int updateChildIdx = 0; + if (binaryPredicate.getChild(0) instanceof LiteralExpr) { + slot = (SlotRef) binaryPredicate.getChildWithoutCast(1); + } else if (binaryPredicate.getChild(1) instanceof LiteralExpr) { + slot = (SlotRef) binaryPredicate.getChildWithoutCast(0); + updateChildIdx = 1; + } else { + Preconditions.checkState(false, "Should contains literal in " + + binaryPredicate.accept(ExprToSqlVisitor.INSTANCE, ToSqlParams.WITH_TABLE)); + } + // not a placeholder to replace + if (!colNameToConjunct.containsKey(slot.getColumnName())) { + continue; + } + binaryPredicate.setChild(updateChildIdx, colNameToConjunct.get(slot.getColumnName())); + } } public void setTimeout(long timeoutMs) { @@ -176,13 +206,21 @@ public void setTimeout(long timeoutMs) { void addKeyTuples( InternalService.PTabletKeyLookupRequest.Builder requestBuilder) throws TException { // TODO handle IN predicates + Map columnExpr = Maps.newHashMap(); KeyTuple.Builder kBuilder = KeyTuple.newBuilder(); + for (Expr expr : shortCircuitQueryContext.scanNode.getConjuncts()) { + BinaryPredicate predicate = (BinaryPredicate) expr; + Expr left = predicate.getChild(0); + Expr right = predicate.getChild(1); + SlotRef columnSlot = left.unwrapSlotRef(); + columnExpr.put(columnSlot.getColumnName(), right); + } // Serialize each literal expr as TExprNode bytes for typed value transfer. // BE deserializes the TExprNode and uses DataType::get_field() to extract // typed Field values directly, avoiding string parsing. TSerializer serializer = new TSerializer(); for (Column column : shortCircuitQueryContext.scanNode.getOlapTable().getBaseSchemaKeyColumns()) { - Expr literalExpr = executionContext.getKeyValues().get(column.getName()); + Expr literalExpr = columnExpr.get(column.getName()); // Ensure the literal type matches the column type for proper TExprNode // deserialization on BE side. Prepared statement parameters may have // mismatched types (e.g., setBigDecimal for INT column produces a @@ -222,13 +260,6 @@ public void cancel(Status cancelReason) { @Override public RowBatch getNext() throws Exception { - // A NULL placeholder or a fixed security constraint that disagrees with the bound - // key makes the full WHERE predicate false/unknown. Return before tablet pruning, - // cloud version lookup, or any BE RPC. - if (executionContext.getDecision() == Decision.EMPTY) { - return new RowBatch(); - } - Preconditions.checkState(executionContext.getDecision() == Decision.LOOKUP); setScanRangeLocations(); // No partition/tablet found return emtpy row batch if (candidateBackends == null || candidateBackends.isEmpty()) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java index 755f72dd5eca24..e71b1878f49b0f 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/ShortCircuitQueryContext.java @@ -20,28 +20,12 @@ import org.apache.doris.analysis.DescriptorToThriftConverter; import org.apache.doris.analysis.Expr; import org.apache.doris.analysis.ExprToThriftVisitor; -import org.apache.doris.analysis.LiteralExpr; -import org.apache.doris.analysis.LiteralExprUtils; import org.apache.doris.analysis.Queriable; -import org.apache.doris.catalog.Column; -import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.OlapTable; import org.apache.doris.catalog.Type; -import org.apache.doris.datasource.CatalogIf; import org.apache.doris.nereids.NereidsPlanner; import org.apache.doris.nereids.SecurityDependencyContext; import org.apache.doris.nereids.StatementContext; -import org.apache.doris.nereids.StatementContext.PointQueryFixedKeyConstraint; -import org.apache.doris.nereids.rules.expression.rules.FoldConstantRuleOnFE; -import org.apache.doris.nereids.trees.expressions.EqualTo; -import org.apache.doris.nereids.trees.expressions.Expression; -import org.apache.doris.nereids.trees.expressions.Placeholder; -import org.apache.doris.nereids.trees.expressions.SlotReference; -import org.apache.doris.nereids.trees.expressions.literal.BooleanLiteral; -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.util.TypeCoercionUtils; import org.apache.doris.planner.OlapScanNode; import org.apache.doris.planner.Planner; import org.apache.doris.thrift.TExpr; @@ -56,12 +40,9 @@ import org.apache.thrift.TSerializer; import java.util.ArrayList; -import java.util.Collections; -import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Objects; -import java.util.TreeMap; import java.util.UUID; import java.util.stream.Collectors; @@ -83,12 +64,10 @@ public class ShortCircuitQueryContext { public final String tableName; private final long fileCacheQueryLimitBytes; private final long partitionTopologyVersion; - private final TableNamespaceSnapshot tableNamespaceSnapshot; private final SecurityDependencyContext securityDependencyContext; public final OlapScanNode scanNode; public final Queriable analzyedQuery; - private final PointQueryKeyTemplate pointQueryKeyTemplate; // Serialized mysql Field, this could avoid serialize mysql field each time sendFields. // Since, serialize fields is too heavy when table is wide Map serializedFields = Maps.newHashMap(); @@ -118,18 +97,13 @@ public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery) throws public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, StatementContext statementContext) throws TException { - this(planner, analzyedQuery, statementContext, + this(planner, analzyedQuery, statementContext == null ? null : statementContext.getSecurityDependencyContext()); } @VisibleForTesting public ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, SecurityDependencyContext securityDependencyContext) throws TException { - this(planner, analzyedQuery, null, securityDependencyContext); - } - - private ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, - StatementContext statementContext, SecurityDependencyContext securityDependencyContext) throws TException { this.planner = planner; this.serializedDescTable = ByteString.copyFrom( new TSerializer().serialize(DescriptorToThriftConverter.toThrift(planner.getDescTable()))); @@ -160,19 +134,11 @@ private ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.tableName = this.scanNode.getTableNameInPlan(); this.schemaVersion = this.tbl.getBaseSchemaVersion(); this.partitionTopologyVersion = this.tbl.getPartitionTopologyVersion(); - this.tableNamespaceSnapshot = TableNamespaceSnapshot.from(this.tbl); this.analzyedQuery = analzyedQuery; - this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(this.scanNode, statementContext); this.securityDependencyContext = securityDependencyContext == null ? null : securityDependencyContext.snapshotForShortCircuit(); } - @VisibleForTesting - ShortCircuitQueryContext(OlapTable tbl, String tableName, int schemaVersion, - long fileCacheQueryLimitBytes) { - this(tbl, tableName, schemaVersion, fileCacheQueryLimitBytes, null); - } - @VisibleForTesting ShortCircuitQueryContext(OlapTable tbl, String tableName, int schemaVersion, long fileCacheQueryLimitBytes, SecurityDependencyContext securityDependencyContext) { @@ -186,74 +152,19 @@ private ShortCircuitQueryContext(Planner planner, Queriable analzyedQuery, this.schemaVersion = schemaVersion; this.fileCacheQueryLimitBytes = fileCacheQueryLimitBytes; this.partitionTopologyVersion = tbl.getPartitionTopologyVersion(); - this.tableNamespaceSnapshot = TableNamespaceSnapshot.from(tbl); this.scanNode = null; this.analzyedQuery = null; - this.pointQueryKeyTemplate = PointQueryKeyTemplate.unsupported(); this.securityDependencyContext = securityDependencyContext; } - @VisibleForTesting - ShortCircuitQueryContext(OlapScanNode scanNode, StatementContext statementContext) { - this.planner = null; - this.serializedDescTable = ByteString.EMPTY; - this.serializedOutputExpr = ByteString.EMPTY; - this.serializedQueryOptions = ByteString.EMPTY; - this.cacheID = UUID.randomUUID(); - this.scanNode = scanNode; - this.tbl = scanNode.getOlapTable(); - this.tableName = scanNode.getTableNameInPlan(); - this.schemaVersion = tbl.getBaseSchemaVersion(); - this.fileCacheQueryLimitBytes = -1; - this.partitionTopologyVersion = tbl.getPartitionTopologyVersion(); - this.tableNamespaceSnapshot = TableNamespaceSnapshot.from(tbl); - this.analzyedQuery = null; - this.pointQueryKeyTemplate = PointQueryKeyTemplate.create(scanNode, statementContext); - this.securityDependencyContext = null; - } - public boolean isReusable(ConnectContext ctx) { return !this.tbl.isDropped && this.tbl.getBaseSchemaVersion() == this.schemaVersion && Objects.equals(this.tableName, this.tbl.getName()) && this.fileCacheQueryLimitBytes == ctx.getSessionVariable().fileCacheQueryLimitBytes && this.tbl.getPartitionTopologyVersion() == this.partitionTopologyVersion - && this.tableNamespaceSnapshot.matches(this.tbl) - && (securityDependencyContext == null || securityDependencyContext.isValid(ctx)); - } - - /** Fence name-scoped grants when a catalog or database is renamed or replaced. */ - private static class TableNamespaceSnapshot { - private final DatabaseIf database; - private final CatalogIf catalog; - private final String databaseName; - private final String catalogName; - - private TableNamespaceSnapshot(DatabaseIf database, CatalogIf catalog, - String databaseName, String catalogName) { - this.database = database; - this.catalog = catalog; - this.databaseName = databaseName; - this.catalogName = catalogName; - } - - private static TableNamespaceSnapshot from(OlapTable table) { - DatabaseIf database = table.getDatabase(); - CatalogIf catalog = database == null ? null : database.getCatalog(); - return new TableNamespaceSnapshot(database, catalog, - database == null ? null : database.getFullName(), - catalog == null ? null : catalog.getName()); - } - - private boolean matches(OlapTable table) { - DatabaseIf currentDatabase = table.getDatabase(); - CatalogIf currentCatalog = currentDatabase == null ? null : currentDatabase.getCatalog(); - return currentDatabase == database - && currentCatalog == catalog - && Objects.equals(databaseName, - currentDatabase == null ? null : currentDatabase.getFullName()) - && Objects.equals(catalogName, currentCatalog == null ? null : currentCatalog.getName()); - } + && securityDependencyContext != null + && securityDependencyContext.isValid(ctx); } public void sanitize() { @@ -263,247 +174,4 @@ public void sanitize() { Preconditions.checkNotNull(tbl); Preconditions.checkNotNull(tableName); } - - /** Build state owned by one execution without modifying the cached plan or scan conjuncts. */ - public PointQueryExecutionContext createPointQueryExecutionContext(StatementContext statementContext) { - return pointQueryKeyTemplate.bind(statementContext); - } - - private static class PointQueryKeyTemplate { - private final List keyColumns; - private final List placeholderBindings; - private final List> fixedConstraints; - private final boolean complete; - - private PointQueryKeyTemplate(List keyColumns, - List placeholderBindings, - List> fixedConstraints, boolean complete) { - this.keyColumns = Collections.unmodifiableList(new ArrayList<>(keyColumns)); - this.placeholderBindings = Collections.unmodifiableList(new ArrayList<>(placeholderBindings)); - List> immutableConstraints = new ArrayList<>(fixedConstraints.size()); - for (List constraints : fixedConstraints) { - immutableConstraints.add(Collections.unmodifiableList(new ArrayList<>(constraints))); - } - this.fixedConstraints = Collections.unmodifiableList(immutableConstraints); - this.complete = complete; - } - - private static PointQueryKeyTemplate unsupported() { - return new PointQueryKeyTemplate(Collections.emptyList(), Collections.emptyList(), - Collections.emptyList(), false); - } - - private static PointQueryKeyTemplate create(OlapScanNode scanNode, StatementContext statementContext) { - if (statementContext == null) { - return unsupported(); - } - List keyColumns = scanNode.getOlapTable().getBaseSchemaKeyColumns(); - if (keyColumns.isEmpty()) { - return new PointQueryKeyTemplate(keyColumns, Collections.emptyList(), - Collections.emptyList(), true); - } - if (!statementContext.arePointQueryFixedKeyConstraintsComplete()) { - return unsupported(); - } - - Map keyOrdinals = new TreeMap<>(String.CASE_INSENSITIVE_ORDER); - List> fixedConstraints = new ArrayList<>(keyColumns.size()); - for (int ordinal = 0; ordinal < keyColumns.size(); ordinal++) { - keyOrdinals.put(keyColumns.get(ordinal).getName(), ordinal); - fixedConstraints.add(new ArrayList<>()); - } - - List placeholderBindings = new ArrayList<>(); - for (Map.Entry entry - : statementContext.getIdToComparisonSlot().entrySet()) { - SlotReference slot = entry.getValue(); - if (!slot.getOriginalColumn().isPresent()) { - return unsupported(); - } - Integer ordinal = keyOrdinals.get(slot.getOriginalColumn().get().getName()); - if (ordinal == null) { - return unsupported(); - } - placeholderBindings.add(new PlaceholderKeyBinding(entry.getKey(), ordinal, slot)); - } - - List placeholders = statementContext.getPlaceholders(); - if (placeholderBindings.size() != placeholders.size()) { - return unsupported(); - } - for (Placeholder placeholder : placeholders) { - if (!statementContext.getIdToComparisonSlot().containsKey(placeholder.getPlaceholderId())) { - return unsupported(); - } - } - - for (PointQueryFixedKeyConstraint constraint - : statementContext.getPointQueryFixedKeyConstraints()) { - SlotReference slot = constraint.getSlot(); - if (!slot.getOriginalColumn().isPresent()) { - return unsupported(); - } - Integer ordinal = keyOrdinals.get(slot.getOriginalColumn().get().getName()); - if (ordinal != null) { - fixedConstraints.get(ordinal).add(constraint.getLiteral()); - } else if (!Column.DELETE_SIGN.equals(slot.getOriginalColumn().get().getName())) { - return unsupported(); - } - } - - boolean[] covered = new boolean[keyColumns.size()]; - for (PlaceholderKeyBinding binding : placeholderBindings) { - covered[binding.keyOrdinal] = true; - } - for (int ordinal = 0; ordinal < fixedConstraints.size(); ordinal++) { - covered[ordinal] |= !fixedConstraints.get(ordinal).isEmpty(); - } - for (boolean keyCovered : covered) { - if (!keyCovered) { - return unsupported(); - } - } - return new PointQueryKeyTemplate(keyColumns, placeholderBindings, fixedConstraints, true); - } - - private PointQueryExecutionContext bind(StatementContext statementContext) { - if (!complete || statementContext == null) { - return PointQueryExecutionContext.fallback(); - } - List> valuesByKey = new ArrayList<>(fixedConstraints.size()); - for (List constraints : fixedConstraints) { - valuesByKey.add(new ArrayList<>(constraints)); - } - for (PlaceholderKeyBinding binding : placeholderBindings) { - Expression value = statementContext.getIdToPlaceholderRealExpr().get(binding.placeholderId); - if (!(value instanceof Literal)) { - return PointQueryExecutionContext.fallback(); - } - Literal typedValue = coerceComparisonLiteral(binding.slot, (Literal) value); - if (typedValue == null) { - return PointQueryExecutionContext.fallback(); - } - if (typedValue instanceof NullLiteral) { - return PointQueryExecutionContext.empty(); - } - valuesByKey.get(binding.keyOrdinal).add(typedValue); - } - - Map keyValues = new LinkedHashMap<>(); - for (int ordinal = 0; ordinal < keyColumns.size(); ordinal++) { - List values = valuesByKey.get(ordinal); - if (values.isEmpty()) { - return PointQueryExecutionContext.fallback(); - } - Literal representative = values.get(0); - if (representative instanceof NullLiteral) { - return PointQueryExecutionContext.empty(); - } - for (int i = 1; i < values.size(); i++) { - Boolean equal = sqlEquals(representative, values.get(i)); - if (equal == null) { - return PointQueryExecutionContext.fallback(); - } - if (!equal) { - return PointQueryExecutionContext.empty(); - } - } - LiteralExpr physicalValue = toPhysicalKeyLiteral(representative, keyColumns.get(ordinal)); - if (physicalValue == null) { - return PointQueryExecutionContext.fallback(); - } - keyValues.put(keyColumns.get(ordinal).getName(), physicalValue); - } - return PointQueryExecutionContext.lookup(keyValues); - } - - private static Literal coerceComparisonLiteral(SlotReference slot, Literal value) { - try { - Expression comparison = TypeCoercionUtils.processComparisonPredicate(new EqualTo(slot, value)); - Expression comparisonSlot = comparison.child(0); - // A cast on the physical key can change equality semantics (for example INT 1 - // compared with string '01'). Normal planning must evaluate such comparisons. - return comparisonSlot instanceof SlotReference && comparison.child(1) instanceof Literal - ? (Literal) comparison.child(1) : null; - } catch (Exception e) { - return null; - } - } - - private static Boolean sqlEquals(Literal left, Literal right) { - if (left instanceof NullLiteral || right instanceof NullLiteral) { - return false; - } - try { - Expression comparison = TypeCoercionUtils.processComparisonPredicate(new EqualTo(left, right)); - Expression result = FoldConstantRuleOnFE.evaluateWithoutContext(comparison); - return result instanceof BooleanLiteral ? ((BooleanLiteral) result).getValue() : null; - } catch (Exception e) { - return null; - } - } - - private static LiteralExpr toPhysicalKeyLiteral(Literal literal, Column column) { - try { - LiteralExpr legacyLiteral = literal.toLegacyLiteral(); - Type columnType = column.getType(); - if (!columnType.equals(legacyLiteral.getType()) - && !columnType.matchesType(legacyLiteral.getType())) { - legacyLiteral = LiteralExprUtils.createLiteral(legacyLiteral.getStringValue(), columnType); - } - return legacyLiteral; - } catch (Exception e) { - return null; - } - } - } - - private static class PlaceholderKeyBinding { - private final PlaceholderId placeholderId; - private final int keyOrdinal; - private final SlotReference slot; - - private PlaceholderKeyBinding(PlaceholderId placeholderId, int keyOrdinal, SlotReference slot) { - this.placeholderId = placeholderId; - this.keyOrdinal = keyOrdinal; - this.slot = slot; - } - } - - /** Immutable outcome and typed key tuple for exactly one point-query execution. */ - public static class PointQueryExecutionContext { - public enum Decision { - LOOKUP, - EMPTY, - FALLBACK - } - - private final Decision decision; - private final Map keyValues; - - private PointQueryExecutionContext(Decision decision, Map keyValues) { - this.decision = decision; - this.keyValues = Collections.unmodifiableMap(new LinkedHashMap<>(keyValues)); - } - - public static PointQueryExecutionContext lookup(Map keyValues) { - return new PointQueryExecutionContext(Decision.LOOKUP, keyValues); - } - - public static PointQueryExecutionContext empty() { - return new PointQueryExecutionContext(Decision.EMPTY, Collections.emptyMap()); - } - - public static PointQueryExecutionContext fallback() { - return new PointQueryExecutionContext(Decision.FALLBACK, Collections.emptyMap()); - } - - public Decision getDecision() { - return decision; - } - - public Map getKeyValues() { - return keyValues; - } - } } 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 7e7e4a83718cc5..cede11185681e5 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 @@ -1541,35 +1541,17 @@ public void executeAndSendResult(boolean isOutfileQuery, boolean isSendFields, // ExecuteCommand publishes this same context after a successful first prepared execution. statementContext.setShortCircuitQueryContext(shortCircuitQueryContext); } - ShortCircuitQueryContext.PointQueryExecutionContext pointQueryExecutionContext = - statementContext.getPointQueryExecutionContext(); - if (pointQueryExecutionContext == null) { - pointQueryExecutionContext = shortCircuitQueryContext - .createPointQueryExecutionContext(statementContext); - statementContext.setPointQueryExecutionContext(pointQueryExecutionContext); - } - if (pointQueryExecutionContext.getDecision() - == ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK) { - // The physical plan is still a valid normal plan. If an execution value cannot be - // safely reduced to an exact typed key, use the Coordinator instead of failing the - // statement or guessing a lookup key. - statementContext.setShortCircuitQuery(false); - statementContext.setShortCircuitQueryContext(null); - } else { - coordBase = new PointQueryExecutor(shortCircuitQueryContext, pointQueryExecutionContext, - context.getSessionVariable().getMaxMsgSizeOfResultReceiver()); - context.getState().setIsQuery(true); - } - } - if (coordBase == null - && planner instanceof NereidsPlanner && ((NereidsPlanner) planner).getDistributedPlans() != null) { + coordBase = new PointQueryExecutor(shortCircuitQueryContext, + context.getSessionVariable().getMaxMsgSizeOfResultReceiver()); + context.getState().setIsQuery(true); + } else if (planner instanceof NereidsPlanner && ((NereidsPlanner) planner).getDistributedPlans() != null) { coord = new NereidsCoordinator(context, (NereidsPlanner) planner, context.getStatsErrorEstimator()); profile.addExecutionProfile(coord.getExecutionProfile()); QeProcessorImpl.INSTANCE.registerQuery(context.queryId(), new QueryInfo(context, originStmt.originStmt, coord)); coordBase = coord; - } else if (coordBase == null) { + } else { coord = EnvFactory.getInstance().createCoordinator( context, planner, context.getStatsErrorEstimator()); profile.addExecutionProfile(coord.getExecutionProfile()); diff --git a/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java b/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java index c5abcf7e70c229..993973a905d8f3 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/mysql/privilege/AuthTest.java @@ -87,17 +87,4 @@ public void testCheckDbPrivWithSessionMappedRoleForTempUser() throws Exception { } } - @Test - public void testAuthorizationVersionAdvancesOnMutation() throws Exception { - Auth auth = Env.getCurrentEnv().getAuth(); - long version = auth.getAuthorizationVersion(); - - addUser("authorization_version_user", true); - - Assertions.assertTrue(auth.getAuthorizationVersion() > version); - long changedVersion = auth.getAuthorizationVersion(); - auth.refreshUserPrivEntriesByResovledIPs(Collections.emptyMap()); - Assertions.assertEquals(changedVersion, auth.getAuthorizationVersion()); - } - } diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java index 2c1b65778667e2..6b2dca601acc50 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/SecurityDependencyContextTest.java @@ -23,10 +23,12 @@ import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.Env; import org.apache.doris.catalog.OlapTable; +import org.apache.doris.common.AnalysisException; import org.apache.doris.datasource.CatalogIf; import org.apache.doris.mysql.privilege.AccessControllerManager; import org.apache.doris.mysql.privilege.Auth; import org.apache.doris.mysql.privilege.InternalAuthorizationPlugin; +import org.apache.doris.mysql.privilege.PrivPredicate; import org.apache.doris.policy.PolicyMgr; import org.apache.doris.qe.ConnectContext; import org.apache.doris.qe.SessionVariable; @@ -48,20 +50,39 @@ public class SecurityDependencyContextTest { private static final String COLUMN = "value"; @Test - public void testBuiltInVersionsAllowConstantTimeReuse() { + public void testCurrentPrivilegeAndNoPolicyAllowReuse() throws Exception { BuiltInFixture fixture = new BuiltInFixture(); SecurityDependencyContext snapshot = fixture.completeDependencies().snapshotForShortCircuit(); + Mockito.clearInvocations(fixture.accessManager, fixture.policyMgr); Assertions.assertTrue(snapshot.isValid(fixture.connectContext)); - Mockito.verify(fixture.accessManager, Mockito.never()).evalRowFilterPolicies( - Mockito.any(), Mockito.anyString(), Mockito.anyString(), Mockito.anyString()); - Mockito.verify(fixture.accessManager, Mockito.never()).evalDataMaskPolicies( - Mockito.any(), Mockito.anyString(), Mockito.anyString(), Mockito.anyString(), Mockito.anySet()); + Mockito.verify(fixture.policyMgr).hasRowPolicy(CATALOG, DATABASE, TABLE); + Mockito.verify(fixture.accessManager).checkColumnsPriv( + fixture.connectContext, CATALOG, DATABASE, TABLE, + ImmutableSet.of(COLUMN), PrivPredicate.SELECT); + } + + @Test + public void testPolicyAddedAfterPlanningInvalidatesReuseBeforePrivilegeCheck() throws Exception { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext snapshot = fixture.completeDependencies().snapshotForShortCircuit(); + Mockito.when(fixture.policyMgr.hasRowPolicy(CATALOG, DATABASE, TABLE)).thenReturn(true); + Mockito.clearInvocations(fixture.accessManager); - Mockito.when(fixture.auth.getAuthorizationVersion()).thenReturn(8L); Assertions.assertFalse(snapshot.isValid(fixture.connectContext)); - Mockito.when(fixture.auth.getAuthorizationVersion()).thenReturn(7L); - Mockito.when(fixture.policyMgr.getRowPolicyVersion()).thenReturn(12L); + Mockito.verify(fixture.accessManager, Mockito.never()).checkColumnsPriv( + Mockito.any(ConnectContext.class), Mockito.anyString(), Mockito.anyString(), Mockito.anyString(), + Mockito.anySet(), Mockito.any()); + } + + @Test + public void testSelectRevocationInvalidatesReuse() throws Exception { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext snapshot = fixture.completeDependencies().snapshotForShortCircuit(); + Mockito.doThrow(new AnalysisException("denied")).when(fixture.accessManager).checkColumnsPriv( + fixture.connectContext, CATALOG, DATABASE, TABLE, + ImmutableSet.of(COLUMN), PrivPredicate.SELECT); + Assertions.assertFalse(snapshot.isValid(fixture.connectContext)); } @@ -84,23 +105,35 @@ public void testDifferentAuthenticatedRolesInvalidateReuse() { } @Test - public void testRowPolicyDisablesShortCircuitReuse() { + public void testEffectiveExternalRowPolicyFailsClosed() { BuiltInFixture fixture = new BuiltInFixture(); SecurityDependencyContext dependencies = fixture.completeDependencies(); - dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, + dependencies.addRowPolicies( ImmutableList.of(RowFilterSpec.restrictive("row:1", "tenant_id = 1"))); - Assertions.assertTrue(dependencies.hasRowPolicy()); + Assertions.assertTrue(dependencies.hasEffectiveRowPolicy()); + Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(fixture.connectContext)); + } + + @Test + public void testNamespaceChangeBeforeSnapshotFailsClosed() { + BuiltInFixture fixture = new BuiltInFixture(); + SecurityDependencyContext dependencies = fixture.completeDependencies(); + Mockito.when(fixture.database.getFullName()).thenReturn("renamed_db"); + Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(fixture.connectContext)); + Mockito.verify(fixture.policyMgr, Mockito.never()).hasRowPolicy(Mockito.anyString(), + Mockito.anyString(), Mockito.anyString()); } @Test - public void testDataMaskDisablesVersionOnlyReuse() { + public void testDataMaskDisablesDirectReuse() { BuiltInFixture fixture = new BuiltInFixture(); SecurityDependencyContext dependencies = fixture.completeDependencies(); dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.of(new DataMaskSpec("mask:1", "null"))); + Assertions.assertTrue(dependencies.hasDataMask()); Assertions.assertFalse(dependencies.snapshotForShortCircuit().isValid(fixture.connectContext)); } @@ -129,12 +162,9 @@ private BuiltInFixture() { Mockito.when(database.getFullName()).thenReturn(DATABASE); Mockito.when(table.getDatabase()).thenReturn((DatabaseIf) database); Mockito.when(table.getName()).thenReturn(TABLE); - Mockito.when(auth.getAuthorizationVersion()).thenReturn(7L); - Mockito.when(auth.isAuthorizationVersionReliable()).thenReturn(true); - Mockito.when(policyMgr.getRowPolicyVersion()).thenReturn(11L); + Mockito.when(policyMgr.hasRowPolicy(CATALOG, DATABASE, TABLE)).thenReturn(false); Mockito.when(accessManager.getAccessControllerOrDefault(CATALOG)) .thenReturn(new InternalAuthorizationPlugin(auth)); - Mockito.when(env.getAuth()).thenReturn(auth); Mockito.when(env.getPolicyMgr()).thenReturn(policyMgr); Mockito.when(env.getAccessManager()).thenReturn(accessManager); Mockito.when(connectContext.getCurrentUserIdentity()).thenReturn(USER); @@ -146,7 +176,6 @@ private BuiltInFixture() { private SecurityDependencyContext completeDependencies() { SecurityDependencyContext dependencies = new SecurityDependencyContext(connectContext); dependencies.addCheckedPrivilege(table, ImmutableSet.of(COLUMN)); - dependencies.setRowPolicies(CATALOG, DATABASE, TABLE, ImmutableList.of()); dependencies.addDataMask(CATALOG, DATABASE, TABLE, COLUMN, Optional.empty()); return dependencies; } diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java index cb5e839c56405f..f02266fcbb3a50 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/ShortCircuitPointQueryTest.java @@ -79,6 +79,7 @@ protected void runBeforeAll() throws Exception { + " \"store_row_column\" = \"true\"\n" + ");"); createView("CREATE VIEW `view_point_query` AS SELECT `key`, `v1` FROM `tbl_point_query`"); + executeSql("CREATE USER point_query_policy_user IDENTIFIED BY 'Point_query_policy_123!'"); } @Test @@ -142,6 +143,32 @@ void testPointQueryWithManualTabletDoesNotUseShortCircuit() throws Exception { Assertions.assertFalse(connectContext.getStatementContext().isShortCircuitQuery()); } + @Test + void testInjectiveCastUsesShortCircuit() { + rewrite("select * from tbl_point_query where cast(`key` as bigint) = 1"); + + Assertions.assertTrue(connectContext.getStatementContext().isShortCircuitQuery()); + } + + @Test + void testDuplicateKeyPredicatesDoNotUseShortCircuit() { + rewrite("select * from tbl_point_query where `key` = 1 and `key` = 2"); + + Assertions.assertFalse(connectContext.getStatementContext().isShortCircuitQuery()); + } + + @Test + void testAnyRowPolicyDisablesShortCircuitAtPlanning() throws Exception { + createPolicy("CREATE ROW POLICY point_query_policy ON test.tbl_point_query " + + "AS RESTRICTIVE TO point_query_policy_user USING (`key` = 1)"); + try { + rewrite("select * from tbl_point_query where `key` = 1"); + Assertions.assertFalse(connectContext.getStatementContext().isShortCircuitQuery()); + } finally { + dropPolicy("DROP ROW POLICY point_query_policy ON test.tbl_point_query"); + } + } + @Test void testViewDoesNotUseShortCircuit() { rewrite("select * from view_point_query where `key` = 1"); diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java index 07a90354ec21ef..2b4c3ebf8ec31e 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommandTest.java @@ -42,7 +42,6 @@ import org.apache.doris.qe.PreparedStatementContext; import org.apache.doris.qe.SessionVariable; import org.apache.doris.qe.ShortCircuitQueryContext; -import org.apache.doris.qe.ShortCircuitQueryContext.PointQueryExecutionContext; import org.apache.doris.qe.StmtExecutor; import org.apache.doris.thrift.TQueryOptions; @@ -266,9 +265,7 @@ public void testFastPathInstallsCachedShortCircuitContextAcrossExecutions() thro PreparedStatementContext preparedStatement = new PreparedStatementContext( prepareCommand, connectContext, statementContext, "stmt"); - // Keep the real point-key binding path, but explicitly model a cache whose security - // dependencies have already been validated. A bare StatementContext intentionally - // cannot produce a reusable security snapshot because production validation is fail-closed. + // Explicitly model a cache whose security dependencies have already been validated. Planner planner = Mockito.mock(Planner.class); Mockito.when(planner.getQueryOptions()).thenReturn(new TQueryOptions()); DescriptorTable descriptorTable = new DescriptorTable(); @@ -278,7 +275,6 @@ public void testFastPathInstallsCachedShortCircuitContextAcrossExecutions() thro OlapTable table = Mockito.spy(new OlapTable()); Mockito.doReturn("tbl").when(table).getName(); Mockito.doReturn(10).when(table).getBaseSchemaVersion(); - Mockito.doReturn(Collections.emptyList()).when(table).getBaseSchemaKeyColumns(); Mockito.when(scanNode.getPointQueryProjectList()).thenReturn(Collections.emptyList()); Mockito.when(scanNode.getOlapTable()).thenReturn(table); Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); @@ -311,45 +307,6 @@ public void testFastPathInstallsCachedShortCircuitContextAcrossExecutions() thro Mockito.any(), Mockito.any(), Mockito.any(), Mockito.any()); } - @Test - public void testUnsafePointKeyFallsBackToNormalPreparedExecution() throws Exception { - String sql = "select 1"; - LogicalPlan logicalPlan = new NereidsParser().parseSingle(sql); - - ConnectContext connectContext = Mockito.mock(ConnectContext.class); - StatementContext statementContext = new StatementContext(); - statementContext.setShortCircuitQuery(true); - PrepareCommand prepareCommand = new PrepareCommand( - "stmt", logicalPlan, Collections.emptyList(), new OriginStatement(sql, 0)); - PreparedStatementContext preparedStatement = new PreparedStatementContext( - prepareCommand, connectContext, statementContext, "stmt"); - ShortCircuitQueryContext cachedPlan = Mockito.mock(ShortCircuitQueryContext.class); - Mockito.when(cachedPlan.isReusable(connectContext)).thenReturn(true); - Mockito.when(cachedPlan.createPointQueryExecutionContext(Mockito.any(StatementContext.class))) - .thenReturn(PointQueryExecutionContext.fallback()); - preparedStatement.shortCircuitQueryContext = Optional.of(cachedPlan); - - StmtExecutor executor = Mockito.mock(StmtExecutor.class); - Mockito.when(connectContext.getPreparedStementContext("stmt")).thenReturn(preparedStatement); - SessionVariable sessionVariable = new SessionVariable(); - sessionVariable.enableGroupCommitFullPrepare = false; - Mockito.when(connectContext.getSessionVariable()).thenReturn(sessionVariable); - Mockito.when(connectContext.getStatementContext()).thenReturn(statementContext); - Mockito.when(executor.getContext()).thenReturn(connectContext); - - new ExecuteCommand("stmt", prepareCommand, statementContext).run(connectContext, executor); - - Mockito.verify(executor).execute(); - Mockito.verify(executor, Mockito.never()).executeAndSendResult(Mockito.anyBoolean(), Mockito.anyBoolean(), - Mockito.any(), Mockito.any(), Mockito.any(), Mockito.any()); - Assertions.assertFalse(preparedStatement.shortCircuitQueryContext.isPresent(), - "an unsafe typed key must discard direct reuse and run the normal planner"); - Assertions.assertFalse(preparedStatement.getStatementContext().isShortCircuitQuery(), - "normal planning must not inherit the cached execution's short-circuit flag"); - Assertions.assertNull(preparedStatement.getStatementContext().getShortCircuitQueryContext()); - Assertions.assertNull(preparedStatement.getStatementContext().getPointQueryExecutionContext()); - } - @Test public void testInvalidSecurityDependenciesRefreshInsteadOfDirectReuse() throws Exception { String sql = "select * from tbl"; diff --git a/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java b/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java index 27fdab19e50836..99ad76556f165e 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/policy/PolicyTest.java @@ -189,17 +189,16 @@ public void testDropPolicy() throws Exception { } @Test - public void testRowPolicyVersionAdvancesOnMutation() throws Exception { + public void testHasRowPolicy() throws Exception { PolicyMgr policyMgr = Env.getCurrentEnv().getPolicyMgr(); - long version = policyMgr.getRowPolicyVersion(); + Assertions.assertFalse(policyMgr.hasRowPolicy("internal", "test", "table1")); - createPolicy("CREATE ROW POLICY test_row_policy_version ON test.table1 AS PERMISSIVE" + createPolicy("CREATE ROW POLICY test_has_row_policy ON test.table1 AS PERMISSIVE" + " TO test_policy USING (k1 = 1)"); - long createdVersion = policyMgr.getRowPolicyVersion(); - Assertions.assertTrue(createdVersion > version); + Assertions.assertTrue(policyMgr.hasRowPolicy("internal", "test", "table1")); - dropPolicy("DROP ROW POLICY test_row_policy_version ON test.table1"); - Assertions.assertTrue(policyMgr.getRowPolicyVersion() > createdVersion); + dropPolicy("DROP ROW POLICY test_has_row_policy ON test.table1"); + Assertions.assertFalse(policyMgr.hasRowPolicy("internal", "test", "table1")); } @Test diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java index e0cef9e3302d96..fcf3bc20c6e220 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/PointQueryExecutorTest.java @@ -17,8 +17,6 @@ package org.apache.doris.qe; -import org.apache.doris.catalog.OlapTable; -import org.apache.doris.nereids.StatementContext; import org.apache.doris.planner.OlapScanNode; import org.junit.jupiter.api.Assertions; @@ -36,21 +34,4 @@ public void testCandidateBackendsShuffleDependsOnQuerySelectionOrder() { Mockito.when(scanNode.isScanBackendOrderBySelection()).thenReturn(true); Assertions.assertFalse(PointQueryExecutor.shouldShuffleCandidateBackends(scanNode)); } - - @Test - public void testEmptyDecisionReturnsBeforeTabletPruning() throws Exception { - OlapTable table = Mockito.mock(OlapTable.class); - Mockito.when(table.getBaseSchemaKeyColumns()).thenReturn(java.util.Collections.emptyList()); - OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); - Mockito.when(scanNode.getOlapTable()).thenReturn(table); - Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); - ShortCircuitQueryContext queryContext = new ShortCircuitQueryContext(scanNode, new StatementContext()); - PointQueryExecutor executor = new PointQueryExecutor(queryContext, - ShortCircuitQueryContext.PointQueryExecutionContext.empty(), 1024); - - Mockito.clearInvocations(scanNode); - Assertions.assertNotNull(executor.getNext()); - // lazyEvaluateRangeLocations is the first operation that can resolve a tablet and lead to a BE RPC. - Mockito.verifyNoInteractions(scanNode); - } } diff --git a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java index 491558a4ff13de..88416ac0e93702 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/qe/ShortCircuitQueryContextTest.java @@ -20,7 +20,6 @@ import org.apache.doris.analysis.DescriptorTable; import org.apache.doris.analysis.Queriable; import org.apache.doris.catalog.Column; -import org.apache.doris.catalog.DatabaseIf; import org.apache.doris.catalog.KeysType; import org.apache.doris.catalog.MaterializedIndex; import org.apache.doris.catalog.OlapTable; @@ -28,16 +27,7 @@ import org.apache.doris.catalog.PrimitiveType; import org.apache.doris.catalog.RandomDistributionInfo; import org.apache.doris.catalog.SinglePartitionInfo; -import org.apache.doris.datasource.CatalogIf; import org.apache.doris.nereids.SecurityDependencyContext; -import org.apache.doris.nereids.StatementContext; -import org.apache.doris.nereids.trees.expressions.Placeholder; -import org.apache.doris.nereids.trees.expressions.SlotReference; -import org.apache.doris.nereids.trees.expressions.StatementScopeIdGenerator; -import org.apache.doris.nereids.trees.expressions.literal.DecimalLiteral; -import org.apache.doris.nereids.trees.expressions.literal.IntegerLiteral; -import org.apache.doris.nereids.trees.expressions.literal.NullLiteral; -import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.planner.OlapScanNode; import org.apache.doris.planner.Planner; import org.apache.doris.thrift.TQueryOptions; @@ -48,10 +38,8 @@ import org.junit.jupiter.api.Test; import org.mockito.Mockito; -import java.math.BigDecimal; import java.util.Collections; import java.util.List; -import java.util.concurrent.atomic.AtomicReference; public class ShortCircuitQueryContextTest { private OlapTable table(String name, int schemaVersion) { @@ -61,14 +49,6 @@ private OlapTable table(String name, int schemaVersion) { return table; } - private OlapTable pointQueryTable(List keyColumns) { - OlapTable table = Mockito.mock(OlapTable.class); - Mockito.when(table.getName()).thenReturn("tbl"); - Mockito.when(table.getBaseSchemaKeyColumns()).thenReturn(keyColumns); - Mockito.when(table.getBaseSchemaVersion()).thenReturn(1); - return table; - } - private ConnectContext connectContext(long fileCacheQueryLimitBytes) { ConnectContext ctx = new ConnectContext(); SessionVariable sessionVariable = new SessionVariable(); @@ -77,10 +57,17 @@ private ConnectContext connectContext(long fileCacheQueryLimitBytes) { return ctx; } + private SecurityDependencyContext validSecurityDependencies() { + SecurityDependencyContext dependencies = Mockito.mock(SecurityDependencyContext.class); + Mockito.when(dependencies.isValid(Mockito.any())).thenReturn(true); + return dependencies; + } + @Test public void testReusableRequiresSameFileCacheQueryLimitBytes() { ShortCircuitQueryContext context = - new ShortCircuitQueryContext(table("tbl", 10), "tbl", 10, -1); + new ShortCircuitQueryContext(table("tbl", 10), "tbl", 10, -1, + validSecurityDependencies()); Assertions.assertTrue(context.isReusable(connectContext(-1))); Assertions.assertFalse(context.isReusable(connectContext(0))); @@ -89,29 +76,12 @@ public void testReusableRequiresSameFileCacheQueryLimitBytes() { @Test public void testReusableStillChecksTableMetadata() { ShortCircuitQueryContext context = - new ShortCircuitQueryContext(table("tbl", 11), "tbl", 10, 0); + new ShortCircuitQueryContext(table("tbl", 11), "tbl", 10, 0, + validSecurityDependencies()); Assertions.assertFalse(context.isReusable(connectContext(0))); } - @Test - @SuppressWarnings({"rawtypes", "unchecked"}) - public void testReusableRequiresSameDatabaseNamespace() { - CatalogIf catalog = Mockito.mock(CatalogIf.class); - DatabaseIf database = Mockito.mock(DatabaseIf.class); - AtomicReference databaseName = new AtomicReference<>("old_db"); - Mockito.when(catalog.getName()).thenReturn("internal"); - Mockito.when(database.getCatalog()).thenReturn(catalog); - Mockito.when(database.getFullName()).thenAnswer(ignored -> databaseName.get()); - OlapTable table = table("tbl", 10); - Mockito.doReturn(database).when(table).getDatabase(); - ShortCircuitQueryContext context = new ShortCircuitQueryContext(table, "tbl", 10, -1); - - Assertions.assertTrue(context.isReusable(connectContext(-1))); - databaseName.set("new_db"); - Assertions.assertFalse(context.isReusable(connectContext(-1))); - } - @Test public void testReusableRequiresSamePartitionTopologyVersion() { long baseIndexId = 2L; @@ -123,7 +93,8 @@ public void testReusableRequiresSamePartitionTopologyVersion() { table.setIndexMeta(baseIndexId, "tbl", baseSchema, 10, 0, (short) 1, TStorageType.COLUMN, KeysType.DUP_KEYS); table.setBaseIndexId(baseIndexId); - ShortCircuitQueryContext context = new ShortCircuitQueryContext(table, "tbl", 10, -1); + ShortCircuitQueryContext context = new ShortCircuitQueryContext( + table, "tbl", 10, -1, validSecurityDependencies()); Assertions.assertTrue(context.isReusable(connectContext(-1))); table.addPartition(new Partition(3L, "p1", @@ -144,6 +115,14 @@ public void testReusableRequiresCurrentSecurityDependencies() { Mockito.verify(securityDependencyContext).isValid(connectContext); } + @Test + public void testMissingSecurityDependenciesFailClosed() { + ShortCircuitQueryContext context = new ShortCircuitQueryContext( + table("tbl", 10), "tbl", 10, -1, null); + + Assertions.assertFalse(context.isReusable(connectContext(-1))); + } + @Test public void testSerializedQueryOptionsKeepBitmapOpCountVersion() throws Exception { TQueryOptions queryOptions = new SessionVariable().toThrift(); @@ -169,171 +148,4 @@ public void testSerializedQueryOptionsKeepBitmapOpCountVersion() throws Exceptio Assertions.assertTrue(serializedQueryOptions.isSetNewVersionBitmapOpCount()); Assertions.assertTrue(serializedQueryOptions.isNewVersionBitmapOpCount()); } - - @Test - public void testPreparedKeyTemplateKeepsFixedConstraintsAcrossExecutions() { - Column parameterKey = new Column("parameter_key", PrimitiveType.INT); - parameterKey.setIsKey(true); - Column fixedKey = new Column("fixed_key", PrimitiveType.INT); - fixedKey.setIsKey(true); - List schema = List.of(parameterKey, fixedKey); - OlapTable table = pointQueryTable(schema); - SlotReference parameterSlot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, parameterKey, Collections.emptyList()); - SlotReference fixedSlot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, fixedKey, Collections.emptyList()); - - PlaceholderId placeholderId = new PlaceholderId(0); - StatementContext templateContext = new StatementContext(); - templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); - templateContext.getIdToComparisonSlot().put(placeholderId, parameterSlot); - // Fixed statement predicates remain distinct from caller-controlled placeholders. - templateContext.addPointQueryFixedKeyConstraint(parameterSlot, new IntegerLiteral(1)); - templateContext.addPointQueryFixedKeyConstraint(fixedSlot, new IntegerLiteral(9)); - - OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); - Mockito.when(scanNode.getOlapTable()).thenReturn(table); - Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); - ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode, templateContext); - - StatementContext first = execution(placeholderId, new IntegerLiteral(1)); - ShortCircuitQueryContext.PointQueryExecutionContext firstExecution = - cached.createPointQueryExecutionContext(first); - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, - firstExecution.getDecision()); - Assertions.assertEquals("1", firstExecution.getKeyValues().get("parameter_key").getStringValue()); - Assertions.assertEquals("9", firstExecution.getKeyValues().get("fixed_key").getStringValue()); - - StatementContext second = execution(placeholderId, new IntegerLiteral(2)); - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.EMPTY, - cached.createPointQueryExecutionContext(second).getDecision()); - - // Reusing the same prepared handle with 1 -> 2 -> 1 must not contaminate the template. - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, - cached.createPointQueryExecutionContext(first).getDecision()); - - StatementContext nullValue = execution(placeholderId, new NullLiteral()); - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.EMPTY, - cached.createPointQueryExecutionContext(nullValue).getDecision()); - Mockito.verify(scanNode, Mockito.never()).getConjuncts(); - } - - @Test - public void testFixedOnlyKeyTemplate() { - Column key = new Column("k", PrimitiveType.INT); - key.setIsKey(true); - OlapTable table = pointQueryTable(Collections.singletonList(key)); - SlotReference slot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); - StatementContext templateContext = new StatementContext(); - templateContext.addPointQueryFixedKeyConstraint(slot, new IntegerLiteral(7)); - ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode(table), templateContext); - - ShortCircuitQueryContext.PointQueryExecutionContext execution = - cached.createPointQueryExecutionContext(new StatementContext()); - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, - execution.getDecision()); - Assertions.assertEquals("7", execution.getKeyValues().get("k").getStringValue()); - } - - @Test - public void testPlaceholderOnlyKeyTemplate() { - Column key = new Column("k", PrimitiveType.INT); - key.setIsKey(true); - OlapTable table = pointQueryTable(Collections.singletonList(key)); - SlotReference slot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); - PlaceholderId placeholderId = new PlaceholderId(0); - StatementContext templateContext = new StatementContext(); - templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); - templateContext.getIdToComparisonSlot().put(placeholderId, slot); - ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode(table), templateContext); - - ShortCircuitQueryContext.PointQueryExecutionContext execution = - cached.createPointQueryExecutionContext(execution(placeholderId, new IntegerLiteral(8))); - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.LOOKUP, - execution.getDecision()); - Assertions.assertEquals("8", execution.getKeyValues().get("k").getStringValue()); - } - - @Test - public void testInexactPhysicalKeyFallsBack() { - Column key = new Column("k", PrimitiveType.INT); - key.setIsKey(true); - OlapTable table = pointQueryTable(Collections.singletonList(key)); - SlotReference slot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); - PlaceholderId placeholderId = new PlaceholderId(0); - StatementContext templateContext = new StatementContext(); - templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); - templateContext.getIdToComparisonSlot().put(placeholderId, slot); - OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); - Mockito.when(scanNode.getOlapTable()).thenReturn(table); - Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); - ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode, templateContext); - - StatementContext execution = execution(placeholderId, new DecimalLiteral(new BigDecimal("1.2"))); - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK, - cached.createPointQueryExecutionContext(execution).getDecision()); - } - - @Test - public void testNonSlotFixedConstraintFallsBack() { - Column key = new Column("k", PrimitiveType.INT); - key.setIsKey(true); - OlapTable table = pointQueryTable(Collections.singletonList(key)); - SlotReference slot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); - PlaceholderId placeholderId = new PlaceholderId(0); - StatementContext templateContext = new StatementContext(); - templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); - templateContext.getIdToComparisonSlot().put(placeholderId, slot); - // ExpressionAnalyzer uses this marker for a fixed predicate such as - // CAST(k AS CHAR(1)) = '1', whose cast cannot identify an exact physical key. - templateContext.markPointQueryFixedKeyConstraintsIncomplete(); - OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); - Mockito.when(scanNode.getOlapTable()).thenReturn(table); - Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); - ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode, templateContext); - - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK, - cached.createPointQueryExecutionContext( - execution(placeholderId, new IntegerLiteral(1))).getDecision()); - } - - @Test - public void testFixedNonKeyConstraintFallsBack() { - Column key = new Column("k", PrimitiveType.INT); - key.setIsKey(true); - Column value = new Column("v", PrimitiveType.INT); - OlapTable table = pointQueryTable(Collections.singletonList(key)); - SlotReference keySlot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, key, Collections.emptyList()); - SlotReference valueSlot = SlotReference.fromColumn( - StatementScopeIdGenerator.newExprId(), table, value, Collections.emptyList()); - PlaceholderId placeholderId = new PlaceholderId(0); - StatementContext templateContext = new StatementContext(); - templateContext.setPlaceholders(Collections.singletonList(new Placeholder(placeholderId))); - templateContext.getIdToComparisonSlot().put(placeholderId, keySlot); - templateContext.addPointQueryFixedKeyConstraint(valueSlot, new IntegerLiteral(1)); - ShortCircuitQueryContext cached = new ShortCircuitQueryContext(scanNode(table), templateContext); - - Assertions.assertEquals(ShortCircuitQueryContext.PointQueryExecutionContext.Decision.FALLBACK, - cached.createPointQueryExecutionContext( - execution(placeholderId, new IntegerLiteral(1))).getDecision()); - } - - private OlapScanNode scanNode(OlapTable table) { - OlapScanNode scanNode = Mockito.mock(OlapScanNode.class); - Mockito.when(scanNode.getOlapTable()).thenReturn(table); - Mockito.when(scanNode.getTableNameInPlan()).thenReturn("tbl"); - return scanNode; - } - - private StatementContext execution(PlaceholderId placeholderId, - org.apache.doris.nereids.trees.expressions.Expression value) { - StatementContext context = new StatementContext(); - context.getIdToPlaceholderRealExpr().put(placeholderId, value); - return context; - } } diff --git a/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy index edb024979cee27..517cb056ab2efa 100644 --- a/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy +++ b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy @@ -94,88 +94,51 @@ suite("prepared_point_query_row_policy", "p0") { assertTrue(explainRows.toString().contains("SHORT-CIRCUIT")) } - // Without a policy, a fixed statement predicate can share a key with a placeholder. The - // immutable key template must preserve that fixed value across every execution. - connect(user, password, url) { - def prepared = prepareStatement """ - SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + // Ambiguous key predicates and non-injective casts are outside the supported point-query shape. + connect(user, password, explainUrl) { + def duplicateKey = sql """ + EXPLAIN SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value FROM prepared_point_query_row_policy - WHERE tenant_id = ? AND tenant_id = 1 AND item_id = ? + WHERE tenant_id = 1 AND tenant_id = 2 AND item_id = 10 """ - assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) - - def readRows = { Integer tenant, int item -> - if (tenant == null) { - prepared.setNull(1, java.sql.Types.INTEGER) - } else { - prepared.setInt(1, tenant) - } - prepared.setInt(2, item) - def rows = [] - prepared.executeQuery().withCloseable { result -> - assertEquals(3, result.getMetaData().getColumnCount()) - assertEquals("tenant_id", result.getMetaData().getColumnLabel(1)) - assertEquals("item_id", result.getMetaData().getColumnLabel(2)) - assertEquals("value", result.getMetaData().getColumnLabel(3)) - while (result.next()) { - rows.add([result.getInt(1), result.getInt(2), result.getString(3)]) - } - } - return rows - } - - assertEquals([[1, 10, "allowed"]], readRows(1, 10)) - assertEquals([], readRows(2, 10)) - assertEquals([[1, 10, "allowed"]], readRows(1, 10)) - assertEquals([], readRows(null, 10)) - - // A non-integral parameter cannot be represented by the physical INT lookup key. It is - // evaluated by the normal planner, exercising direct-reuse FALLBACK without an error. - prepared.setBigDecimal(1, new BigDecimal("1.2")) - prepared.setInt(2, 10) - def fallbackRows = [] - prepared.executeQuery().withCloseable { result -> - while (result.next()) { - fallbackRows.add([result.getInt(1), result.getInt(2), result.getString(3)]) - } - } - assertEquals([], fallbackRows) - prepared.close() - } - - // A cast row policy also keeps the statement on the normal path. - sql """ - CREATE ROW POLICY ${policyName} ON ${dbName}.prepared_point_query_row_policy - AS RESTRICTIVE TO ${user} USING (CAST(tenant_id AS CHAR(1)) = '1') - """ - sql "SYNC" - connect(user, password, explainUrl) { - def explainRows = sql """ + assertFalse(duplicateKey.toString().contains("SHORT-CIRCUIT")) + def nonInjectiveCast = sql """ EXPLAIN SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value - FROM prepared_point_query_row_policy WHERE tenant_id = 1 AND item_id = 10 + FROM prepared_point_query_row_policy + WHERE CAST(tenant_id AS CHAR(1)) = '1' AND item_id = 10 """ - assertFalse(explainRows.toString().contains("SHORT-CIRCUIT")) + assertFalse(nonInjectiveCast.toString().contains("SHORT-CIRCUIT")) } + // A placeholder outside the final filter cannot be rebound in cached output expressions. + // Keep this shape on the normal path and verify both lookup and projection values across executions. connect(user, password, url) { def prepared = prepareStatement """ - SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id = ? AS matches_parameter FROM prepared_point_query_row_policy WHERE tenant_id = ? AND item_id = ? """ assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) + + prepared.setInt(1, 2) + prepared.setInt(2, 1) + prepared.setInt(3, 10) + prepared.executeQuery().withCloseable { result -> + assertTrue(result.next()) + assertFalse(result.getBoolean(1)) + assertFalse(result.next()) + } + prepared.setInt(1, 10) prepared.setInt(2, 10) - def rows = [] + prepared.setInt(3, 10) prepared.executeQuery().withCloseable { result -> - while (result.next()) { - rows.add([result.getInt(1), result.getInt(2), result.getString(3)]) - } + assertTrue(result.next()) + assertTrue(result.getBoolean(1)) + assertFalse(result.next()) } prepared.close() - assertEquals([[10, 10, "cast-match"]], rows) } - sql "DROP ROW POLICY IF EXISTS ${policyName} ON ${dbName}.prepared_point_query_row_policy FOR ${user}" sql "DROP USER IF EXISTS ${user}" } From 4ccff1315ce1ad8ce8c89dd84e0750e42ecc8342 Mon Sep 17 00:00:00 2001 From: morrySnow Date: Tue, 15 Sep 2026 17:20:21 +0800 Subject: [PATCH 6/6] [fix](point query) Scope placeholder tracking to filters --- .../doris/nereids/StatementContext.java | 9 ----- .../rules/analysis/ExpressionAnalyzer.java | 39 ++++++++++++------- ...calResultSinkToShortCircuitPointQuery.java | 12 ++++-- .../trees/plans/commands/ExecuteCommand.java | 10 ----- .../prepared_point_query_row_policy.groovy | 28 +++++++++++++ 5 files changed, 63 insertions(+), 35 deletions(-) diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java index 2b75cf43557ed0..a8caf670852c8b 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java @@ -184,7 +184,6 @@ public enum TableFrom { private final Set viewDdlSqlSet = Sets.newHashSet(); private final SqlCacheContext sqlCacheContext; private final SecurityDependencyContext securityDependencyContext; - private boolean hasNonFilterPlaceholder; // generate for next id for prepared statement's placeholders, which is // connection level @@ -699,14 +698,6 @@ public SecurityDependencyContext getSecurityDependencyContext() { return securityDependencyContext; } - public boolean hasNonFilterPlaceholder() { - return hasNonFilterPlaceholder; - } - - public void setHasNonFilterPlaceholder(boolean hasNonFilterPlaceholder) { - this.hasNonFilterPlaceholder = hasNonFilterPlaceholder; - } - public boolean isDpHyp() { return isDpHyp; } 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 f8fd0910ae12db..41a70d8d8c5af6 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 @@ -23,7 +23,6 @@ import org.apache.doris.common.DdlException; import org.apache.doris.common.Pair; import org.apache.doris.common.util.Util; -import org.apache.doris.mysql.MysqlCommand; import org.apache.doris.nereids.CascadesContext; import org.apache.doris.nereids.SqlCacheContext; import org.apache.doris.nereids.StatementContext; @@ -92,6 +91,7 @@ import org.apache.doris.nereids.trees.expressions.typecoercion.ImplicitCastInputTypes; import org.apache.doris.nereids.trees.plans.PlaceholderId; import org.apache.doris.nereids.trees.plans.Plan; +import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalJoin; import org.apache.doris.nereids.trees.plans.logical.LogicalPlan; import org.apache.doris.nereids.types.ArrayType; @@ -920,21 +920,34 @@ public Expression visitPlaceholder(Placeholder placeholder, ExpressionRewriteCon return visit(realExpr, context); } - // Register prepared statement placeholder id to related slot in comparison predicate. - // Used to replace expression in ShortCircuit plan + // Register point-query filter placeholders so cached conjuncts can be rebound on EXECUTE. + // Restrict this registry to LogicalFilter: placeholders in projections or other plan nodes + // cannot be updated by the short-circuit executor and must keep the statement on the normal path. private void registerPlaceholderIdToSlot(ComparisonPredicate cp, ExpressionRewriteContext context, Expression left, Expression right) { - if (ConnectContext.get() != null - && ConnectContext.get().getCommand() == MysqlCommand.COM_STMT_EXECUTE) { - // Used to replace expression in ShortCircuit plan - if (cp.right() instanceof Placeholder && left instanceof SlotReference) { - PlaceholderId id = ((Placeholder) cp.right()).getPlaceholderId(); - context.cascadesContext.getStatementContext().getIdToComparisonSlot().put(id, (SlotReference) left); - } else if (cp.left() instanceof Placeholder && right instanceof SlotReference) { - PlaceholderId id = ((Placeholder) cp.left()).getPlaceholderId(); - context.cascadesContext.getStatementContext().getIdToComparisonSlot().put(id, (SlotReference) right); - } + if (context == null || !(currentPlan instanceof LogicalFilter)) { + return; + } + SlotReference leftSlot = extractInjectiveCastSlot(left); + SlotReference rightSlot = extractInjectiveCastSlot(right); + if (cp.right() instanceof Placeholder && leftSlot != null) { + PlaceholderId id = ((Placeholder) cp.right()).getPlaceholderId(); + context.cascadesContext.getStatementContext().getIdToComparisonSlot().put(id, leftSlot); + } else if (cp.left() instanceof Placeholder && rightSlot != null) { + PlaceholderId id = ((Placeholder) cp.left()).getPlaceholderId(); + context.cascadesContext.getStatementContext().getIdToComparisonSlot().put(id, rightSlot); + } + } + + private SlotReference extractInjectiveCastSlot(Expression expression) { + if (expression instanceof SlotReference) { + return (SlotReference) expression; + } + if (expression instanceof Cast && expression.child(0) instanceof SlotReference + && expression.child(0).getDataType().isInjectiveCastTo(expression.getDataType())) { + return (SlotReference) expression.child(0); } + return null; } @Override diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java index cfe01c1299ee1b..0a266acc4eea6e 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/LogicalResultSinkToShortCircuitPointQuery.java @@ -82,6 +82,12 @@ private boolean hasRowPolicy(OlapTable table, StatementContext statementContext) } } + private boolean allPlaceholdersBoundByFilter(StatementContext statementContext) { + return statementContext.getPlaceholders().stream() + .allMatch(placeholder -> statementContext.getIdToComparisonSlot() + .containsKey(placeholder.getPlaceholderId())); + } + @VisibleForTesting boolean scanMatchShortCircuitCondition(LogicalOlapScan olapScan) { ConnectContext connectContext = ConnectContext.get(); @@ -122,12 +128,12 @@ boolean scanMatchShortCircuitCondition(LogicalOlapScan olapScan) { // set short circuit flag and return the original plan private Plan shortCircuit(Plan root, OlapTable olapTable, Set conjuncts, StatementContext statementContext) { - // Keep policy-bearing tables, inlined views, and placeholders outside the final filter on - // the normal path. A cached no-policy plan repeats the table-level lookup before reuse. + // Keep policy-bearing tables, inlined views, and placeholders that the final filter cannot + // rebind on the normal path. A cached no-policy plan repeats the table-level lookup before reuse. if (hasRowPolicy(olapTable, statementContext) || statementContext.getSecurityDependencyContext().hasEffectiveRowPolicy() || statementContext.getSecurityDependencyContext().hasDataMask() - || statementContext.hasNonFilterPlaceholder() + || !allPlaceholdersBoundByFilter(statementContext) || !statementContext.getViewDdlSqls().isEmpty()) { return root; } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java index d9f558e10f0b40..e7342b38eb11ca 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/commands/ExecuteCommand.java @@ -35,7 +35,6 @@ import org.apache.doris.nereids.trees.plans.commands.insert.InsertOverwriteTableCommand; import org.apache.doris.nereids.trees.plans.commands.insert.OlapGroupCommitInsertExecutor; import org.apache.doris.nereids.trees.plans.commands.merge.MergeIntoCommand; -import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalPlan; import org.apache.doris.nereids.trees.plans.logical.LogicalSqlCache; import org.apache.doris.nereids.trees.plans.visitor.PlanVisitor; @@ -118,7 +117,6 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { } // Commands hide their retained query trees from normal plan traversal. Reset every exposed // root so a later EXECUTE cannot reuse a relation-local snapshot from an earlier execution. - boolean hasNonFilterPlaceholder = false; for (int rootIndex = 0; rootIndex < relationRoots.size(); rootIndex++) { LogicalPlan relationRoot = relationRoots.get(rootIndex); for (UnboundRelation relation : relationRoot.collectToList( @@ -130,10 +128,6 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { } for (LogicalPlan plan : relationRoot.collectToList(node -> true)) { for (Expression expression : plan.getExpressions()) { - if (!(plan instanceof LogicalFilter) - && expression.anyMatch(Placeholder.class::isInstance)) { - hasNonFilterPlaceholder = true; - } for (SubqueryExpr subquery : expression.collectToList( SubqueryExpr.class::isInstance)) { // SubqueryExpr owns its query plan outside Plan.children(), so retained prepared @@ -143,10 +137,6 @@ public void run(ConnectContext ctx, StmtExecutor executor) throws Exception { } } } - statementContext.setHasNonFilterPlaceholder(hasNonFilterPlaceholder); - if (hasNonFilterPlaceholder) { - statementContext.setShortCircuitQuery(false); - } if (logicalPlan instanceof LogicalSqlCache) { throw new AnalysisException("Unsupported sql cache for server prepared statement"); } diff --git a/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy index 517cb056ab2efa..1e967b5be66941 100644 --- a/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy +++ b/regression-test/suites/prepared_stmt_p0/prepared_point_query_row_policy.groovy @@ -140,5 +140,33 @@ suite("prepared_point_query_row_policy", "p0") { prepared.close() } + // Injective casts on key columns remain eligible, and their filter placeholders must be + // rebound when the cached point-query plan is reused. + connect(user, password, url) { + def prepared = prepareStatement """ + SELECT /*+ SET_VAR(enable_short_circuit_query=true) */ tenant_id, item_id, value + FROM prepared_point_query_row_policy + WHERE CAST(tenant_id AS BIGINT) = ? AND item_id = ? + """ + assertEquals(com.mysql.cj.jdbc.ServerPreparedStatement, prepared.class) + + prepared.setLong(1, 1) + prepared.setInt(2, 10) + prepared.executeQuery().withCloseable { result -> + assertTrue(result.next()) + assertEquals("allowed", result.getString(3)) + assertFalse(result.next()) + } + + prepared.setLong(1, 10) + prepared.setInt(2, 10) + prepared.executeQuery().withCloseable { result -> + assertTrue(result.next()) + assertEquals("cast-match", result.getString(3)) + assertFalse(result.next()) + } + prepared.close() + } + sql "DROP USER IF EXISTS ${user}" }