From c02e2871440d7a3d8dbb2928ee724984d4518a9b Mon Sep 17 00:00:00 2001 From: Claire Date: Thu, 16 Jul 2026 14:51:36 +0800 Subject: [PATCH 1/8] support OPENQUERY function --- .../assignment/EncryptOpenQueryUtils.java | 119 ++++++++++++++++++ .../token/pojo/EncryptOpenQuerySQLToken.java | 44 +++++++ .../assignment/EncryptOpenQueryUtilsTest.java | 75 +++++++++++ 3 files changed, 238 insertions(+) create mode 100644 features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java create mode 100644 features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java create mode 100644 features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java new file mode 100644 index 0000000000000..8b737e699bb53 --- /dev/null +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java @@ -0,0 +1,119 @@ +/* + * 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.shardingsphere.encrypt.rewrite.token.generator.assignment; + +import lombok.AccessLevel; +import lombok.NoArgsConstructor; +import org.apache.shardingsphere.database.connector.core.metadata.database.enums.QuoteCharacter; +import org.apache.shardingsphere.encrypt.rule.EncryptRule; +import org.apache.shardingsphere.encrypt.rule.table.EncryptTable; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.ExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.FunctionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.FunctionTableSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableSegment; + +import java.util.Optional; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +/** + * Encrypt OPENQUERY utility. + */ +@NoArgsConstructor(access = AccessLevel.PRIVATE) +public final class EncryptOpenQueryUtils { + + private static final Pattern FROM_TABLE_PATTERN = Pattern.compile("\\bFROM\\s+([^\\s,;]+)", Pattern.CASE_INSENSITIVE); + + /** + * Whether table segment is OPENQUERY function table. + * + * @param tableSegment table segment + * @return whether OPENQUERY function table + */ + public static boolean isOpenQueryFunctionTable(final TableSegment tableSegment) { + return tableSegment instanceof FunctionTableSegment && ((FunctionTableSegment) tableSegment).getTableFunction() instanceof FunctionSegment + && "OPENQUERY".equalsIgnoreCase(((FunctionSegment) ((FunctionTableSegment) tableSegment).getTableFunction()).getFunctionName()); + } + + /** + * Find OPENQUERY SQL literal. + * + * @param tableSegment table segment + * @return OPENQUERY SQL literal + */ + public static Optional findOpenQuerySQLLiteral(final TableSegment tableSegment) { + if (!isOpenQueryFunctionTable(tableSegment)) { + return Optional.empty(); + } + int parameterIndex = 0; + for (ExpressionSegment each : ((FunctionSegment) ((FunctionTableSegment) tableSegment).getTableFunction()).getParameters()) { + if (1 == parameterIndex) { + return each instanceof LiteralExpressionSegment ? Optional.of((LiteralExpressionSegment) each) : Optional.empty(); + } + parameterIndex++; + } + return Optional.empty(); + } + + /** + * Find encrypt table from OPENQUERY target. + * + * @param rule encrypt rule + * @param tableSegment table segment + * @return encrypt table + */ + public static Optional findEncryptTable(final EncryptRule rule, final TableSegment tableSegment) { + Optional openQuerySQL = findOpenQuerySQLLiteral(tableSegment); + if (!openQuerySQL.isPresent()) { + return Optional.empty(); + } + Matcher matcher = FROM_TABLE_PATTERN.matcher(openQuerySQL.get().getText()); + if (!matcher.find()) { + return Optional.empty(); + } + String actualTableName = QuoteCharacter.unwrapAndTrimText(matcher.group(1).substring(matcher.group(1).lastIndexOf('.') + 1)); + for (String each : rule.getAllTableNames()) { + Optional encryptTable = rule.findEncryptTable(each); + if (encryptTable.isPresent() && each.equalsIgnoreCase(actualTableName)) { + return encryptTable; + } + } + return Optional.empty(); + } + + /** + * Find schema name from OPENQUERY SQL. + * + * @param tableSegment table segment + * @return schema name + */ + public static Optional findSchemaName(final TableSegment tableSegment) { + Optional openQuerySQL = findOpenQuerySQLLiteral(tableSegment); + if (!openQuerySQL.isPresent()) { + return Optional.empty(); + } + Matcher matcher = FROM_TABLE_PATTERN.matcher(openQuerySQL.get().getText()); + if (!matcher.find()) { + return Optional.empty(); + } + String tableExpression = matcher.group(1); + int delimiterIndex = tableExpression.lastIndexOf('.'); + return -1 == delimiterIndex ? Optional.empty() : Optional.of(QuoteCharacter.unwrapAndTrimText(tableExpression.substring(0, delimiterIndex))); + } +} diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java new file mode 100644 index 0000000000000..11f2dc659295d --- /dev/null +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java @@ -0,0 +1,44 @@ +/* + * 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.shardingsphere.encrypt.rewrite.token.pojo; + +import lombok.Getter; +import org.apache.shardingsphere.infra.rewrite.sql.token.common.pojo.SQLToken; +import org.apache.shardingsphere.infra.rewrite.sql.token.common.pojo.Substitutable; + +/** + * OPENQUERY SQL token for encrypt. + */ +public final class EncryptOpenQuerySQLToken extends SQLToken implements Substitutable { + + @Getter + private final int stopIndex; + + private final String sql; + + public EncryptOpenQuerySQLToken(final int startIndex, final int stopIndex, final String sql) { + super(startIndex); + this.stopIndex = stopIndex; + this.sql = sql; + } + + @Override + public String toString() { + return "'" + sql + "'"; + } +} diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java new file mode 100644 index 0000000000000..40af9f1f5f51b --- /dev/null +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java @@ -0,0 +1,75 @@ +/* + * 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.shardingsphere.encrypt.rewrite.token.generator.assignment; + +import org.apache.shardingsphere.encrypt.rule.EncryptRule; +import org.apache.shardingsphere.encrypt.rule.table.EncryptTable; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.column.ColumnSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.FunctionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.FunctionTableSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.SimpleTableSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableNameSegment; +import org.apache.shardingsphere.sql.parser.statement.core.value.identifier.IdentifierValue; +import org.junit.jupiter.api.Test; + +import java.util.Collections; +import java.util.Optional; + +import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.is; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +class EncryptOpenQueryUtilsTest { + + @Test + void assertIsOpenQueryFunctionTable() { + assertTrue(EncryptOpenQueryUtils.isOpenQueryFunctionTable(createOpenQueryTableSegment("SELECT foo_col FROM foo_schema.foo_tbl"))); + assertFalse(EncryptOpenQueryUtils.isOpenQueryFunctionTable(new SimpleTableSegment(new TableNameSegment(0, 0, new IdentifierValue("foo_tbl"))))); + } + + @Test + void assertFindEncryptTable() { + EncryptRule rule = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + when(rule.getAllTableNames()).thenReturn(Collections.singleton("foo_tbl")); + when(rule.findEncryptTable("foo_tbl")).thenReturn(Optional.of(encryptTable)); + Optional actual = EncryptOpenQueryUtils.findEncryptTable(rule, createOpenQueryTableSegment("SELECT foo_col FROM foo_schema.foo_tbl")); + assertTrue(actual.isPresent()); + assertThat(actual.get(), is(encryptTable)); + assertFalse(EncryptOpenQueryUtils.findEncryptTable(rule, createOpenQueryTableSegment("SELECT foo_col FROM foo_schema.bar_tbl")).isPresent()); + } + + @Test + void assertFindSchemaName() { + Optional actual = EncryptOpenQueryUtils.findSchemaName(createOpenQueryTableSegment("SELECT foo_col FROM foo_schema.foo_tbl")); + assertTrue(actual.isPresent()); + assertThat(actual.get(), is("foo_schema")); + assertFalse(EncryptOpenQueryUtils.findSchemaName(createOpenQueryTableSegment("SELECT foo_col FROM foo_tbl")).isPresent()); + } + + private FunctionTableSegment createOpenQueryTableSegment(final String openQuerySQL) { + FunctionSegment functionSegment = new FunctionSegment(0, 0, "OPENQUERY", "OPENQUERY (foo_server, '" + openQuerySQL + "')"); + functionSegment.getParameters().add(new ColumnSegment(0, 0, new IdentifierValue("foo_server"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(0, 0, openQuerySQL)); + return new FunctionTableSegment(0, 0, functionSegment); + } +} From 3b05a6bad03c6bd04358b27bb13d7a1bd74ee141 Mon Sep 17 00:00:00 2001 From: Claire Date: Thu, 16 Jul 2026 14:51:46 +0800 Subject: [PATCH 2/8] support OPENQUERY function --- .../EncryptSQLRewriteContextDecorator.java | 14 ++- .../EncryptAssignmentTokenGenerator.java | 119 +++++++++++++----- ...EncryptUpdateAssignmentTokenGenerator.java | 5 +- ...EncryptSQLRewriteContextDecoratorTest.java | 42 ++++++- .../EncryptAssignmentTokenGeneratorTest.java | 46 +++++++ ...yptUpdateAssignmentTokenGeneratorTest.java | 31 +++-- .../engine/scenario/EncryptSQLRewriterIT.java | 12 ++ .../query-with-cipher/dml/update/update.xml | 5 + .../encrypt/config/query-with-cipher.yaml | 6 + 9 files changed, 226 insertions(+), 54 deletions(-) diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecorator.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecorator.java index e84fd65dddf0e..50bc94a746ddf 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecorator.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecorator.java @@ -22,12 +22,14 @@ import org.apache.shardingsphere.encrypt.rewrite.condition.EncryptConditionEngine; import org.apache.shardingsphere.encrypt.rewrite.parameter.EncryptParameterRewritersRegistry; import org.apache.shardingsphere.encrypt.rewrite.token.EncryptTokenGenerateBuilder; +import org.apache.shardingsphere.encrypt.rewrite.token.generator.assignment.EncryptOpenQueryUtils; import org.apache.shardingsphere.encrypt.rule.EncryptRule; import org.apache.shardingsphere.infra.annotation.HighFrequencyInvocation; import org.apache.shardingsphere.infra.binder.context.available.WhereContextAvailable; import org.apache.shardingsphere.infra.binder.context.extractor.SQLStatementContextExtractor; import org.apache.shardingsphere.infra.binder.context.statement.SQLStatementContext; import org.apache.shardingsphere.infra.binder.context.statement.type.dml.SelectStatementContext; +import org.apache.shardingsphere.infra.binder.context.statement.type.dml.UpdateStatementContext; import org.apache.shardingsphere.infra.config.props.ConfigurationProperties; import org.apache.shardingsphere.infra.rewrite.context.SQLRewriteContext; import org.apache.shardingsphere.infra.rewrite.context.SQLRewriteContextDecorator; @@ -37,6 +39,7 @@ import org.apache.shardingsphere.infra.route.context.RouteContext; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.predicate.WhereSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.SimpleTableSegment; +import org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.UpdateStatement; import java.util.Collection; import java.util.Collections; @@ -69,7 +72,16 @@ private boolean containsEncryptTable(final EncryptRule rule, final SQLStatementC return true; } } - return false; + return containsOpenQueryEncryptTable(rule, sqlStatementContext); + } + + private boolean containsOpenQueryEncryptTable(final EncryptRule rule, final SQLStatementContext sqlStatementContext) { + if (!(sqlStatementContext instanceof UpdateStatementContext)) { + return false; + } + UpdateStatement updateStatement = ((UpdateStatementContext) sqlStatementContext).getSqlStatement(); + return EncryptOpenQueryUtils.isOpenQueryFunctionTable(updateStatement.getTable()) + && EncryptOpenQueryUtils.findEncryptTable(rule, updateStatement.getTable()).isPresent(); } private Collection createEncryptConditions(final EncryptRule rule, final SQLStatementContext sqlStatementContext) { diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java index 25e949454a423..d5716fd879a36 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java @@ -17,14 +17,15 @@ package org.apache.shardingsphere.encrypt.rewrite.token.generator.assignment; -import lombok.AllArgsConstructor; -import lombok.extern.slf4j.Slf4j; +import lombok.Getter; +import lombok.RequiredArgsConstructor; import org.apache.shardingsphere.database.connector.core.metadata.database.enums.QuoteCharacter; import org.apache.shardingsphere.database.connector.core.type.DatabaseType; import org.apache.shardingsphere.database.connector.core.type.DatabaseTypeRegistry; import org.apache.shardingsphere.encrypt.enums.EncryptDerivedColumnSuffix; import org.apache.shardingsphere.encrypt.rewrite.token.pojo.EncryptAssignmentToken; import org.apache.shardingsphere.encrypt.rewrite.token.pojo.EncryptLiteralAssignmentToken; +import org.apache.shardingsphere.encrypt.rewrite.token.pojo.EncryptOpenQuerySQLToken; import org.apache.shardingsphere.encrypt.rewrite.token.pojo.EncryptParameterAssignmentToken; import org.apache.shardingsphere.encrypt.rule.EncryptRule; import org.apache.shardingsphere.encrypt.rule.column.EncryptColumn; @@ -40,19 +41,21 @@ import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.ExpressionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.ParameterMarkerExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableSegment; import java.util.Collection; import java.util.Collections; import java.util.LinkedList; import java.util.List; import java.util.Optional; +import java.util.regex.Matcher; +import java.util.regex.Pattern; /** * Assignment generator for encrypt. */ -@Slf4j @HighFrequencyInvocation -@AllArgsConstructor +@RequiredArgsConstructor public final class EncryptAssignmentTokenGenerator { private final EncryptRule rule; @@ -69,53 +72,102 @@ public final class EncryptAssignmentTokenGenerator { * @return generated SQL tokens */ public Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment) { + return generateSQLTokens(tablesContext, setAssignmentSegment, Optional.empty()); + } + + /** + * Generate SQL tokens. + * + * @param tablesContext SQL statement context + * @param setAssignmentSegment set assignment segment + * @param targetTable target table segment + * @return generated SQL tokens + */ + public Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final TableSegment targetTable) { + return generateSQLTokens(tablesContext, setAssignmentSegment, Optional.of(targetTable)); + } + + private Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final Optional targetTable) { Collection result = new LinkedList<>(); - DatabaseTypeRegistry databaseTypeRegistry = new DatabaseTypeRegistry(databaseType); - String schemaName = tablesContext.getSchemaName().orElseGet(() -> databaseTypeRegistry.getDefaultSchemaName(database.getName())); - QuoteCharacter quoteCharacter = databaseTypeRegistry.getDialectDatabaseMetaData().getQuoteCharacter(); for (ColumnAssignmentSegment each : setAssignmentSegment.getAssignments()) { ColumnSegment assignedColumn = getAssignedColumn(each); - findEncryptTable(assignedColumn).ifPresent(encryptTable -> { - String columnName = assignedColumn.getIdentifier().getValue(); - if (encryptTable.isEncryptColumn(columnName)) { - result.addAll(generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptTable.getEncryptColumn(columnName), each, quoteCharacter)); - } - }); + Optional openQueryContext = findOpenQueryContext(targetTable, assignedColumn.getIdentifier().getValue()); + findEncryptTable(assignedColumn, openQueryContext) + .ifPresent(encryptTable -> appendEncryptAssignmentTokens(result, tablesContext, targetTable, each, assignedColumn, openQueryContext, encryptTable)); } return result; } + private void appendEncryptAssignmentTokens(final Collection result, final TablesContext tablesContext, final Optional targetTable, + final ColumnAssignmentSegment assignmentSegment, final ColumnSegment assignedColumn, + final Optional openQueryContext, final EncryptTable encryptTable) { + String columnName = assignedColumn.getIdentifier().getValue(); + if (!encryptTable.isEncryptColumn(columnName)) { + return; + } + EncryptColumn encryptColumn = encryptTable.getEncryptColumn(columnName); + DatabaseTypeRegistry databaseTypeRegistry = new DatabaseTypeRegistry(databaseType); + String schemaName = targetTable.flatMap(EncryptOpenQueryUtils::findSchemaName) + .orElseGet(() -> tablesContext.getSchemaName().orElseGet(() -> databaseTypeRegistry.getDefaultSchemaName(database.getName()))); + QuoteCharacter quoteCharacter = databaseTypeRegistry.getDialectDatabaseMetaData().getQuoteCharacter(); + result.addAll(generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptColumn, assignmentSegment, quoteCharacter, openQueryContext.isPresent())); + openQueryContext.ifPresent(optional -> result.add(generateOpenQuerySQLToken(optional, encryptColumn))); + } + + private Optional findEncryptTable(final ColumnSegment assignedColumn, final Optional openQueryContext) { + Optional result = rule.findEncryptTable(assignedColumn.getColumnBoundInfo().getOriginalTable().getValue()); + return result.isPresent() ? result : openQueryContext.map(OpenQueryContext::getEncryptTable); + } + + private Optional findOpenQueryContext(final Optional targetTable, final String columnName) { + if (!targetTable.isPresent()) { + return Optional.empty(); + } + Optional openQuerySQL = EncryptOpenQueryUtils.findOpenQuerySQLLiteral(targetTable.get()); + Optional encryptTable = EncryptOpenQueryUtils.findEncryptTable(rule, targetTable.get()).filter(optional -> optional.isEncryptColumn(columnName)); + return openQuerySQL.isPresent() && encryptTable.isPresent() ? Optional.of(new OpenQueryContext(openQuerySQL.get(), encryptTable.get())) : Optional.empty(); + } + + private EncryptOpenQuerySQLToken generateOpenQuerySQLToken(final OpenQueryContext openQueryContext, final EncryptColumn encryptColumn) { + String openQuerySQL = openQueryContext.getOpenQuerySQL().getText(); + String rewrittenSQL = Pattern.compile(String.join("", "\\b", Pattern.quote(encryptColumn.getName()), "\\b"), Pattern.CASE_INSENSITIVE).matcher(openQuerySQL) + .replaceAll(Matcher.quoteReplacement(encryptColumn.getCipher().getName())); + return new EncryptOpenQuerySQLToken(openQueryContext.getOpenQuerySQL().getStartIndex(), openQueryContext.getOpenQuerySQL().getStopIndex(), rewrittenSQL); + } + private Collection generateAssignmentSQLTokens(final String schemaName, final String tableName, final EncryptColumn encryptColumn, - final ColumnAssignmentSegment segment, final QuoteCharacter quoteCharacter) { + final ColumnAssignmentSegment segment, final QuoteCharacter quoteCharacter, final boolean useActualColumnName) { ExpressionSegment value = segment.getValue(); if (value instanceof ParameterMarkerExpressionSegment) { - return Collections.singleton(generateParameterSQLToken(encryptColumn, segment, quoteCharacter)); + return Collections.singleton(generateParameterSQLToken(encryptColumn, segment, quoteCharacter, useActualColumnName)); } if (value instanceof LiteralExpressionSegment) { - return Collections.singleton(generateLiteralSQLToken(schemaName, tableName, encryptColumn, segment, quoteCharacter)); + return Collections.singleton(generateLiteralSQLToken(schemaName, tableName, encryptColumn, segment, quoteCharacter, useActualColumnName)); } return Collections.emptyList(); } - private EncryptAssignmentToken generateParameterSQLToken(final EncryptColumn encryptColumn, final ColumnAssignmentSegment segment, final QuoteCharacter quoteCharacter) { + private EncryptAssignmentToken generateParameterSQLToken(final EncryptColumn encryptColumn, final ColumnAssignmentSegment segment, final QuoteCharacter quoteCharacter, + final boolean useActualColumnName) { ColumnSegment leftColumn = getAssignedColumn(segment); EncryptParameterAssignmentToken result = new EncryptParameterAssignmentToken(leftColumn.getStartIndex(), segment.getStopIndex(), quoteCharacter); - appendEncryptColumnTokens(leftColumn, encryptColumn, (targetName, suffix) -> result.addColumnName(targetName)); + appendEncryptColumnTokens(leftColumn, encryptColumn, useActualColumnName, (targetName, suffix) -> result.addColumnName(targetName)); return result; } - private String getColumnName(final ColumnSegment columnSegment, final EncryptDerivedColumnSuffix derivedColumnSuffix, final String actualColumnName) { - return TableSourceType.TEMPORARY_TABLE == columnSegment.getColumnBoundInfo().getTableSourceType() + private String getColumnName(final ColumnSegment columnSegment, final EncryptDerivedColumnSuffix derivedColumnSuffix, final String actualColumnName, final boolean useActualColumnName) { + return !useActualColumnName && TableSourceType.TEMPORARY_TABLE == columnSegment.getColumnBoundInfo().getTableSourceType() ? derivedColumnSuffix.getDerivedColumnName(columnSegment.getIdentifier().getValue(), databaseType) : actualColumnName; } private EncryptAssignmentToken generateLiteralSQLToken(final String schemaName, final String tableName, final EncryptColumn encryptColumn, final ColumnAssignmentSegment segment, - final QuoteCharacter quoteCharacter) { + final QuoteCharacter quoteCharacter, final boolean useActualColumnName) { ColumnSegment leftColumn = getAssignedColumn(segment); EncryptLiteralAssignmentToken result = new EncryptLiteralAssignmentToken(leftColumn.getStartIndex(), segment.getStopIndex(), quoteCharacter); Object literalValue = ((LiteralExpressionSegment) segment.getValue()).getLiterals(); - appendEncryptColumnTokens(leftColumn, encryptColumn, (targetName, suffix) -> addLiteralSQLToken(schemaName, tableName, encryptColumn, targetName, suffix, result, literalValue)); + appendEncryptColumnTokens(leftColumn, encryptColumn, useActualColumnName, + (targetName, suffix) -> addLiteralSQLToken(schemaName, tableName, encryptColumn, targetName, suffix, result, literalValue)); return result; } @@ -144,16 +196,14 @@ private Object encrypt(final EncryptColumn encryptColumn, final EncryptDerivedCo } } - private Optional findEncryptTable(final ColumnSegment columnSegment) { - return rule.findEncryptTable(columnSegment.getColumnBoundInfo().getOriginalTable().getValue()); - } - - private void appendEncryptColumnTokens(final ColumnSegment leftColumn, final EncryptColumn encryptColumn, final EncryptColumnConsumer consumer) { - consumer.accept(getColumnName(leftColumn, EncryptDerivedColumnSuffix.CIPHER, encryptColumn.getCipher().getName()), EncryptDerivedColumnSuffix.CIPHER); + private void appendEncryptColumnTokens(final ColumnSegment leftColumn, final EncryptColumn encryptColumn, final boolean useActualColumnName, final EncryptColumnConsumer consumer) { + consumer.accept(getColumnName(leftColumn, EncryptDerivedColumnSuffix.CIPHER, encryptColumn.getCipher().getName(), useActualColumnName), EncryptDerivedColumnSuffix.CIPHER); encryptColumn.getAssistedQuery() - .ifPresent(optional -> consumer.accept(getColumnName(leftColumn, EncryptDerivedColumnSuffix.ASSISTED_QUERY, optional.getName()), EncryptDerivedColumnSuffix.ASSISTED_QUERY)); + .ifPresent(optional -> consumer.accept(getColumnName(leftColumn, EncryptDerivedColumnSuffix.ASSISTED_QUERY, optional.getName(), useActualColumnName), + EncryptDerivedColumnSuffix.ASSISTED_QUERY)); encryptColumn.getLikeQuery() - .ifPresent(optional -> consumer.accept(getColumnName(leftColumn, EncryptDerivedColumnSuffix.LIKE_QUERY, optional.getName()), EncryptDerivedColumnSuffix.LIKE_QUERY)); + .ifPresent(optional -> consumer.accept(getColumnName(leftColumn, EncryptDerivedColumnSuffix.LIKE_QUERY, optional.getName(), useActualColumnName), + EncryptDerivedColumnSuffix.LIKE_QUERY)); } private ColumnSegment getAssignedColumn(final ColumnAssignmentSegment assignmentSegment) { @@ -165,4 +215,13 @@ private interface EncryptColumnConsumer { void accept(String targetColumnName, EncryptDerivedColumnSuffix derivedColumnSuffix); } + + @Getter + @RequiredArgsConstructor + private static final class OpenQueryContext { + + private final LiteralExpressionSegment openQuerySQL; + + private final EncryptTable encryptTable; + } } diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java index e070296811892..03e71a0dd8ee2 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java @@ -44,7 +44,8 @@ public final class EncryptUpdateAssignmentTokenGenerator implements CollectionSQ @Override public boolean isGenerateSQLToken(final SQLStatementContext sqlStatementContext) { - return sqlStatementContext instanceof UpdateStatementContext && containsEncryptTable(sqlStatementContext.getTablesContext().getSimpleTables()); + return sqlStatementContext instanceof UpdateStatementContext && (containsEncryptTable(sqlStatementContext.getTablesContext().getSimpleTables()) + || EncryptOpenQueryUtils.isOpenQueryFunctionTable(((UpdateStatementContext) sqlStatementContext).getSqlStatement().getTable())); } private boolean containsEncryptTable(final Collection simpleTableSegments) { @@ -59,6 +60,6 @@ private boolean containsEncryptTable(final Collection simple @Override public Collection generateSQLTokens(final UpdateStatementContext sqlStatementContext) { return new EncryptAssignmentTokenGenerator(rule, database, sqlStatementContext.getSqlStatement().getDatabaseType()) - .generateSQLTokens(sqlStatementContext.getTablesContext(), sqlStatementContext.getSqlStatement().getSetAssignment()); + .generateSQLTokens(sqlStatementContext.getTablesContext(), sqlStatementContext.getSqlStatement().getSetAssignment(), sqlStatementContext.getSqlStatement().getTable()); } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java index b51721d653d51..a78661f215d2e 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java @@ -21,19 +21,26 @@ import org.apache.shardingsphere.encrypt.config.rule.EncryptColumnItemRuleConfiguration; import org.apache.shardingsphere.encrypt.config.rule.EncryptColumnRuleConfiguration; import org.apache.shardingsphere.encrypt.config.rule.EncryptTableRuleConfiguration; -import org.apache.shardingsphere.encrypt.rule.changed.EncryptTableChangedProcessor; import org.apache.shardingsphere.encrypt.rule.EncryptRule; +import org.apache.shardingsphere.encrypt.rule.changed.EncryptTableChangedProcessor; import org.apache.shardingsphere.infra.algorithm.core.config.AlgorithmConfiguration; import org.apache.shardingsphere.infra.binder.context.statement.SQLStatementContext; import org.apache.shardingsphere.infra.binder.context.statement.type.dml.InsertStatementContext; import org.apache.shardingsphere.infra.binder.context.statement.type.dml.SelectStatementContext; +import org.apache.shardingsphere.infra.binder.context.statement.type.dml.UpdateStatementContext; import org.apache.shardingsphere.infra.config.props.ConfigurationProperties; +import org.apache.shardingsphere.infra.metadata.database.ShardingSphereDatabase; import org.apache.shardingsphere.infra.rewrite.context.SQLRewriteContext; import org.apache.shardingsphere.infra.rewrite.context.SQLRewriteContextDecorator; import org.apache.shardingsphere.infra.route.context.RouteContext; import org.apache.shardingsphere.infra.spi.type.ordered.OrderedSPILoader; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.column.ColumnSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.FunctionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.FunctionTableSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.SimpleTableSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableNameSegment; +import org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.UpdateStatement; import org.apache.shardingsphere.sql.parser.statement.core.value.identifier.IdentifierValue; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -88,7 +95,7 @@ void assertDecorateWithoutEncryptTable() { @Test void assertDecorateWithoutDroppedEncryptTable() { - EncryptRuleConfiguration ruleConfig = getEncryptRuleConfiguration(); + EncryptRuleConfiguration ruleConfig = getEncryptRuleConfiguration(true); SQLRewriteContext sqlRewriteContext = mock(SQLRewriteContext.class); InsertStatementContext insertStatementContext = mock(InsertStatementContext.class, RETURNS_DEEP_STUBS); when(insertStatementContext.getTablesContext().getSimpleTables()).thenReturn(Collections.singleton( @@ -98,12 +105,39 @@ void assertDecorateWithoutDroppedEncryptTable() { verify(sqlRewriteContext, never()).addSQLTokenGenerators(any()); } - private EncryptRuleConfiguration getEncryptRuleConfiguration() { + @Test + void assertDecorateWithOpenQueryEncryptTable() { + SQLRewriteContext sqlRewriteContext = mock(SQLRewriteContext.class); + when(sqlRewriteContext.getParameters()).thenReturn(Collections.emptyList()); + UpdateStatementContext updateStatementContext = mock(UpdateStatementContext.class, RETURNS_DEEP_STUBS); + when(updateStatementContext.getTablesContext().getSimpleTables()).thenReturn(Collections.emptyList()); + when(updateStatementContext.getTablesContext().getTableNames()).thenReturn(Collections.emptyList()); + when(updateStatementContext.getWhereSegments()).thenReturn(Collections.emptyList()); + UpdateStatement updateStatement = mock(UpdateStatement.class); + when(updateStatement.getTable()).thenReturn(createOpenQueryTableSegment()); + when(updateStatementContext.getSqlStatement()).thenReturn(updateStatement); + when(sqlRewriteContext.getSqlStatementContext()).thenReturn(updateStatementContext); + when(sqlRewriteContext.getDatabase()).thenReturn(mock(ShardingSphereDatabase.class)); + EncryptRule encryptRule = new EncryptRule("foo_db", getEncryptRuleConfiguration(false)); + decorator.decorate(encryptRule, mock(ConfigurationProperties.class), sqlRewriteContext, mock(RouteContext.class)); + verify(sqlRewriteContext).addSQLTokenGenerators(any()); + } + + private EncryptRuleConfiguration getEncryptRuleConfiguration(final boolean dropEncryptTable) { EncryptColumnRuleConfiguration columnConfig = new EncryptColumnRuleConfiguration("pwd", new EncryptColumnItemRuleConfiguration("pwd_cipher", "standard_encryptor")); EncryptTableRuleConfiguration tableConfig = new EncryptTableRuleConfiguration("t_encrypt", Collections.singleton(columnConfig)); EncryptRuleConfiguration result = new EncryptRuleConfiguration(new LinkedList<>(Collections.singleton(tableConfig)), Collections.singletonMap("standard_encryptor", new AlgorithmConfiguration("CORE.FIXTURE", new Properties()))); - new EncryptTableChangedProcessor().dropRuleItemConfiguration("t_encrypt", result); + if (dropEncryptTable) { + new EncryptTableChangedProcessor().dropRuleItemConfiguration("t_encrypt", result); + } return result; } + + private FunctionTableSegment createOpenQueryTableSegment() { + FunctionSegment functionSegment = new FunctionSegment(0, 0, "OPENQUERY", "OPENQUERY (foo_server, 'SELECT pwd FROM foo_schema.t_encrypt')"); + functionSegment.getParameters().add(new ColumnSegment(0, 0, new IdentifierValue("foo_server"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(0, 0, "SELECT pwd FROM foo_schema.t_encrypt")); + return new FunctionTableSegment(0, 0, functionSegment); + } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java index 4ded88a893a42..6730690f8b5e1 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java @@ -26,15 +26,18 @@ import org.apache.shardingsphere.encrypt.rule.table.EncryptTable; import org.apache.shardingsphere.infra.binder.context.segment.table.TablesContext; import org.apache.shardingsphere.infra.metadata.database.ShardingSphereDatabase; +import org.apache.shardingsphere.infra.rewrite.sql.token.common.pojo.SQLToken; import org.apache.shardingsphere.infra.spi.type.typed.TypedSPILoader; import org.apache.shardingsphere.sql.parser.statement.core.enums.TableSourceType; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.ColumnAssignmentSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.SetAssignmentSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.column.ColumnSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.FunctionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.ParameterMarkerExpressionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.bound.ColumnSegmentBoundInfo; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.bound.TableSegmentBoundInfo; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.FunctionTableSegment; import org.apache.shardingsphere.sql.parser.statement.core.value.identifier.IdentifierValue; import org.junit.jupiter.api.AfterAll; import org.junit.jupiter.api.BeforeAll; @@ -46,7 +49,9 @@ import org.mockito.MockedConstruction; import org.mockito.junit.jupiter.MockitoExtension; +import java.util.Collection; import java.util.Collections; +import java.util.Iterator; import java.util.Optional; import static org.hamcrest.Matchers.is; @@ -136,4 +141,45 @@ void assertGenerateSQLTokenWithInsertLiteralExpressionSegment() { when(assignmentSegment.getValue()).thenReturn(mock(LiteralExpressionSegment.class)); assertThat(tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment).size(), is(1)); } + + @Test + void assertGenerateSQLTokenWithOpenQueryLiteralExpressionSegment() { + ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); + when(database.getName()).thenReturn("foo_db"); + tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRule(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(new LiteralExpressionSegment(124, 144, "Sales and Marketing")); + when(assignmentSegment.getStopIndex()).thenReturn(144); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + Collection actual = tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegment()); + Iterator iterator = actual.iterator(); + assertThat(actual.size(), is(2)); + assertThat(iterator.next().toString(), is("group_name_cipher = 'encryptValue'")); + assertThat(iterator.next().toString(), is("'SELECT group_name_cipher FROM dbo.Department WHERE DepartmentID = 4'")); + } + + private EncryptRule mockOpenQueryEncryptRule() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.getAllTableNames()).thenReturn(Collections.singleton("Department")); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); + when(encryptTable.getTable()).thenReturn("Department"); + when(encryptTable.getEncryptColumn("GroupName")).thenReturn(encryptColumn); + when(encryptColumn.getName()).thenReturn("GroupName"); + when(encryptColumn.getCipher().getName()).thenReturn("group_name_cipher"); + when(encryptColumn.getCipher().encrypt("foo_db", "dbo", "Department", "GroupName", Collections.singletonList("Sales and Marketing"))) + .thenReturn(Collections.singletonList("encryptValue")); + return result; + } + + private FunctionTableSegment createOpenQueryTableSegment() { + FunctionSegment functionSegment = new FunctionSegment(7, 106, "OPENQUERY", "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE DepartmentID = 4')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 95, "SELECT GroupName FROM dbo.Department WHERE DepartmentID = 4")); + return new FunctionTableSegment(7, 106, functionSegment); + } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGeneratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGeneratorTest.java index 6bb48adf2f551..c4c606f16fab1 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGeneratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGeneratorTest.java @@ -22,11 +22,12 @@ import org.apache.shardingsphere.infra.binder.context.statement.type.dml.UpdateStatementContext; import org.apache.shardingsphere.infra.metadata.database.ShardingSphereDatabase; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.ColumnAssignmentSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.FunctionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.bound.TableSegmentBoundInfo; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.FunctionTableSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.SimpleTableSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableNameSegment; import org.apache.shardingsphere.sql.parser.statement.core.value.identifier.IdentifierValue; -import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Answers; @@ -43,31 +44,27 @@ @ExtendWith(MockitoExtension.class) class EncryptUpdateAssignmentTokenGeneratorTest { - private EncryptUpdateAssignmentTokenGenerator tokenGenerator; - @Mock(answer = Answers.RETURNS_DEEP_STUBS) private UpdateStatementContext updateStatementContext; - @BeforeEach - void setup() { - EncryptRule encryptRule = mockEncryptRule(); - tokenGenerator = new EncryptUpdateAssignmentTokenGenerator(encryptRule, mock(ShardingSphereDatabase.class)); + @Test + void assertIsGenerateSQLTokenUpdateSQLSuccess() { + EncryptRule encryptRule = mock(EncryptRule.class); + when(encryptRule.findEncryptTable("table")).thenReturn(Optional.of(mock(EncryptTable.class))); TableNameSegment tableNameSegment = new TableNameSegment(0, 0, new IdentifierValue("table")); tableNameSegment.setTableBoundInfo(new TableSegmentBoundInfo(new IdentifierValue("foo_db"), new IdentifierValue("foo_db"))); when(updateStatementContext.getTablesContext().getSimpleTables()).thenReturn(Collections.singleton(new SimpleTableSegment(tableNameSegment))); - ColumnAssignmentSegment assignmentSegment = mock(ColumnAssignmentSegment.class); - when(updateStatementContext.getSqlStatement().getSetAssignment().getAssignments()).thenReturn(Collections.singleton(assignmentSegment)); - } - - private EncryptRule mockEncryptRule() { - EncryptRule result = mock(EncryptRule.class); - EncryptTable encryptTable = mock(EncryptTable.class); - when(result.findEncryptTable("table")).thenReturn(Optional.of(encryptTable)); - return result; + when(updateStatementContext.getSqlStatement().getSetAssignment().getAssignments()).thenReturn(Collections.singleton(mock(ColumnAssignmentSegment.class))); + EncryptUpdateAssignmentTokenGenerator tokenGenerator = new EncryptUpdateAssignmentTokenGenerator(encryptRule, mock(ShardingSphereDatabase.class)); + assertTrue(tokenGenerator.isGenerateSQLToken(updateStatementContext)); } @Test - void assertIsGenerateSQLTokenUpdateSQLSuccess() { + void assertIsGenerateSQLTokenWithOpenQueryTarget() { + EncryptUpdateAssignmentTokenGenerator tokenGenerator = new EncryptUpdateAssignmentTokenGenerator(mock(EncryptRule.class), mock(ShardingSphereDatabase.class)); + when(updateStatementContext.getTablesContext().getSimpleTables()).thenReturn(Collections.emptyList()); + FunctionSegment functionSegment = new FunctionSegment(0, 0, "OPENQUERY", "OPENQUERY (foo_server, 'SELECT foo_col FROM foo_schema.foo_tbl')"); + when(updateStatementContext.getSqlStatement().getTable()).thenReturn(new FunctionTableSegment(0, 0, functionSegment)); assertTrue(tokenGenerator.isGenerateSQLToken(updateStatementContext)); } } diff --git a/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java b/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java index b30e44788fe7c..a1d7cd5bff7be 100644 --- a/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java +++ b/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java @@ -83,6 +83,10 @@ protected Collection mockSchemas(final String schemaName) new ShardingSphereColumn("WorkOrderID", Types.INTEGER, false, false, false, true, false, false), new ShardingSphereColumn("ScrapReasonID", Types.INTEGER, false, false, false, true, false, false), new ShardingSphereColumn("ScrappedQty", Types.INTEGER, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); + tables.add(new ShardingSphereTable("Department", Arrays.asList( + new ShardingSphereColumn("DepartmentID", Types.INTEGER, false, false, false, true, false, false), + new ShardingSphereColumn("Name", Types.VARCHAR, false, false, false, true, false, false), + new ShardingSphereColumn("GroupName", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); tables.add(new ShardingSphereTable("StateRegion", Arrays.asList( new ShardingSphereColumn("StateCode", Types.INTEGER, false, false, false, true, false, false), new ShardingSphereColumn("CountryRegionName", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); @@ -107,6 +111,12 @@ protected Collection mockSchemas(final String schemaName) new ShardingSphereColumn("YearToDateAmt", Types.DECIMAL, false, false, false, true, false, false), new ShardingSphereColumn("RegionCode", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); result.add(new ShardingSphereSchema("Sales", mock(DatabaseType.class), salesSchemaTables, Collections.emptyList())); + Collection humanResourcesSchemaTables = new LinkedList<>(); + humanResourcesSchemaTables.add(new ShardingSphereTable("Department", Arrays.asList( + new ShardingSphereColumn("DepartmentID", Types.INTEGER, false, false, false, true, false, false), + new ShardingSphereColumn("Name", Types.VARCHAR, false, false, false, true, false, false), + new ShardingSphereColumn("GroupName", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); + result.add(new ShardingSphereSchema("HumanResources", mock(DatabaseType.class), humanResourcesSchemaTables, Collections.emptyList())); return result; } @@ -121,11 +131,13 @@ protected void mockDatabaseRules(final Collection rules, fin singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "t_user"); singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "ScrapReason"); singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "WorkOrder"); + singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "Department"); singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "StateRegion"); singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "vStateProvinceCountryRegion"); singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "SalesPerson"); singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", schemaName, "SalesOrderHeader"); singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", "Sales", "SalesPerson"); + singleRule.get().getAttributes().getAttribute(MutableDataNodeRuleAttribute.class).put("encrypt_ds", "HumanResources", "Department"); } } } diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml index 3093abf3c9650..2e8082e6493d2 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml @@ -115,6 +115,11 @@ + + + + + diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml index 98ce1bed99047..5e9b3cf3af396 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml @@ -158,6 +158,12 @@ rules: cipher: name: remark_cipher encryptorName: rewrite_normal_fixture + Department: + columns: + GroupName: + cipher: + name: group_name_cipher + encryptorName: rewrite_normal_fixture StateRegion: columns: CountryRegionName: From d04c4f8d45adf970516b1f4dc3c3fdc9e2a5e39f Mon Sep 17 00:00:00 2001 From: Claire Date: Thu, 16 Jul 2026 14:56:41 +0800 Subject: [PATCH 3/8] support OPENQUERY function --- RELEASE-NOTES.md | 1 + 1 file changed, 1 insertion(+) diff --git a/RELEASE-NOTES.md b/RELEASE-NOTES.md index a5af0bc22b30b..7a939663fe337 100644 --- a/RELEASE-NOTES.md +++ b/RELEASE-NOTES.md @@ -76,6 +76,7 @@ 1. Proxy: Support basic Firebird batch operations: create, send, execute, cancel, and release - [#38605](https://github.com/apache/shardingsphere/pull/38605) 1. JDBC & Proxy: Add a check to verify database name naming conventions. - [#38883](https://github.com/apache/shardingsphere/pull/38883) 1. Encrypt: Support SqlServer update statement for Updating data in a remote table by using a linked server when use encrypt feature - [#39122](https://github.com/apache/shardingsphere/pull/39122) +1. Encrypt: Support SqlServer update statement for Updating data in a remote table by using the OPENQUERY function when use encrypt feature - [#39156](https://github.com/apache/shardingsphere/pull/39156) 1. Encrypt: Support SqlServer update statement for Specifying a table alias as the target object when use encrypt feature - [#38733](https://github.com/apache/shardingsphere/pull/38733) 1. Encrypt: Support SqlServer update statement for Specifying a view as the target object when use encrypt feature - [#38896](https://github.com/apache/shardingsphere/pull/38896) 1. Encrypt: Support SqlServer for Using the UPDATE statement with information from another table when use encrypt feature - [#38926](https://github.com/apache/shardingsphere/pull/38926) From f115808b981a886bb39f593200871c1b30099a31 Mon Sep 17 00:00:00 2001 From: Claire Date: Mon, 20 Jul 2026 12:27:46 +0800 Subject: [PATCH 4/8] fix --- .../EncryptAssignmentParameterRewriter.java | 58 ++- .../EncryptAssignmentTokenGenerator.java | 99 +++-- .../EncryptOpenQueryPassThroughSQL.java | 400 ++++++++++++++++++ .../assignment/EncryptOpenQueryUtils.java | 27 +- ...EncryptUpdateAssignmentTokenGenerator.java | 9 +- ...EncryptSQLRewriteContextDecoratorTest.java | 23 + .../EncryptAssignmentTokenGeneratorTest.java | 171 +++++++- .../EncryptOpenQueryPassThroughSQLTest.java | 146 +++++++ .../assignment/EncryptOpenQueryUtilsTest.java | 27 +- .../engine/scenario/EncryptSQLRewriterIT.java | 6 +- .../query-with-cipher/dml/update/update.xml | 57 ++- .../encrypt/config/query-with-cipher.yaml | 14 + 12 files changed, 949 insertions(+), 88 deletions(-) create mode 100644 features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java create mode 100644 features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/parameter/rewriter/EncryptAssignmentParameterRewriter.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/parameter/rewriter/EncryptAssignmentParameterRewriter.java index 364a8a5421474..9b4855ca32efa 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/parameter/rewriter/EncryptAssignmentParameterRewriter.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/parameter/rewriter/EncryptAssignmentParameterRewriter.java @@ -20,8 +20,10 @@ import com.google.common.base.Preconditions; import lombok.RequiredArgsConstructor; import org.apache.shardingsphere.database.connector.core.type.DatabaseTypeRegistry; +import org.apache.shardingsphere.encrypt.rewrite.token.generator.assignment.EncryptOpenQueryUtils; import org.apache.shardingsphere.encrypt.rule.EncryptRule; import org.apache.shardingsphere.encrypt.rule.column.EncryptColumn; +import org.apache.shardingsphere.encrypt.rule.table.EncryptTable; import org.apache.shardingsphere.infra.binder.context.statement.SQLStatementContext; import org.apache.shardingsphere.infra.binder.context.statement.type.dml.InsertStatementContext; import org.apache.shardingsphere.infra.binder.context.statement.type.dml.UpdateStatementContext; @@ -33,6 +35,7 @@ import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.SetAssignmentSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.ExpressionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.ParameterMarkerExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableSegment; import org.apache.shardingsphere.sql.parser.statement.core.statement.SQLStatement; import org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.InsertStatement; import org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.UpdateStatement; @@ -66,23 +69,54 @@ public boolean isNeedRewrite(final SQLStatementContext sqlStatementContext) { @Override public void rewrite(final ParameterBuilder paramBuilder, final SQLStatementContext sqlStatementContext, final List params) { - String schemaName = sqlStatementContext.getTablesContext().getSchemaName() - .orElseGet(() -> new DatabaseTypeRegistry(sqlStatementContext.getSqlStatement().getDatabaseType()).getDefaultSchemaName(databaseName)); + TableSegment openQueryTarget = findOpenQueryTarget(sqlStatementContext); for (ColumnAssignmentSegment each : getSetAssignmentSegment(sqlStatementContext.getSqlStatement()).getAssignments()) { String columnName = each.getColumns().get(0).getIdentifier().getValue(); - String tableName = each.getColumns().get(0).getColumnBoundInfo().getOriginalTable().getValue(); - if (!rule.findEncryptTable(tableName).map(optional -> optional.isEncryptColumn(columnName)).orElse(false)) { + String originalTableName = each.getColumns().get(0).getColumnBoundInfo().getOriginalTable().getValue(); + EncryptTable encryptTable = resolveEncryptTable(originalTableName, columnName, openQueryTarget); + if (null == encryptTable) { continue; } - EncryptColumn encryptColumn = rule.getEncryptTable(tableName).getEncryptColumn(columnName); - StandardParameterBuilder standardParamBuilder = paramBuilder instanceof StandardParameterBuilder - ? (StandardParameterBuilder) paramBuilder - : ((GroupedParameterBuilder) paramBuilder).getParameterBuilders().get(0); - ExpressionSegment valueExpression = each.getValue(); - if (valueExpression instanceof ParameterMarkerExpressionSegment) { - encryptParameters(standardParamBuilder, schemaName, tableName, encryptColumn, ((ParameterMarkerExpressionSegment) valueExpression).getParameterMarkerIndex(), params); - } + rewriteEncryptAssignment(paramBuilder, sqlStatementContext, openQueryTarget, encryptTable, columnName, each, params); + } + } + + private void rewriteEncryptAssignment(final ParameterBuilder paramBuilder, final SQLStatementContext sqlStatementContext, final TableSegment openQueryTarget, + final EncryptTable encryptTable, final String columnName, final ColumnAssignmentSegment assignmentSegment, final List params) { + String schemaName = resolveSchemaName(sqlStatementContext, openQueryTarget); + EncryptColumn encryptColumn = encryptTable.getEncryptColumn(columnName); + StandardParameterBuilder standardParamBuilder = paramBuilder instanceof StandardParameterBuilder + ? (StandardParameterBuilder) paramBuilder + : ((GroupedParameterBuilder) paramBuilder).getParameterBuilders().get(0); + ExpressionSegment valueExpression = assignmentSegment.getValue(); + if (valueExpression instanceof ParameterMarkerExpressionSegment) { + encryptParameters(standardParamBuilder, schemaName, encryptTable.getTable(), encryptColumn, ((ParameterMarkerExpressionSegment) valueExpression).getParameterMarkerIndex(), params); + } + } + + private String resolveSchemaName(final SQLStatementContext sqlStatementContext, final TableSegment openQueryTarget) { + String defaultSchemaName = sqlStatementContext.getTablesContext().getSchemaName() + .orElseGet(() -> new DatabaseTypeRegistry(sqlStatementContext.getSqlStatement().getDatabaseType()).getDefaultSchemaName(databaseName)); + return null == openQueryTarget ? defaultSchemaName : EncryptOpenQueryUtils.findSchemaName(openQueryTarget).orElse(defaultSchemaName); + } + + private TableSegment findOpenQueryTarget(final SQLStatementContext sqlStatementContext) { + if (!(sqlStatementContext instanceof UpdateStatementContext)) { + return null; + } + TableSegment table = ((UpdateStatementContext) sqlStatementContext).getSqlStatement().getTable(); + return EncryptOpenQueryUtils.isOpenQueryFunctionTable(table) ? table : null; + } + + private EncryptTable resolveEncryptTable(final String originalTableName, final String columnName, final TableSegment openQueryTarget) { + Optional fromOriginal = rule.findEncryptTable(originalTableName).filter(t -> t.isEncryptColumn(columnName)); + if (fromOriginal.isPresent()) { + return fromOriginal.get(); + } + if (null == openQueryTarget) { + return null; } + return EncryptOpenQueryUtils.findEncryptTable(rule, openQueryTarget).filter(t -> t.isEncryptColumn(columnName)).orElse(null); } private SetAssignmentSegment getSetAssignmentSegment(final SQLStatement sqlStatement) { diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java index d5716fd879a36..f17dccfa9be5a 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java @@ -17,7 +17,6 @@ package org.apache.shardingsphere.encrypt.rewrite.token.generator.assignment; -import lombok.Getter; import lombok.RequiredArgsConstructor; import org.apache.shardingsphere.database.connector.core.metadata.database.enums.QuoteCharacter; import org.apache.shardingsphere.database.connector.core.type.DatabaseType; @@ -48,8 +47,6 @@ import java.util.LinkedList; import java.util.List; import java.util.Optional; -import java.util.regex.Matcher; -import java.util.regex.Pattern; /** * Assignment generator for encrypt. @@ -72,7 +69,7 @@ public final class EncryptAssignmentTokenGenerator { * @return generated SQL tokens */ public Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment) { - return generateSQLTokens(tablesContext, setAssignmentSegment, Optional.empty()); + return generateNormalUpdateTokens(tablesContext, setAssignmentSegment); } /** @@ -80,59 +77,84 @@ public Collection generateSQLTokens(final TablesContext tablesContext, * * @param tablesContext SQL statement context * @param setAssignmentSegment set assignment segment - * @param targetTable target table segment + * @param openQueryTable OPENQUERY function table segment * @return generated SQL tokens */ - public Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final TableSegment targetTable) { - return generateSQLTokens(tablesContext, setAssignmentSegment, Optional.of(targetTable)); + Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final TableSegment openQueryTable) { + return generateOpenQueryUpdateTokens(tablesContext, setAssignmentSegment, openQueryTable); } - private Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final Optional targetTable) { + private Collection generateNormalUpdateTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment) { Collection result = new LinkedList<>(); for (ColumnAssignmentSegment each : setAssignmentSegment.getAssignments()) { ColumnSegment assignedColumn = getAssignedColumn(each); - Optional openQueryContext = findOpenQueryContext(targetTable, assignedColumn.getIdentifier().getValue()); - findEncryptTable(assignedColumn, openQueryContext) - .ifPresent(encryptTable -> appendEncryptAssignmentTokens(result, tablesContext, targetTable, each, assignedColumn, openQueryContext, encryptTable)); + String columnName = assignedColumn.getIdentifier().getValue(); + Optional encryptTable = rule.findEncryptTable(assignedColumn.getColumnBoundInfo().getOriginalTable().getValue()); + if (!encryptTable.isPresent() || !encryptTable.get().isEncryptColumn(columnName)) { + continue; + } + EncryptColumn encryptColumn = encryptTable.get().getEncryptColumn(columnName); + appendNormalAssignmentTokens(result, tablesContext, each, encryptTable.get(), encryptColumn); } return result; } - private void appendEncryptAssignmentTokens(final Collection result, final TablesContext tablesContext, final Optional targetTable, - final ColumnAssignmentSegment assignmentSegment, final ColumnSegment assignedColumn, - final Optional openQueryContext, final EncryptTable encryptTable) { - String columnName = assignedColumn.getIdentifier().getValue(); - if (!encryptTable.isEncryptColumn(columnName)) { + private Collection generateOpenQueryUpdateTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final TableSegment openQueryTable) { + Collection result = new LinkedList<>(); + Collection openQueryEncryptColumns = new LinkedList<>(); + for (ColumnAssignmentSegment each : setAssignmentSegment.getAssignments()) { + ColumnSegment assignedColumn = getAssignedColumn(each); + String columnName = assignedColumn.getIdentifier().getValue(); + Optional encryptTable = findOpenQueryEncryptTable(openQueryTable, columnName); + if (!encryptTable.isPresent()) { + continue; + } + EncryptColumn encryptColumn = encryptTable.get().getEncryptColumn(columnName); + appendOpenQueryAssignmentTokens(result, tablesContext, openQueryTable, each, encryptTable.get(), encryptColumn); + openQueryEncryptColumns.add(encryptColumn); + } + appendComposedOpenQuerySQLToken(result, openQueryTable, openQueryEncryptColumns); + return result; + } + + private void appendComposedOpenQuerySQLToken(final Collection result, final TableSegment openQueryTable, final Collection encryptColumns) { + if (encryptColumns.isEmpty()) { return; } - EncryptColumn encryptColumn = encryptTable.getEncryptColumn(columnName); + Optional openQuerySQL = EncryptOpenQueryUtils.findOpenQuerySQLLiteral(openQueryTable); + if (!openQuerySQL.isPresent()) { + return; + } + result.add(generateOpenQuerySQLToken(openQuerySQL.get(), encryptColumns)); + } + + private void appendNormalAssignmentTokens(final Collection result, final TablesContext tablesContext, + final ColumnAssignmentSegment assignmentSegment, final EncryptTable encryptTable, final EncryptColumn encryptColumn) { DatabaseTypeRegistry databaseTypeRegistry = new DatabaseTypeRegistry(databaseType); - String schemaName = targetTable.flatMap(EncryptOpenQueryUtils::findSchemaName) - .orElseGet(() -> tablesContext.getSchemaName().orElseGet(() -> databaseTypeRegistry.getDefaultSchemaName(database.getName()))); + String schemaName = tablesContext.getSchemaName().orElseGet(() -> databaseTypeRegistry.getDefaultSchemaName(database.getName())); QuoteCharacter quoteCharacter = databaseTypeRegistry.getDialectDatabaseMetaData().getQuoteCharacter(); - result.addAll(generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptColumn, assignmentSegment, quoteCharacter, openQueryContext.isPresent())); - openQueryContext.ifPresent(optional -> result.add(generateOpenQuerySQLToken(optional, encryptColumn))); + result.addAll(generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptColumn, assignmentSegment, quoteCharacter, false)); } - private Optional findEncryptTable(final ColumnSegment assignedColumn, final Optional openQueryContext) { - Optional result = rule.findEncryptTable(assignedColumn.getColumnBoundInfo().getOriginalTable().getValue()); - return result.isPresent() ? result : openQueryContext.map(OpenQueryContext::getEncryptTable); + private void appendOpenQueryAssignmentTokens(final Collection result, final TablesContext tablesContext, final TableSegment openQueryTable, + final ColumnAssignmentSegment assignmentSegment, final EncryptTable encryptTable, final EncryptColumn encryptColumn) { + DatabaseTypeRegistry databaseTypeRegistry = new DatabaseTypeRegistry(databaseType); + String schemaName = EncryptOpenQueryUtils.findSchemaName(openQueryTable) + .orElseGet(() -> tablesContext.getSchemaName().orElseGet(() -> databaseTypeRegistry.getDefaultSchemaName(database.getName()))); + QuoteCharacter quoteCharacter = databaseTypeRegistry.getDialectDatabaseMetaData().getQuoteCharacter(); + result.addAll(generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptColumn, assignmentSegment, quoteCharacter, true)); } - private Optional findOpenQueryContext(final Optional targetTable, final String columnName) { - if (!targetTable.isPresent()) { + private Optional findOpenQueryEncryptTable(final TableSegment openQueryTable, final String columnName) { + if (!EncryptOpenQueryUtils.findOpenQuerySQLLiteral(openQueryTable).isPresent()) { return Optional.empty(); } - Optional openQuerySQL = EncryptOpenQueryUtils.findOpenQuerySQLLiteral(targetTable.get()); - Optional encryptTable = EncryptOpenQueryUtils.findEncryptTable(rule, targetTable.get()).filter(optional -> optional.isEncryptColumn(columnName)); - return openQuerySQL.isPresent() && encryptTable.isPresent() ? Optional.of(new OpenQueryContext(openQuerySQL.get(), encryptTable.get())) : Optional.empty(); + return EncryptOpenQueryUtils.findEncryptTable(rule, openQueryTable).filter(optional -> optional.isEncryptColumn(columnName)); } - private EncryptOpenQuerySQLToken generateOpenQuerySQLToken(final OpenQueryContext openQueryContext, final EncryptColumn encryptColumn) { - String openQuerySQL = openQueryContext.getOpenQuerySQL().getText(); - String rewrittenSQL = Pattern.compile(String.join("", "\\b", Pattern.quote(encryptColumn.getName()), "\\b"), Pattern.CASE_INSENSITIVE).matcher(openQuerySQL) - .replaceAll(Matcher.quoteReplacement(encryptColumn.getCipher().getName())); - return new EncryptOpenQuerySQLToken(openQueryContext.getOpenQuerySQL().getStartIndex(), openQueryContext.getOpenQuerySQL().getStopIndex(), rewrittenSQL); + private EncryptOpenQuerySQLToken generateOpenQuerySQLToken(final LiteralExpressionSegment openQuerySQL, final Collection encryptColumns) { + String rewrittenSQL = EncryptOpenQueryPassThroughSQL.parse(openQuerySQL.getText()).rewrite(encryptColumns); + return new EncryptOpenQuerySQLToken(openQuerySQL.getStartIndex(), openQuerySQL.getStopIndex(), rewrittenSQL); } private Collection generateAssignmentSQLTokens(final String schemaName, final String tableName, final EncryptColumn encryptColumn, @@ -215,13 +237,4 @@ private interface EncryptColumnConsumer { void accept(String targetColumnName, EncryptDerivedColumnSuffix derivedColumnSuffix); } - - @Getter - @RequiredArgsConstructor - private static final class OpenQueryContext { - - private final LiteralExpressionSegment openQuerySQL; - - private final EncryptTable encryptTable; - } } diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java new file mode 100644 index 0000000000000..fd995dcc0a8f1 --- /dev/null +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java @@ -0,0 +1,400 @@ +/* + * 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.shardingsphere.encrypt.rewrite.token.generator.assignment; + +import lombok.Getter; +import org.apache.shardingsphere.encrypt.exception.syntax.UnsupportedEncryptSQLException; +import org.apache.shardingsphere.encrypt.rule.column.EncryptColumn; + +import java.util.ArrayList; +import java.util.Collection; +import java.util.List; +import java.util.Locale; +import java.util.Optional; + +/** + * Encrypt OPENQUERY pass-through SQL. + */ +@Getter +final class EncryptOpenQueryPassThroughSQL { + + private static final String UNSUPPORTED_SHAPE = "OPENQUERY with query that is not SELECT ... FROM"; + + private static final String UNSUPPORTED_SELECT_LITERAL = "OPENQUERY SELECT list with string literal"; + + private static final String UNSUPPORTED_SELECT_EXPRESSION = "OPENQUERY SELECT list with expression"; + + private static final String UNSUPPORTED_SPACE_DELIMITED_IDENTIFIER = "OPENQUERY with space-delimited identifier"; + + private static final String UNSUPPORTED_MULTIPART_TABLE = "OPENQUERY with three-part table name"; + + private static final String UNSUPPORTED_JOIN = "OPENQUERY with JOIN statement"; + + private final String selectList; + + private final String tableExpression; + + private final String tableName; + + private final Optional schemaName; + + private final String remainder; + + private EncryptOpenQueryPassThroughSQL(final String selectList, final String tableExpression, final String tableName, + final Optional schemaName, final String remainder) { + this.selectList = selectList; + this.tableExpression = tableExpression; + this.tableName = tableName; + this.schemaName = schemaName; + this.remainder = remainder; + } + + /** + * Find table name from pass-through SQL without validating supported shape. + * + * @param passThroughSQL pass-through SQL + * @return table name + */ + static Optional findTableName(final String passThroughSQL) { + String trimmedSQL = passThroughSQL.trim(); + if (!startsWithKeyword(trimmedSQL, "SELECT")) { + return Optional.empty(); + } + Optional fromIndex = findFromKeywordIndexIfPresent(trimmedSQL); + if (!fromIndex.isPresent()) { + return Optional.empty(); + } + return extractTableNameAfterFrom(trimmedSQL, fromIndex.get() + "FROM".length()); + } + + /** + * Parse and validate pass-through SQL. + * + * @param passThroughSQL pass-through SQL + * @return parsed pass-through SQL + * @throws UnsupportedEncryptSQLException if pass-through SQL shape is unsupported + */ + static EncryptOpenQueryPassThroughSQL parse(final String passThroughSQL) { + String trimmedSQL = passThroughSQL.trim(); + if (!startsWithKeyword(trimmedSQL, "SELECT")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); + } + int fromIndex = findFromKeywordIndex(trimmedSQL); + String selectList = trimmedSQL.substring("SELECT".length(), fromIndex).trim(); + if (selectList.isEmpty()) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); + } + validateSelectList(selectList); + int tableStartIndex = fromIndex + "FROM".length(); + TableReference tableReference = parseTableReference(trimmedSQL, tableStartIndex); + String remainder = trimmedSQL.substring(tableReference.getStopIndex()); + validateRemainder(remainder); + return new EncryptOpenQueryPassThroughSQL(selectList, tableReference.getExpression(), tableReference.getTableName(), tableReference.getSchemaName(), remainder); + } + + /** + * Rewrite pass-through SQL with encrypted physical columns in the SELECT list. + * + * @param encryptColumns encrypt columns + * @return rewritten pass-through SQL + */ + String rewrite(final Collection encryptColumns) { + List rewrittenItems = new ArrayList<>(); + for (String each : splitSelectList(selectList)) { + String trimmedItem = each.trim(); + Optional matchedEncryptColumn = findEncryptColumn(encryptColumns, unwrapIdentifier(trimmedItem)); + rewrittenItems.add(matchedEncryptColumn.map(EncryptOpenQueryPassThroughSQL::getPhysicalColumnNames).orElse(trimmedItem)); + } + return "SELECT " + String.join(", ", rewrittenItems) + " FROM " + tableExpression + remainder; + } + + private static void validateSelectList(final String selectList) { + for (String each : splitSelectList(selectList)) { + String trimmedItem = each.trim(); + if (trimmedItem.contains("'")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SELECT_LITERAL); + } + if (trimmedItem.contains("(") || trimmedItem.contains(")")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SELECT_EXPRESSION); + } + validateColumnIdentifier(trimmedItem); + } + } + + private static void validateColumnIdentifier(final String identifier) { + if (identifier.isEmpty()) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); + } + if (identifier.startsWith("[") && identifier.endsWith("]")) { + String inner = identifier.substring(1, identifier.length() - 1); + if (inner.contains(" ") || inner.contains("]") || inner.contains("[")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SPACE_DELIMITED_IDENTIFIER); + } + return; + } + for (int index = 0; index < identifier.length(); index++) { + char current = identifier.charAt(index); + if (!Character.isLetterOrDigit(current) && '_' != current) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SELECT_EXPRESSION); + } + } + } + + private static void validateRemainder(final String remainder) { + String trimmedRemainder = remainder.trim().toUpperCase(Locale.ENGLISH); + if (trimmedRemainder.startsWith("JOIN ") || trimmedRemainder.contains(" JOIN ")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_JOIN); + } + } + + private static TableReference parseTableReference(final String passThroughSQL, final int startIndex) { + int index = skipWhitespace(passThroughSQL, startIndex); + IdentifierPart schemaPart = readIdentifierPart(passThroughSQL, index); + index = schemaPart.getStopIndex(); + if (index < passThroughSQL.length() && '.' == passThroughSQL.charAt(index)) { + IdentifierPart tablePart = readIdentifierPart(passThroughSQL, index + 1); + if (tablePart.getStopIndex() < passThroughSQL.length() && '.' == passThroughSQL.charAt(tablePart.getStopIndex())) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_MULTIPART_TABLE); + } + String expression = passThroughSQL.substring(schemaPart.getStartIndex(), tablePart.getStopIndex()); + return new TableReference(expression, tablePart.getValue(), Optional.of(schemaPart.getValue()), tablePart.getStopIndex()); + } + String expression = passThroughSQL.substring(schemaPart.getStartIndex(), schemaPart.getStopIndex()); + return new TableReference(expression, schemaPart.getValue(), Optional.empty(), schemaPart.getStopIndex()); + } + + private static IdentifierPart readIdentifierPart(final String passThroughSQL, final int startIndex) { + int index = skipWhitespace(passThroughSQL, startIndex); + if (index >= passThroughSQL.length()) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); + } + if ('[' == passThroughSQL.charAt(index)) { + int closeIndex = passThroughSQL.indexOf(']', index + 1); + if (-1 == closeIndex) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); + } + String inner = passThroughSQL.substring(index + 1, closeIndex); + if (inner.contains(" ")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SPACE_DELIMITED_IDENTIFIER); + } + return new IdentifierPart(inner, index, closeIndex + 1); + } + int stopIndex = index; + while (stopIndex < passThroughSQL.length()) { + char current = passThroughSQL.charAt(stopIndex); + if (Character.isLetterOrDigit(current) || '_' == current) { + stopIndex++; + continue; + } + break; + } + if (stopIndex == index) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); + } + return new IdentifierPart(passThroughSQL.substring(index, stopIndex), index, stopIndex); + } + + private static Optional extractTableNameAfterFrom(final String passThroughSQL, final int startIndex) { + int index = skipWhitespace(passThroughSQL, startIndex); + Optional identifierPart = readIdentifierPartIfPresent(passThroughSQL, index); + if (!identifierPart.isPresent()) { + return Optional.empty(); + } + String tableName = identifierPart.get().getValue(); + index = identifierPart.get().getStopIndex(); + while (index < passThroughSQL.length() && '.' == passThroughSQL.charAt(index)) { + Optional nextPart = readIdentifierPartIfPresent(passThroughSQL, index + 1); + if (!nextPart.isPresent()) { + break; + } + tableName = nextPart.get().getValue(); + index = nextPart.get().getStopIndex(); + } + return Optional.of(tableName); + } + + private static Optional findFromKeywordIndexIfPresent(final String passThroughSQL) { + int index = 0; + boolean inString = false; + while (index <= passThroughSQL.length() - 4) { + if (!inString && isFromKeywordAt(passThroughSQL, index)) { + return Optional.of(index); + } + char current = passThroughSQL.charAt(index); + if ('\'' != current) { + index++; + continue; + } + if (!inString) { + inString = true; + index++; + continue; + } + if (index + 1 < passThroughSQL.length() && '\'' == passThroughSQL.charAt(index + 1)) { + index += 2; + continue; + } + inString = false; + index++; + } + return Optional.empty(); + } + + private static Optional readIdentifierPartIfPresent(final String passThroughSQL, final int startIndex) { + int index = skipWhitespace(passThroughSQL, startIndex); + if (index >= passThroughSQL.length()) { + return Optional.empty(); + } + if ('[' == passThroughSQL.charAt(index)) { + int closeIndex = passThroughSQL.indexOf(']', index + 1); + if (-1 == closeIndex) { + return Optional.empty(); + } + return Optional.of(new IdentifierPart(passThroughSQL.substring(index + 1, closeIndex), index, closeIndex + 1)); + } + int stopIndex = index; + while (stopIndex < passThroughSQL.length()) { + char current = passThroughSQL.charAt(stopIndex); + if (Character.isLetterOrDigit(current) || '_' == current) { + stopIndex++; + continue; + } + break; + } + if (stopIndex == index) { + return Optional.empty(); + } + return Optional.of(new IdentifierPart(passThroughSQL.substring(index, stopIndex), index, stopIndex)); + } + + private static int findFromKeywordIndex(final String passThroughSQL) { + Optional result = findFromKeywordIndexIfPresent(passThroughSQL); + if (!result.isPresent()) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); + } + return result.get(); + } + + private static boolean isFromKeywordAt(final String passThroughSQL, final int index) { + return matchesKeyword(passThroughSQL, index, "FROM") && isWordBoundary(passThroughSQL, index - 1) && isWordBoundary(passThroughSQL, index + 4); + } + + private static boolean startsWithKeyword(final String passThroughSQL, final String keyword) { + if (passThroughSQL.length() < keyword.length()) { + return false; + } + if (!passThroughSQL.regionMatches(true, 0, keyword, 0, keyword.length())) { + return false; + } + return passThroughSQL.length() == keyword.length() || isWordBoundary(passThroughSQL, keyword.length()); + } + + private static boolean matchesKeyword(final String passThroughSQL, final int startIndex, final String keyword) { + return passThroughSQL.regionMatches(true, startIndex, keyword, 0, keyword.length()); + } + + private static boolean isWordBoundary(final String passThroughSQL, final int index) { + if (index < 0 || index >= passThroughSQL.length()) { + return true; + } + char current = passThroughSQL.charAt(index); + return !Character.isLetterOrDigit(current) && '_' != current; + } + + private static int skipWhitespace(final String passThroughSQL, final int startIndex) { + int index = startIndex; + while (index < passThroughSQL.length() && Character.isWhitespace(passThroughSQL.charAt(index))) { + index++; + } + return index; + } + + private static List splitSelectList(final String selectList) { + List result = new ArrayList<>(); + StringBuilder current = new StringBuilder(); + for (int index = 0; index < selectList.length(); index++) { + char currentChar = selectList.charAt(index); + if (',' == currentChar) { + result.add(current.toString()); + current = new StringBuilder(); + continue; + } + current.append(currentChar); + } + result.add(current.toString()); + return result; + } + + private static String unwrapIdentifier(final String identifier) { + if (identifier.startsWith("[") && identifier.endsWith("]")) { + return identifier.substring(1, identifier.length() - 1); + } + return identifier; + } + + private static Optional findEncryptColumn(final Collection encryptColumns, final String logicColumnName) { + for (EncryptColumn each : encryptColumns) { + if (each.getName().equalsIgnoreCase(logicColumnName)) { + return Optional.of(each); + } + } + return Optional.empty(); + } + + private static String getPhysicalColumnNames(final EncryptColumn encryptColumn) { + StringBuilder result = new StringBuilder(encryptColumn.getCipher().getName()); + encryptColumn.getAssistedQuery().ifPresent(optional -> result.append(", ").append(optional.getName())); + encryptColumn.getLikeQuery().ifPresent(optional -> result.append(", ").append(optional.getName())); + return result.toString(); + } + + @Getter + private static final class TableReference { + + private final String expression; + + private final String tableName; + + private final Optional schemaName; + + private final int stopIndex; + + private TableReference(final String expression, final String tableName, final Optional schemaName, final int stopIndex) { + this.expression = expression; + this.tableName = tableName; + this.schemaName = schemaName; + this.stopIndex = stopIndex; + } + } + + @Getter + private static final class IdentifierPart { + + private final String value; + + private final int startIndex; + + private final int stopIndex; + + private IdentifierPart(final String value, final int startIndex, final int stopIndex) { + this.value = value; + this.startIndex = startIndex; + this.stopIndex = stopIndex; + } + } +} diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java index 8b737e699bb53..057659fdbad17 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java @@ -19,7 +19,6 @@ import lombok.AccessLevel; import lombok.NoArgsConstructor; -import org.apache.shardingsphere.database.connector.core.metadata.database.enums.QuoteCharacter; import org.apache.shardingsphere.encrypt.rule.EncryptRule; import org.apache.shardingsphere.encrypt.rule.table.EncryptTable; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.ExpressionSegment; @@ -29,8 +28,6 @@ import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableSegment; import java.util.Optional; -import java.util.regex.Matcher; -import java.util.regex.Pattern; /** * Encrypt OPENQUERY utility. @@ -38,8 +35,6 @@ @NoArgsConstructor(access = AccessLevel.PRIVATE) public final class EncryptOpenQueryUtils { - private static final Pattern FROM_TABLE_PATTERN = Pattern.compile("\\bFROM\\s+([^\\s,;]+)", Pattern.CASE_INSENSITIVE); - /** * Whether table segment is OPENQUERY function table. * @@ -83,18 +78,8 @@ public static Optional findEncryptTable(final EncryptRule rule, fi if (!openQuerySQL.isPresent()) { return Optional.empty(); } - Matcher matcher = FROM_TABLE_PATTERN.matcher(openQuerySQL.get().getText()); - if (!matcher.find()) { - return Optional.empty(); - } - String actualTableName = QuoteCharacter.unwrapAndTrimText(matcher.group(1).substring(matcher.group(1).lastIndexOf('.') + 1)); - for (String each : rule.getAllTableNames()) { - Optional encryptTable = rule.findEncryptTable(each); - if (encryptTable.isPresent() && each.equalsIgnoreCase(actualTableName)) { - return encryptTable; - } - } - return Optional.empty(); + Optional tableName = EncryptOpenQueryPassThroughSQL.findTableName(openQuerySQL.get().getText()); + return tableName.isPresent() ? rule.findEncryptTable(tableName.get()) : Optional.empty(); } /** @@ -108,12 +93,6 @@ public static Optional findSchemaName(final TableSegment tableSegment) { if (!openQuerySQL.isPresent()) { return Optional.empty(); } - Matcher matcher = FROM_TABLE_PATTERN.matcher(openQuerySQL.get().getText()); - if (!matcher.find()) { - return Optional.empty(); - } - String tableExpression = matcher.group(1); - int delimiterIndex = tableExpression.lastIndexOf('.'); - return -1 == delimiterIndex ? Optional.empty() : Optional.of(QuoteCharacter.unwrapAndTrimText(tableExpression.substring(0, delimiterIndex))); + return EncryptOpenQueryPassThroughSQL.parse(openQuerySQL.get().getText()).getSchemaName(); } } diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java index 03e71a0dd8ee2..13d2ab1677254 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptUpdateAssignmentTokenGenerator.java @@ -27,6 +27,7 @@ import org.apache.shardingsphere.infra.rewrite.sql.token.common.generator.CollectionSQLTokenGenerator; import org.apache.shardingsphere.infra.rewrite.sql.token.common.pojo.SQLToken; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.SimpleTableSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableSegment; import java.util.Collection; @@ -59,7 +60,11 @@ private boolean containsEncryptTable(final Collection simple @Override public Collection generateSQLTokens(final UpdateStatementContext sqlStatementContext) { - return new EncryptAssignmentTokenGenerator(rule, database, sqlStatementContext.getSqlStatement().getDatabaseType()) - .generateSQLTokens(sqlStatementContext.getTablesContext(), sqlStatementContext.getSqlStatement().getSetAssignment(), sqlStatementContext.getSqlStatement().getTable()); + EncryptAssignmentTokenGenerator tokenGenerator = new EncryptAssignmentTokenGenerator(rule, database, sqlStatementContext.getSqlStatement().getDatabaseType()); + TableSegment table = sqlStatementContext.getSqlStatement().getTable(); + if (EncryptOpenQueryUtils.isOpenQueryFunctionTable(table)) { + return tokenGenerator.generateSQLTokens(sqlStatementContext.getTablesContext(), sqlStatementContext.getSqlStatement().getSetAssignment(), table); + } + return tokenGenerator.generateSQLTokens(sqlStatementContext.getTablesContext(), sqlStatementContext.getSqlStatement().getSetAssignment()); } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java index a78661f215d2e..36112dab50abc 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java @@ -123,6 +123,21 @@ void assertDecorateWithOpenQueryEncryptTable() { verify(sqlRewriteContext).addSQLTokenGenerators(any()); } + @Test + void assertDecorateWithOpenQueryUnrelatedTableAndUnsupportedPassThroughShape() { + SQLRewriteContext sqlRewriteContext = mock(SQLRewriteContext.class); + UpdateStatementContext updateStatementContext = mock(UpdateStatementContext.class, RETURNS_DEEP_STUBS); + when(updateStatementContext.getTablesContext().getSimpleTables()).thenReturn(Collections.emptyList()); + when(updateStatementContext.getTablesContext().getTableNames()).thenReturn(Collections.emptyList()); + UpdateStatement updateStatement = mock(UpdateStatement.class); + when(updateStatement.getTable()).thenReturn(createOpenQueryTableSegmentWithJoinOnUnrelatedTable()); + when(updateStatementContext.getSqlStatement()).thenReturn(updateStatement); + when(sqlRewriteContext.getSqlStatementContext()).thenReturn(updateStatementContext); + EncryptRule encryptRule = new EncryptRule("foo_db", getEncryptRuleConfiguration(false)); + decorator.decorate(encryptRule, mock(ConfigurationProperties.class), sqlRewriteContext, mock(RouteContext.class)); + verify(sqlRewriteContext, never()).addSQLTokenGenerators(any()); + } + private EncryptRuleConfiguration getEncryptRuleConfiguration(final boolean dropEncryptTable) { EncryptColumnRuleConfiguration columnConfig = new EncryptColumnRuleConfiguration("pwd", new EncryptColumnItemRuleConfiguration("pwd_cipher", "standard_encryptor")); EncryptTableRuleConfiguration tableConfig = new EncryptTableRuleConfiguration("t_encrypt", Collections.singleton(columnConfig)); @@ -140,4 +155,12 @@ private FunctionTableSegment createOpenQueryTableSegment() { functionSegment.getParameters().add(new LiteralExpressionSegment(0, 0, "SELECT pwd FROM foo_schema.t_encrypt")); return new FunctionTableSegment(0, 0, functionSegment); } + + private FunctionTableSegment createOpenQueryTableSegmentWithJoinOnUnrelatedTable() { + FunctionSegment functionSegment = new FunctionSegment(0, 0, "OPENQUERY", + "OPENQUERY (foo_server, 'SELECT col FROM db.schema.bar_tbl JOIN foo_schema.t_encrypt ON bar_tbl.id = t_encrypt.id')"); + functionSegment.getParameters().add(new ColumnSegment(0, 0, new IdentifierValue("foo_server"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(0, 0, "SELECT col FROM db.schema.bar_tbl JOIN foo_schema.t_encrypt ON bar_tbl.id = t_encrypt.id")); + return new FunctionTableSegment(0, 0, functionSegment); + } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java index 6730690f8b5e1..b68d233f28484 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java @@ -22,6 +22,9 @@ import org.apache.shardingsphere.database.connector.core.type.DatabaseType; import org.apache.shardingsphere.database.connector.core.type.DatabaseTypeRegistry; import org.apache.shardingsphere.encrypt.rule.EncryptRule; +import org.apache.shardingsphere.encrypt.rule.column.item.AssistedQueryColumnItem; +import org.apache.shardingsphere.encrypt.rule.column.item.LikeQueryColumnItem; +import org.apache.shardingsphere.encrypt.spi.EncryptAlgorithm; import org.apache.shardingsphere.encrypt.rule.column.EncryptColumn; import org.apache.shardingsphere.encrypt.rule.table.EncryptTable; import org.apache.shardingsphere.infra.binder.context.segment.table.TablesContext; @@ -49,6 +52,7 @@ import org.mockito.MockedConstruction; import org.mockito.junit.jupiter.MockitoExtension; +import java.util.Arrays; import java.util.Collection; import java.util.Collections; import java.util.Iterator; @@ -58,8 +62,10 @@ import static org.hamcrest.MatcherAssert.assertThat; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.Mockito.RETURNS_DEEP_STUBS; -import static org.mockito.Mockito.mock; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.mockConstruction; +import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @ExtendWith(MockitoExtension.class) @@ -160,11 +166,60 @@ void assertGenerateSQLTokenWithOpenQueryLiteralExpressionSegment() { assertThat(iterator.next().toString(), is("'SELECT group_name_cipher FROM dbo.Department WHERE DepartmentID = 4'")); } + @Test + void assertGenerateSQLTokenWithOpenQueryMultipleAssignments() { + ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); + when(database.getName()).thenReturn("foo_db"); + tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryMultiColumnEncryptRule(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + FunctionTableSegment openQueryTable = createMultiColumnOpenQueryTableSegment(); + SetAssignmentSegment multiAssignment = createMultiColumnSetAssignment(); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + Collection actual = tokenGenerator.generateSQLTokens(tablesContext, multiAssignment, openQueryTable); + assertThat(actual.size(), is(3)); + Iterator iterator = actual.iterator(); + assertThat(iterator.next().toString(), is("group_name_cipher = 'groupEncryptValue'")); + assertThat(iterator.next().toString(), is("dept_code_cipher = 'deptEncryptValue'")); + assertThat(iterator.next().toString(), is("'SELECT group_name_cipher, dept_code_cipher FROM dbo.Department WHERE DepartmentID = 4'")); + } + + @Test + void assertGenerateSQLTokenWithOpenQueryPreservesWhereClauseColumnRef() { + ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); + when(database.getName()).thenReturn("foo_db"); + tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRule(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(new LiteralExpressionSegment(124, 144, "Sales and Marketing")); + when(assignmentSegment.getStopIndex()).thenReturn(144); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + Collection actual = tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithColumnInWhere()); + assertThat(actual.size(), is(2)); + Iterator iterator = actual.iterator(); + iterator.next(); + assertThat(iterator.next().toString(), is("'SELECT group_name_cipher FROM dbo.Department WHERE GroupName IS NOT NULL'")); + } + + @Test + void assertGenerateSQLTokenWithOpenQueryDerivedColumns() { + ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); + when(database.getName()).thenReturn("foo_db"); + tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryDerivedColumnsEncryptRule(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + FunctionTableSegment openQueryTable = createDerivedColumnsOpenQueryTableSegment(); + SetAssignmentSegment multiAssignment = createDerivedColumnsSetAssignment(); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + Collection actual = tokenGenerator.generateSQLTokens(tablesContext, multiAssignment, openQueryTable); + assertThat(actual.size(), is(3)); + Iterator iterator = actual.iterator(); + assertThat(iterator.next().toString(), is("group_name_cipher = 'groupEncryptValue'")); + assertThat(iterator.next().toString(), is("remark_cipher = 'remarkCipherValue', assisted_query_remark = 'assistedValue', like_query_remark = 'likeValue'")); + assertThat(iterator.next().toString(), is("'SELECT group_name_cipher, remark_cipher, assisted_query_remark, like_query_remark FROM dbo.Department WHERE DepartmentID = 4'")); + } + private EncryptRule mockOpenQueryEncryptRule() { EncryptRule result = mock(EncryptRule.class); EncryptTable encryptTable = mock(EncryptTable.class); EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); - when(result.getAllTableNames()).thenReturn(Collections.singleton("Department")); when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); when(encryptTable.getTable()).thenReturn("Department"); @@ -182,4 +237,116 @@ private FunctionTableSegment createOpenQueryTableSegment() { functionSegment.getParameters().add(new LiteralExpressionSegment(34, 95, "SELECT GroupName FROM dbo.Department WHERE DepartmentID = 4")); return new FunctionTableSegment(7, 106, functionSegment); } + + private FunctionTableSegment createOpenQueryTableSegmentWithColumnInWhere() { + FunctionSegment functionSegment = new FunctionSegment(7, 115, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE GroupName IS NOT NULL')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 104, "SELECT GroupName FROM dbo.Department WHERE GroupName IS NOT NULL")); + return new FunctionTableSegment(7, 115, functionSegment); + } + + private EncryptRule mockOpenQueryMultiColumnEncryptRule() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn groupNameColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + EncryptColumn deptCodeColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); + when(encryptTable.isEncryptColumn("DeptCode")).thenReturn(true); + when(encryptTable.getTable()).thenReturn("Department"); + when(encryptTable.getEncryptColumn("GroupName")).thenReturn(groupNameColumn); + when(encryptTable.getEncryptColumn("DeptCode")).thenReturn(deptCodeColumn); + when(groupNameColumn.getName()).thenReturn("GroupName"); + when(groupNameColumn.getCipher().getName()).thenReturn("group_name_cipher"); + when(groupNameColumn.getCipher().encrypt("foo_db", "dbo", "Department", "GroupName", Collections.singletonList("Sales"))) + .thenReturn(Collections.singletonList("groupEncryptValue")); + when(deptCodeColumn.getName()).thenReturn("DeptCode"); + when(deptCodeColumn.getCipher().getName()).thenReturn("dept_code_cipher"); + when(deptCodeColumn.getCipher().encrypt("foo_db", "dbo", "Department", "DeptCode", Collections.singletonList("D001"))) + .thenReturn(Collections.singletonList("deptEncryptValue")); + return result; + } + + private SetAssignmentSegment createMultiColumnSetAssignment() { + ColumnSegment groupNameCol = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + groupNameCol.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + ColumnAssignmentSegment groupNameAssignment = mock(ColumnAssignmentSegment.class); + when(groupNameAssignment.getColumns()).thenReturn(Collections.singletonList(groupNameCol)); + when(groupNameAssignment.getValue()).thenReturn(new LiteralExpressionSegment(124, 128, "Sales")); + when(groupNameAssignment.getStopIndex()).thenReturn(128); + ColumnSegment deptCodeCol = new ColumnSegment(132, 139, new IdentifierValue("DeptCode")); + deptCodeCol.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("DeptCode"), TableSourceType.TEMPORARY_TABLE)); + ColumnAssignmentSegment deptCodeAssignment = mock(ColumnAssignmentSegment.class); + when(deptCodeAssignment.getColumns()).thenReturn(Collections.singletonList(deptCodeCol)); + when(deptCodeAssignment.getValue()).thenReturn(new LiteralExpressionSegment(143, 146, "D001")); + when(deptCodeAssignment.getStopIndex()).thenReturn(146); + SetAssignmentSegment result = mock(SetAssignmentSegment.class); + when(result.getAssignments()).thenReturn(Arrays.asList(groupNameAssignment, deptCodeAssignment)); + return result; + } + + private FunctionTableSegment createMultiColumnOpenQueryTableSegment() { + FunctionSegment functionSegment = new FunctionSegment(7, 110, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName, DeptCode FROM dbo.Department WHERE DepartmentID = 4')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 99, "SELECT GroupName, DeptCode FROM dbo.Department WHERE DepartmentID = 4")); + return new FunctionTableSegment(7, 110, functionSegment); + } + + private EncryptRule mockOpenQueryDerivedColumnsEncryptRule() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn groupNameColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + EncryptColumn remarkColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + EncryptAlgorithm assistedEncryptor = mock(EncryptAlgorithm.class); + EncryptAlgorithm likeEncryptor = mock(EncryptAlgorithm.class); + when(assistedEncryptor.encrypt(eq("note"), any())).thenReturn("assistedValue"); + when(likeEncryptor.encrypt(eq("note"), any())).thenReturn("likeValue"); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); + when(encryptTable.isEncryptColumn("Remark")).thenReturn(true); + when(encryptTable.getTable()).thenReturn("Department"); + when(encryptTable.getEncryptColumn("GroupName")).thenReturn(groupNameColumn); + when(encryptTable.getEncryptColumn("Remark")).thenReturn(remarkColumn); + when(groupNameColumn.getName()).thenReturn("GroupName"); + when(groupNameColumn.getCipher().getName()).thenReturn("group_name_cipher"); + when(groupNameColumn.getCipher().encrypt("foo_db", "dbo", "Department", "GroupName", Collections.singletonList("Sales"))) + .thenReturn(Collections.singletonList("groupEncryptValue")); + when(groupNameColumn.getAssistedQuery()).thenReturn(Optional.empty()); + when(groupNameColumn.getLikeQuery()).thenReturn(Optional.empty()); + when(remarkColumn.getName()).thenReturn("Remark"); + when(remarkColumn.getCipher().getName()).thenReturn("remark_cipher"); + when(remarkColumn.getCipher().encrypt("foo_db", "dbo", "Department", "Remark", Collections.singletonList("note"))) + .thenReturn(Collections.singletonList("remarkCipherValue")); + when(remarkColumn.getAssistedQuery()).thenReturn(Optional.of(new AssistedQueryColumnItem("assisted_query_remark", assistedEncryptor))); + when(remarkColumn.getLikeQuery()).thenReturn(Optional.of(new LikeQueryColumnItem("like_query_remark", likeEncryptor))); + return result; + } + + private SetAssignmentSegment createDerivedColumnsSetAssignment() { + ColumnSegment groupNameCol = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + groupNameCol.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + ColumnAssignmentSegment groupNameAssignment = mock(ColumnAssignmentSegment.class); + when(groupNameAssignment.getColumns()).thenReturn(Collections.singletonList(groupNameCol)); + when(groupNameAssignment.getValue()).thenReturn(new LiteralExpressionSegment(124, 128, "Sales")); + when(groupNameAssignment.getStopIndex()).thenReturn(128); + ColumnSegment remarkCol = new ColumnSegment(132, 137, new IdentifierValue("Remark")); + remarkCol.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("Remark"), TableSourceType.TEMPORARY_TABLE)); + ColumnAssignmentSegment remarkAssignment = mock(ColumnAssignmentSegment.class); + when(remarkAssignment.getColumns()).thenReturn(Collections.singletonList(remarkCol)); + when(remarkAssignment.getValue()).thenReturn(new LiteralExpressionSegment(141, 144, "note")); + when(remarkAssignment.getStopIndex()).thenReturn(144); + SetAssignmentSegment result = mock(SetAssignmentSegment.class); + when(result.getAssignments()).thenReturn(Arrays.asList(groupNameAssignment, remarkAssignment)); + return result; + } + + private FunctionTableSegment createDerivedColumnsOpenQueryTableSegment() { + FunctionSegment functionSegment = new FunctionSegment(7, 108, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName, Remark FROM dbo.Department WHERE DepartmentID = 4')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 97, "SELECT GroupName, Remark FROM dbo.Department WHERE DepartmentID = 4")); + return new FunctionTableSegment(7, 108, functionSegment); + } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java new file mode 100644 index 0000000000000..58fbc035a29d8 --- /dev/null +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java @@ -0,0 +1,146 @@ +/* + * 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.shardingsphere.encrypt.rewrite.token.generator.assignment; + +import org.apache.shardingsphere.encrypt.exception.syntax.UnsupportedEncryptSQLException; +import org.apache.shardingsphere.encrypt.rule.column.EncryptColumn; +import org.apache.shardingsphere.encrypt.rule.column.item.AssistedQueryColumnItem; +import org.apache.shardingsphere.encrypt.rule.column.item.LikeQueryColumnItem; +import org.junit.jupiter.api.Test; + +import java.util.Collections; +import java.util.Optional; + +import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.is; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.RETURNS_DEEP_STUBS; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +class EncryptOpenQueryPassThroughSQLTest { + + @Test + void assertParseWithMultipartTableName() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE foo_col = 1"); + assertThat(actual.getTableName(), is("foo_tbl")); + assertThat(actual.getSchemaName(), is(Optional.of("foo_schema"))); + assertThat(actual.getRemainder(), is(" WHERE foo_col = 1")); + } + + @Test + void assertParseWithDelimitedMultipartTableName() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM [foo_schema].[foo_tbl]"); + assertThat(actual.getTableName(), is("foo_tbl")); + assertThat(actual.getSchemaName(), is(Optional.of("foo_schema"))); + } + + @Test + void assertParsePreservesPredicateOnEncryptedColumn() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE GroupName IS NOT NULL"); + assertThat(actual.getRemainder(), is(" WHERE GroupName IS NOT NULL")); + } + + @Test + void assertParsePreservesStringLiteralContainingColumnNameInPredicate() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE Note = 'GroupName'"); + assertThat(actual.getRemainder(), is(" WHERE Note = 'GroupName'")); + } + + @Test + void assertRewritePreservesPredicateOnEncryptedColumn() { + EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(encryptColumn.getName()).thenReturn("GroupName"); + when(encryptColumn.getCipher().getName()).thenReturn("group_name_cipher"); + when(encryptColumn.getAssistedQuery()).thenReturn(Optional.empty()); + when(encryptColumn.getLikeQuery()).thenReturn(Optional.empty()); + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE GroupName IS NOT NULL"); + String actual = passThroughSQL.rewrite(Collections.singletonList(encryptColumn)); + assertThat(actual, is("SELECT group_name_cipher FROM dbo.Department WHERE GroupName IS NOT NULL")); + } + + @Test + void assertRewritePreservesStringLiteralContainingColumnNameInPredicate() { + EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(encryptColumn.getName()).thenReturn("GroupName"); + when(encryptColumn.getCipher().getName()).thenReturn("group_name_cipher"); + when(encryptColumn.getAssistedQuery()).thenReturn(Optional.empty()); + when(encryptColumn.getLikeQuery()).thenReturn(Optional.empty()); + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE Note = 'GroupName'"); + String actual = passThroughSQL.rewrite(Collections.singletonList(encryptColumn)); + assertThat(actual, is("SELECT group_name_cipher FROM dbo.Department WHERE Note = 'GroupName'")); + } + + @Test + void assertFindTableNameWithThreePartTableName() { + Optional actual = EncryptOpenQueryPassThroughSQL.findTableName("SELECT GroupName FROM db.schema.Department WHERE DepartmentID = 4"); + assertTrue(actual.isPresent()); + assertThat(actual.get(), is("Department")); + } + + @Test + void assertFindTableNameWithJoin() { + Optional actual = EncryptOpenQueryPassThroughSQL.findTableName( + "SELECT GroupName FROM dbo.Department JOIN dbo.Employee ON Department.DepartmentID = Employee.DepartmentID"); + assertTrue(actual.isPresent()); + assertThat(actual.get(), is("Department")); + } + + @Test + void assertParseWithSelectListLiteralExpectsException() { + assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT 'GroupName' FROM dbo.Department")); + } + + @Test + void assertParseWithSelectListExpressionExpectsException() { + assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT UPPER(GroupName) FROM dbo.Department")); + } + + @Test + void assertParseWithSpaceDelimitedIdentifierExpectsException() { + assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM [Human Resources].[Department]")); + } + + @Test + void assertParseWithThreePartTableNameExpectsException() { + assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM db.schema.Department")); + } + + @Test + void assertParseWithJoinExpectsException() { + String passThroughSQL = "SELECT GroupName FROM dbo.Department JOIN dbo.Employee ON Department.DepartmentID = Employee.DepartmentID"; + assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse(passThroughSQL)); + } + + @Test + void assertRewriteWithDerivedColumns() { + EncryptColumn remarkColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + AssistedQueryColumnItem assistedQuery = mock(AssistedQueryColumnItem.class); + LikeQueryColumnItem likeQuery = mock(LikeQueryColumnItem.class); + when(remarkColumn.getName()).thenReturn("Remark"); + when(remarkColumn.getCipher().getName()).thenReturn("remark_cipher"); + when(assistedQuery.getName()).thenReturn("assisted_query_remark"); + when(likeQuery.getName()).thenReturn("like_query_remark"); + when(remarkColumn.getAssistedQuery()).thenReturn(Optional.of(assistedQuery)); + when(remarkColumn.getLikeQuery()).thenReturn(Optional.of(likeQuery)); + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName, Remark FROM dbo.Department WHERE DepartmentID = 4"); + String actual = passThroughSQL.rewrite(Collections.singletonList(remarkColumn)); + assertThat(actual, is("SELECT GroupName, remark_cipher, assisted_query_remark, like_query_remark FROM dbo.Department WHERE DepartmentID = 4")); + } +} diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java index 40af9f1f5f51b..1eaa46f3a406e 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java @@ -28,7 +28,6 @@ import org.apache.shardingsphere.sql.parser.statement.core.value.identifier.IdentifierValue; import org.junit.jupiter.api.Test; -import java.util.Collections; import java.util.Optional; import static org.hamcrest.MatcherAssert.assertThat; @@ -50,7 +49,6 @@ void assertIsOpenQueryFunctionTable() { void assertFindEncryptTable() { EncryptRule rule = mock(EncryptRule.class); EncryptTable encryptTable = mock(EncryptTable.class); - when(rule.getAllTableNames()).thenReturn(Collections.singleton("foo_tbl")); when(rule.findEncryptTable("foo_tbl")).thenReturn(Optional.of(encryptTable)); Optional actual = EncryptOpenQueryUtils.findEncryptTable(rule, createOpenQueryTableSegment("SELECT foo_col FROM foo_schema.foo_tbl")); assertTrue(actual.isPresent()); @@ -66,6 +64,31 @@ void assertFindSchemaName() { assertFalse(EncryptOpenQueryUtils.findSchemaName(createOpenQueryTableSegment("SELECT foo_col FROM foo_tbl")).isPresent()); } + @Test + void assertFindEncryptTableWithDelimitedMultipartTableName() { + EncryptRule rule = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + when(rule.findEncryptTable("foo_tbl")).thenReturn(Optional.of(encryptTable)); + Optional actual = EncryptOpenQueryUtils.findEncryptTable(rule, createOpenQueryTableSegment("SELECT foo_col FROM [foo_schema].[foo_tbl]")); + assertTrue(actual.isPresent()); + assertThat(actual.get(), is(encryptTable)); + } + + @Test + void assertFindEncryptTableWithThreePartUnrelatedTableReturnsEmpty() { + EncryptRule rule = mock(EncryptRule.class); + when(rule.findEncryptTable("Employee")).thenReturn(Optional.empty()); + assertFalse(EncryptOpenQueryUtils.findEncryptTable(rule, createOpenQueryTableSegment("SELECT col FROM db.schema.Employee")).isPresent()); + } + + @Test + void assertFindEncryptTableWithJoinUnrelatedTableReturnsEmpty() { + EncryptRule rule = mock(EncryptRule.class); + when(rule.findEncryptTable("Employee")).thenReturn(Optional.empty()); + assertFalse(EncryptOpenQueryUtils.findEncryptTable(rule, + createOpenQueryTableSegment("SELECT col FROM dbo.Employee JOIN dbo.Department ON Employee.DepartmentID = Department.DepartmentID")).isPresent()); + } + private FunctionTableSegment createOpenQueryTableSegment(final String openQuerySQL) { FunctionSegment functionSegment = new FunctionSegment(0, 0, "OPENQUERY", "OPENQUERY (foo_server, '" + openQuerySQL + "')"); functionSegment.getParameters().add(new ColumnSegment(0, 0, new IdentifierValue("foo_server"))); diff --git a/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java b/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java index a1d7cd5bff7be..c02db3d7446be 100644 --- a/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java +++ b/test/it/rewriter/src/test/java/org/apache/shardingsphere/test/it/rewriter/engine/scenario/EncryptSQLRewriterIT.java @@ -86,7 +86,8 @@ protected Collection mockSchemas(final String schemaName) tables.add(new ShardingSphereTable("Department", Arrays.asList( new ShardingSphereColumn("DepartmentID", Types.INTEGER, false, false, false, true, false, false), new ShardingSphereColumn("Name", Types.VARCHAR, false, false, false, true, false, false), - new ShardingSphereColumn("GroupName", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); + new ShardingSphereColumn("GroupName", Types.VARCHAR, false, false, false, true, false, false), + new ShardingSphereColumn("Remark", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); tables.add(new ShardingSphereTable("StateRegion", Arrays.asList( new ShardingSphereColumn("StateCode", Types.INTEGER, false, false, false, true, false, false), new ShardingSphereColumn("CountryRegionName", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); @@ -115,7 +116,8 @@ protected Collection mockSchemas(final String schemaName) humanResourcesSchemaTables.add(new ShardingSphereTable("Department", Arrays.asList( new ShardingSphereColumn("DepartmentID", Types.INTEGER, false, false, false, true, false, false), new ShardingSphereColumn("Name", Types.VARCHAR, false, false, false, true, false, false), - new ShardingSphereColumn("GroupName", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); + new ShardingSphereColumn("GroupName", Types.VARCHAR, false, false, false, true, false, false), + new ShardingSphereColumn("Remark", Types.VARCHAR, false, false, false, true, false, false)), Collections.emptyList(), Collections.emptyList())); result.add(new ShardingSphereSchema("HumanResources", mock(DatabaseType.class), humanResourcesSchemaTables, Collections.emptyList())); return result; } diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml index 2e8082e6493d2..eb10ebf9bb3ce 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml @@ -119,7 +119,62 @@ - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml index 5e9b3cf3af396..3f7412d141a63 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml @@ -164,6 +164,20 @@ rules: cipher: name: group_name_cipher encryptorName: rewrite_normal_fixture + Name: + cipher: + name: name_cipher + encryptorName: rewrite_normal_fixture + Remark: + cipher: + name: remark_cipher + encryptorName: rewrite_normal_fixture + assistedQuery: + name: assisted_query_remark + encryptorName: rewrite_assisted_query_fixture + likeQuery: + name: like_query_remark + encryptorName: rewrite_like_query_fixture StateRegion: columns: CountryRegionName: From 498acaacf6dc2a85ed66060bc1993988edee285e Mon Sep 17 00:00:00 2001 From: Claire Date: Mon, 20 Jul 2026 23:40:38 +0800 Subject: [PATCH 5/8] update --- RELEASE-NOTES.md | 2 +- .../features/encrypt/limitations.cn.md | 8 + .../features/encrypt/limitations.en.md | 8 + .../EncryptOpenQueryPassThroughSQL.java | 267 ++++++++++++++++-- ...eneratorOpenQueryUnsupportedShapeTest.java | 92 ++++++ .../EncryptAssignmentTokenGeneratorTest.java | 54 +++- .../EncryptOpenQueryPassThroughSQLTest.java | 149 ++++++---- .../assignment/EncryptOpenQueryUtilsTest.java | 21 ++ .../query-with-cipher/dml/update/update.xml | 70 +++-- .../encrypt/config/query-with-cipher.yaml | 24 ++ 10 files changed, 596 insertions(+), 99 deletions(-) create mode 100644 features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java diff --git a/RELEASE-NOTES.md b/RELEASE-NOTES.md index 7a939663fe337..3612eebcbffbf 100644 --- a/RELEASE-NOTES.md +++ b/RELEASE-NOTES.md @@ -76,7 +76,7 @@ 1. Proxy: Support basic Firebird batch operations: create, send, execute, cancel, and release - [#38605](https://github.com/apache/shardingsphere/pull/38605) 1. JDBC & Proxy: Add a check to verify database name naming conventions. - [#38883](https://github.com/apache/shardingsphere/pull/38883) 1. Encrypt: Support SqlServer update statement for Updating data in a remote table by using a linked server when use encrypt feature - [#39122](https://github.com/apache/shardingsphere/pull/39122) -1. Encrypt: Support SqlServer update statement for Updating data in a remote table by using the OPENQUERY function when use encrypt feature - [#39156](https://github.com/apache/shardingsphere/pull/39156) +1. Encrypt: Support encrypt rewrite for SQL Server OPENQUERY UPDATE with narrow SELECT ... FROM shape - [#39156](https://github.com/apache/shardingsphere/pull/39156) 1. Encrypt: Support SqlServer update statement for Specifying a table alias as the target object when use encrypt feature - [#38733](https://github.com/apache/shardingsphere/pull/38733) 1. Encrypt: Support SqlServer update statement for Specifying a view as the target object when use encrypt feature - [#38896](https://github.com/apache/shardingsphere/pull/38896) 1. Encrypt: Support SqlServer for Using the UPDATE statement with information from another table when use encrypt feature - [#38926](https://github.com/apache/shardingsphere/pull/38926) diff --git a/docs/document/content/features/encrypt/limitations.cn.md b/docs/document/content/features/encrypt/limitations.cn.md index bcc7079efea9e..f0169a7518ebc 100644 --- a/docs/document/content/features/encrypt/limitations.cn.md +++ b/docs/document/content/features/encrypt/limitations.cn.md @@ -10,3 +10,11 @@ weight = 2 - 加密字段无法支持计算操作,如:AVG、SUM 以及计算表达式; - 不支持使用 `;` 分隔的多条 SQL 同时执行; - 当投影子查询中包含加密字段时,必须使用别名。 + +## SQL Server OPENQUERY 加密功能 + +`OPENQUERY` 函数的加密改写不支持以下场景: + +- `WHERE` 后引用加密列; +- 物理列名中间包含 `]` 或 `[]`; +- 透传查询中使用 `JOIN`、`CROSS APPLY`、`OUTER APPLY`、`UNION`、`UNION ALL`、`EXCEPT`、`INTERSECT`。 diff --git a/docs/document/content/features/encrypt/limitations.en.md b/docs/document/content/features/encrypt/limitations.en.md index ce6d3a831564e..10ce6f3a308d4 100644 --- a/docs/document/content/features/encrypt/limitations.en.md +++ b/docs/document/content/features/encrypt/limitations.en.md @@ -10,3 +10,11 @@ weight = 2 - Calculation operations are not supported for encrypted fields, such as `AVG`, `SUM`, and computation expressions. - Not support simultaneous execution of multiple SQL statements separated by `;`. - When projection subquery contains encrypt column, you must use alias. + +## SQL Server OPENQUERY encryption + +The following are not supported for encrypt rewrite with the `OPENQUERY` function: + +- Predicates after `WHERE` that reference encrypted columns. +- Physical column names that contain `]` or `[]` in the middle. +- `JOIN`, `CROSS APPLY`, `OUTER APPLY`, `UNION`, `UNION ALL`, `EXCEPT`, and `INTERSECT` in the pass-through query. diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java index fd995dcc0a8f1..5a88a824cbeb4 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java @@ -18,13 +18,13 @@ package org.apache.shardingsphere.encrypt.rewrite.token.generator.assignment; import lombok.Getter; +import org.apache.shardingsphere.database.connector.core.metadata.database.enums.QuoteCharacter; import org.apache.shardingsphere.encrypt.exception.syntax.UnsupportedEncryptSQLException; import org.apache.shardingsphere.encrypt.rule.column.EncryptColumn; import java.util.ArrayList; import java.util.Collection; import java.util.List; -import java.util.Locale; import java.util.Optional; /** @@ -45,6 +45,14 @@ final class EncryptOpenQueryPassThroughSQL { private static final String UNSUPPORTED_JOIN = "OPENQUERY with JOIN statement"; + private static final String UNSUPPORTED_COMMA_TABLE_SOURCE = "OPENQUERY with comma-separated table sources"; + + private static final String UNSUPPORTED_APPLY = "OPENQUERY with APPLY statement"; + + private static final String UNSUPPORTED_SET_OPERATION = "OPENQUERY with set operation"; + + private static final String UNSUPPORTED_ENCRYPTED_PREDICATE = "OPENQUERY with predicate on encrypted column"; + private final String selectList; private final String tableExpression; @@ -112,8 +120,10 @@ static EncryptOpenQueryPassThroughSQL parse(final String passThroughSQL) { * * @param encryptColumns encrypt columns * @return rewritten pass-through SQL + * @throws UnsupportedEncryptSQLException if remainder references an encrypted logic column */ String rewrite(final Collection encryptColumns) { + validateRemainderHasNoEncryptColumnReference(remainder, encryptColumns); List rewrittenItems = new ArrayList<>(); for (String each : splitSelectList(selectList)) { String trimmedItem = each.trim(); @@ -123,6 +133,72 @@ String rewrite(final Collection encryptColumns) { return "SELECT " + String.join(", ", rewrittenItems) + " FROM " + tableExpression + remainder; } + private static void validateRemainderHasNoEncryptColumnReference(final String remainder, final Collection encryptColumns) { + if (remainder.isEmpty()) { + return; + } + for (EncryptColumn each : encryptColumns) { + if (containsLogicColumnReference(remainder, each.getName())) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_ENCRYPTED_PREDICATE); + } + } + } + + private static boolean containsLogicColumnReference(final String sqlFragment, final String logicColumnName) { + int index = 0; + boolean inString = false; + while (index < sqlFragment.length()) { + char current = sqlFragment.charAt(index); + if ('\'' == current) { + if (!inString) { + inString = true; + index++; + continue; + } + if (index + 1 < sqlFragment.length() && '\'' == sqlFragment.charAt(index + 1)) { + index += 2; + continue; + } + inString = false; + index++; + continue; + } + if (inString) { + index++; + continue; + } + if ('[' == current || '"' == current) { + Optional delimitedPart = readDelimitedIdentifierPartIfPresent(sqlFragment, index); + if (!delimitedPart.isPresent()) { + return false; + } + if (delimitedPart.get().getValue().equalsIgnoreCase(logicColumnName)) { + return true; + } + index = delimitedPart.get().getStopIndex(); + continue; + } + if (Character.isLetter(current) || '_' == current) { + int stopIndex = index + 1; + while (stopIndex < sqlFragment.length()) { + char stopChar = sqlFragment.charAt(stopIndex); + if (Character.isLetterOrDigit(stopChar) || '_' == stopChar) { + stopIndex++; + continue; + } + break; + } + if (sqlFragment.substring(index, stopIndex).equalsIgnoreCase(logicColumnName)) { + return true; + } + index = stopIndex; + continue; + } + index++; + } + return false; + } + private static void validateSelectList(final String selectList) { for (String each : splitSelectList(selectList)) { String trimmedItem = each.trim(); @@ -147,6 +223,12 @@ private static void validateColumnIdentifier(final String identifier) { } return; } + if (identifier.startsWith("\"") && identifier.endsWith("\"")) { + if (!isClosedDelimitedIdentifier(identifier, '"')) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SELECT_EXPRESSION); + } + return; + } for (int index = 0; index < identifier.length(); index++) { char current = identifier.charAt(index); if (!Character.isLetterOrDigit(current) && '_' != current) { @@ -156,10 +238,102 @@ private static void validateColumnIdentifier(final String identifier) { } private static void validateRemainder(final String remainder) { - String trimmedRemainder = remainder.trim().toUpperCase(Locale.ENGLISH); - if (trimmedRemainder.startsWith("JOIN ") || trimmedRemainder.contains(" JOIN ")) { + if (remainder.isEmpty()) { + return; + } + validateNoCommaSeparatedTableSource(remainder); + if (containsKeywordOutsideString(remainder, "JOIN")) { throw new UnsupportedEncryptSQLException(UNSUPPORTED_JOIN); } + if (containsKeywordOutsideString(remainder, "APPLY")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_APPLY); + } + if (containsKeywordOutsideString(remainder, "UNION") + || containsKeywordOutsideString(remainder, "EXCEPT") + || containsKeywordOutsideString(remainder, "INTERSECT")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SET_OPERATION); + } + } + + private static void validateNoCommaSeparatedTableSource(final String remainder) { + int index = skipWhitespace(remainder, 0); + if (index >= remainder.length()) { + return; + } + if (',' == remainder.charAt(index)) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_COMMA_TABLE_SOURCE); + } + if (isClauseKeywordAt(remainder, index)) { + return; + } + if (isKeywordAt(remainder, index, "AS")) { + index = skipWhitespace(remainder, index + "AS".length()); + Optional aliasPart = readIdentifierPartIfPresent(remainder, index); + if (!aliasPart.isPresent()) { + return; + } + index = skipWhitespace(remainder, aliasPart.get().getStopIndex()); + } else { + Optional aliasPart = readIdentifierPartIfPresent(remainder, index); + if (!aliasPart.isPresent()) { + return; + } + index = skipWhitespace(remainder, aliasPart.get().getStopIndex()); + } + if (index < remainder.length() && ',' == remainder.charAt(index)) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_COMMA_TABLE_SOURCE); + } + } + + private static boolean isClauseKeywordAt(final String sqlFragment, final int index) { + return isKeywordAt(sqlFragment, index, "WHERE") + || isKeywordAt(sqlFragment, index, "ORDER") + || isKeywordAt(sqlFragment, index, "GROUP") + || isKeywordAt(sqlFragment, index, "HAVING") + || isKeywordAt(sqlFragment, index, "JOIN") + || isKeywordAt(sqlFragment, index, "UNION") + || isKeywordAt(sqlFragment, index, "EXCEPT") + || isKeywordAt(sqlFragment, index, "INTERSECT") + || isKeywordAt(sqlFragment, index, "CROSS") + || isKeywordAt(sqlFragment, index, "OUTER") + || isKeywordAt(sqlFragment, index, "APPLY"); + } + + private static boolean isKeywordAt(final String sqlFragment, final int index, final String keyword) { + return matchesKeyword(sqlFragment, index, keyword) + && isWordBoundary(sqlFragment, index - 1) + && isWordBoundary(sqlFragment, index + keyword.length()); + } + + private static boolean containsKeywordOutsideString(final String sqlFragment, final String keyword) { + int index = 0; + boolean inString = false; + while (index <= sqlFragment.length() - keyword.length()) { + char current = sqlFragment.charAt(index); + if ('\'' == current) { + if (!inString) { + inString = true; + index++; + continue; + } + if (index + 1 < sqlFragment.length() && '\'' == sqlFragment.charAt(index + 1)) { + index += 2; + continue; + } + inString = false; + index++; + continue; + } + if (inString) { + index++; + continue; + } + if (isKeywordAt(sqlFragment, index, keyword)) { + return true; + } + index++; + } + return false; } private static TableReference parseTableReference(final String passThroughSQL, final int startIndex) { @@ -183,16 +357,12 @@ private static IdentifierPart readIdentifierPart(final String passThroughSQL, fi if (index >= passThroughSQL.length()) { throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); } - if ('[' == passThroughSQL.charAt(index)) { - int closeIndex = passThroughSQL.indexOf(']', index + 1); - if (-1 == closeIndex) { - throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); - } - String inner = passThroughSQL.substring(index + 1, closeIndex); - if (inner.contains(" ")) { + Optional delimitedPart = readDelimitedIdentifierPartIfPresent(passThroughSQL, index); + if (delimitedPart.isPresent()) { + if ('[' == passThroughSQL.charAt(index) && delimitedPart.get().getValue().contains(" ")) { throw new UnsupportedEncryptSQLException(UNSUPPORTED_SPACE_DELIMITED_IDENTIFIER); } - return new IdentifierPart(inner, index, closeIndex + 1); + return delimitedPart.get(); } int stopIndex = index; while (stopIndex < passThroughSQL.length()) { @@ -260,12 +430,9 @@ private static Optional readIdentifierPartIfPresent(final String if (index >= passThroughSQL.length()) { return Optional.empty(); } - if ('[' == passThroughSQL.charAt(index)) { - int closeIndex = passThroughSQL.indexOf(']', index + 1); - if (-1 == closeIndex) { - return Optional.empty(); - } - return Optional.of(new IdentifierPart(passThroughSQL.substring(index + 1, closeIndex), index, closeIndex + 1)); + Optional delimitedPart = readDelimitedIdentifierPartIfPresent(passThroughSQL, index); + if (delimitedPart.isPresent()) { + return delimitedPart; } int stopIndex = index; while (stopIndex < passThroughSQL.length()) { @@ -282,6 +449,52 @@ private static Optional readIdentifierPartIfPresent(final String return Optional.of(new IdentifierPart(passThroughSQL.substring(index, stopIndex), index, stopIndex)); } + private static Optional readDelimitedIdentifierPartIfPresent(final String passThroughSQL, final int startIndex) { + if (startIndex >= passThroughSQL.length()) { + return Optional.empty(); + } + char openDelimiter = passThroughSQL.charAt(startIndex); + if ('[' != openDelimiter && '"' != openDelimiter) { + return Optional.empty(); + } + char closeDelimiter = '[' == openDelimiter ? ']' : '"'; + StringBuilder value = new StringBuilder(); + int index = startIndex + 1; + while (index < passThroughSQL.length()) { + char current = passThroughSQL.charAt(index); + if (closeDelimiter == current) { + if (index + 1 < passThroughSQL.length() && closeDelimiter == passThroughSQL.charAt(index + 1)) { + value.append(closeDelimiter); + index += 2; + continue; + } + return Optional.of(new IdentifierPart(value.toString(), startIndex, index + 1)); + } + value.append(current); + index++; + } + return Optional.empty(); + } + + private static boolean isClosedDelimitedIdentifier(final String identifier, final char delimiter) { + if (identifier.length() < 2 || delimiter != identifier.charAt(0) || delimiter != identifier.charAt(identifier.length() - 1)) { + return false; + } + int index = 1; + while (index < identifier.length() - 1) { + char current = identifier.charAt(index); + if (delimiter == current) { + if (index + 1 < identifier.length() - 1 && delimiter == identifier.charAt(index + 1)) { + index += 2; + continue; + } + return false; + } + index++; + } + return true; + } + private static int findFromKeywordIndex(final String passThroughSQL) { Optional result = findFromKeywordIndexIfPresent(passThroughSQL); if (!result.isPresent()) { @@ -342,11 +555,19 @@ private static List splitSelectList(final String selectList) { private static String unwrapIdentifier(final String identifier) { if (identifier.startsWith("[") && identifier.endsWith("]")) { - return identifier.substring(1, identifier.length() - 1); + return unescapeDelimitedContent(identifier.substring(1, identifier.length() - 1), ']'); + } + if (identifier.startsWith("\"") && identifier.endsWith("\"")) { + return unescapeDelimitedContent(identifier.substring(1, identifier.length() - 1), '"'); } return identifier; } + private static String unescapeDelimitedContent(final String content, final char delimiter) { + String escapedDelimiter = String.valueOf(delimiter) + delimiter; + return content.replace(escapedDelimiter, String.valueOf(delimiter)); + } + private static Optional findEncryptColumn(final Collection encryptColumns, final String logicColumnName) { for (EncryptColumn each : encryptColumns) { if (each.getName().equalsIgnoreCase(logicColumnName)) { @@ -357,12 +578,16 @@ private static Optional findEncryptColumn(final Collection result.append(", ").append(optional.getName())); - encryptColumn.getLikeQuery().ifPresent(optional -> result.append(", ").append(optional.getName())); + StringBuilder result = new StringBuilder(quotePhysicalColumnName(encryptColumn.getCipher().getName())); + encryptColumn.getAssistedQuery().ifPresent(optional -> result.append(", ").append(quotePhysicalColumnName(optional.getName()))); + encryptColumn.getLikeQuery().ifPresent(optional -> result.append(", ").append(quotePhysicalColumnName(optional.getName()))); return result.toString(); } + private static String quotePhysicalColumnName(final String physicalColumnName) { + return QuoteCharacter.BRACKETS.wrap(physicalColumnName); + } + @Getter private static final class TableReference { diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java new file mode 100644 index 0000000000000..d0b1acfef726d --- /dev/null +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java @@ -0,0 +1,92 @@ +/* + * 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.shardingsphere.encrypt.rewrite.token.generator.assignment; + +import org.apache.shardingsphere.database.connector.core.type.DatabaseType; +import org.apache.shardingsphere.encrypt.exception.syntax.UnsupportedEncryptSQLException; +import org.apache.shardingsphere.encrypt.rule.EncryptRule; +import org.apache.shardingsphere.encrypt.rule.column.EncryptColumn; +import org.apache.shardingsphere.encrypt.rule.table.EncryptTable; +import org.apache.shardingsphere.infra.binder.context.segment.table.TablesContext; +import org.apache.shardingsphere.infra.metadata.database.ShardingSphereDatabase; +import org.apache.shardingsphere.infra.spi.type.typed.TypedSPILoader; +import org.apache.shardingsphere.sql.parser.statement.core.enums.TableSourceType; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.ColumnAssignmentSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.SetAssignmentSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.column.ColumnSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.FunctionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.bound.ColumnSegmentBoundInfo; +import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.FunctionTableSegment; +import org.apache.shardingsphere.sql.parser.statement.core.value.identifier.IdentifierValue; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Answers; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import java.util.Collections; +import java.util.Optional; + +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.Mockito.RETURNS_DEEP_STUBS; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +@ExtendWith(MockitoExtension.class) +class EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest { + + @Mock(answer = Answers.RETURNS_DEEP_STUBS) + private TablesContext tablesContext; + + @Mock(answer = Answers.RETURNS_DEEP_STUBS) + private ColumnAssignmentSegment assignmentSegment; + + @Mock(answer = Answers.RETURNS_DEEP_STUBS) + private SetAssignmentSegment setAssignmentSegment; + + @Test + void assertGenerateSQLTokenWithCommaTableSourcesExpectsException() { + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(setAssignmentSegment.getAssignments()).thenReturn(Collections.singleton(assignmentSegment)); + EncryptAssignmentTokenGenerator tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRule(), mock(ShardingSphereDatabase.class), + TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + assertThrows(UnsupportedEncryptSQLException.class, + () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithCommaTableSources())); + } + + private EncryptRule mockOpenQueryEncryptRule() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); + when(encryptTable.getEncryptColumn("GroupName")).thenReturn(encryptColumn); + return result; + } + + private FunctionTableSegment createOpenQueryTableSegmentWithCommaTableSources() { + FunctionSegment functionSegment = new FunctionSegment(7, 95, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department, dbo.Other')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 84, "SELECT GroupName FROM dbo.Department, dbo.Other")); + return new FunctionTableSegment(7, 95, functionSegment); + } +} diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java index b68d233f28484..e7127b70ab485 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java @@ -21,6 +21,7 @@ import org.apache.shardingsphere.database.connector.core.metadata.database.metadata.DialectDatabaseMetaData; import org.apache.shardingsphere.database.connector.core.type.DatabaseType; import org.apache.shardingsphere.database.connector.core.type.DatabaseTypeRegistry; +import org.apache.shardingsphere.encrypt.exception.syntax.UnsupportedEncryptSQLException; import org.apache.shardingsphere.encrypt.rule.EncryptRule; import org.apache.shardingsphere.encrypt.rule.column.item.AssistedQueryColumnItem; import org.apache.shardingsphere.encrypt.rule.column.item.LikeQueryColumnItem; @@ -60,6 +61,7 @@ import static org.hamcrest.Matchers.is; import static org.hamcrest.MatcherAssert.assertThat; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.Mockito.RETURNS_DEEP_STUBS; import static org.mockito.ArgumentMatchers.any; @@ -163,7 +165,7 @@ void assertGenerateSQLTokenWithOpenQueryLiteralExpressionSegment() { Iterator iterator = actual.iterator(); assertThat(actual.size(), is(2)); assertThat(iterator.next().toString(), is("group_name_cipher = 'encryptValue'")); - assertThat(iterator.next().toString(), is("'SELECT group_name_cipher FROM dbo.Department WHERE DepartmentID = 4'")); + assertThat(iterator.next().toString(), is("'SELECT [group_name_cipher] FROM dbo.Department WHERE DepartmentID = 4'")); } @Test @@ -179,11 +181,11 @@ void assertGenerateSQLTokenWithOpenQueryMultipleAssignments() { Iterator iterator = actual.iterator(); assertThat(iterator.next().toString(), is("group_name_cipher = 'groupEncryptValue'")); assertThat(iterator.next().toString(), is("dept_code_cipher = 'deptEncryptValue'")); - assertThat(iterator.next().toString(), is("'SELECT group_name_cipher, dept_code_cipher FROM dbo.Department WHERE DepartmentID = 4'")); + assertThat(iterator.next().toString(), is("'SELECT [group_name_cipher], [dept_code_cipher] FROM dbo.Department WHERE DepartmentID = 4'")); } @Test - void assertGenerateSQLTokenWithOpenQueryPreservesWhereClauseColumnRef() { + void assertGenerateSQLTokenWithOpenQueryEncryptedColumnInWhereExpectsException() { ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); when(database.getName()).thenReturn("foo_db"); tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRule(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); @@ -193,11 +195,26 @@ void assertGenerateSQLTokenWithOpenQueryPreservesWhereClauseColumnRef() { when(assignmentSegment.getValue()).thenReturn(new LiteralExpressionSegment(124, 144, "Sales and Marketing")); when(assignmentSegment.getStopIndex()).thenReturn(144); when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); - Collection actual = tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithColumnInWhere()); + assertThrows(UnsupportedEncryptSQLException.class, + () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithColumnInWhere())); + } + + @Test + void assertGenerateSQLTokenWithOpenQuerySpaceDelimitedPhysicalColumnName() { + ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); + when(database.getName()).thenReturn("foo_db"); + tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQuerySpaceDelimitedPhysicalColumnEncryptRule(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + ColumnSegment columnSegment = new ColumnSegment(112, 123, new IdentifierValue("SecureLabel")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("SecureLabel"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(new LiteralExpressionSegment(127, 134, "secret")); + when(assignmentSegment.getStopIndex()).thenReturn(134); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + Collection actual = tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQuerySecureLabelTableSegment()); assertThat(actual.size(), is(2)); Iterator iterator = actual.iterator(); - iterator.next(); - assertThat(iterator.next().toString(), is("'SELECT group_name_cipher FROM dbo.Department WHERE GroupName IS NOT NULL'")); + assertThat(iterator.next().toString(), is("cipher name = 'encryptValue'")); + assertThat(iterator.next().toString(), is("'SELECT [cipher name] FROM dbo.Department WHERE DepartmentID = 4'")); } @Test @@ -213,7 +230,7 @@ void assertGenerateSQLTokenWithOpenQueryDerivedColumns() { Iterator iterator = actual.iterator(); assertThat(iterator.next().toString(), is("group_name_cipher = 'groupEncryptValue'")); assertThat(iterator.next().toString(), is("remark_cipher = 'remarkCipherValue', assisted_query_remark = 'assistedValue', like_query_remark = 'likeValue'")); - assertThat(iterator.next().toString(), is("'SELECT group_name_cipher, remark_cipher, assisted_query_remark, like_query_remark FROM dbo.Department WHERE DepartmentID = 4'")); + assertThat(iterator.next().toString(), is("'SELECT [group_name_cipher], [remark_cipher], [assisted_query_remark], [like_query_remark] FROM dbo.Department WHERE DepartmentID = 4'")); } private EncryptRule mockOpenQueryEncryptRule() { @@ -231,6 +248,21 @@ private EncryptRule mockOpenQueryEncryptRule() { return result; } + private EncryptRule mockOpenQuerySpaceDelimitedPhysicalColumnEncryptRule() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("SecureLabel")).thenReturn(true); + when(encryptTable.getTable()).thenReturn("Department"); + when(encryptTable.getEncryptColumn("SecureLabel")).thenReturn(encryptColumn); + when(encryptColumn.getName()).thenReturn("SecureLabel"); + when(encryptColumn.getCipher().getName()).thenReturn("cipher name"); + when(encryptColumn.getCipher().encrypt("foo_db", "dbo", "Department", "SecureLabel", Collections.singletonList("secret"))) + .thenReturn(Collections.singletonList("encryptValue")); + return result; + } + private FunctionTableSegment createOpenQueryTableSegment() { FunctionSegment functionSegment = new FunctionSegment(7, 106, "OPENQUERY", "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE DepartmentID = 4')"); functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); @@ -246,6 +278,14 @@ private FunctionTableSegment createOpenQueryTableSegmentWithColumnInWhere() { return new FunctionTableSegment(7, 115, functionSegment); } + private FunctionTableSegment createOpenQuerySecureLabelTableSegment() { + FunctionSegment functionSegment = new FunctionSegment(7, 108, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT SecureLabel FROM dbo.Department WHERE DepartmentID = 4')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 97, "SELECT SecureLabel FROM dbo.Department WHERE DepartmentID = 4")); + return new FunctionTableSegment(7, 108, functionSegment); + } + private EncryptRule mockOpenQueryMultiColumnEncryptRule() { EncryptRule result = mock(EncryptRule.class); EncryptTable encryptTable = mock(EncryptTable.class); diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java index 58fbc035a29d8..ba128b45d0f17 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java @@ -22,9 +22,13 @@ import org.apache.shardingsphere.encrypt.rule.column.item.AssistedQueryColumnItem; import org.apache.shardingsphere.encrypt.rule.column.item.LikeQueryColumnItem; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; import java.util.Collections; import java.util.Optional; +import java.util.stream.Stream; import static org.hamcrest.MatcherAssert.assertThat; import static org.hamcrest.Matchers.is; @@ -52,95 +56,140 @@ void assertParseWithDelimitedMultipartTableName() { } @Test - void assertParsePreservesPredicateOnEncryptedColumn() { - EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE GroupName IS NOT NULL"); - assertThat(actual.getRemainder(), is(" WHERE GroupName IS NOT NULL")); + void assertParseWithDoubleQuotedMultipartTableName() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM \"foo_schema\".\"foo_tbl\""); + assertThat(actual.getTableName(), is("foo_tbl")); + assertThat(actual.getSchemaName(), is(Optional.of("foo_schema"))); + assertThat(actual.getTableExpression(), is("\"foo_schema\".\"foo_tbl\"")); } @Test - void assertParsePreservesStringLiteralContainingColumnNameInPredicate() { - EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE Note = 'GroupName'"); - assertThat(actual.getRemainder(), is(" WHERE Note = 'GroupName'")); + void assertParseKeepsLogicColumnInWhereClause() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE foo_col IS NOT NULL"); + assertThat(actual.getRemainder(), is(" WHERE foo_col IS NOT NULL")); } @Test - void assertRewritePreservesPredicateOnEncryptedColumn() { - EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); - when(encryptColumn.getName()).thenReturn("GroupName"); - when(encryptColumn.getCipher().getName()).thenReturn("group_name_cipher"); - when(encryptColumn.getAssistedQuery()).thenReturn(Optional.empty()); - when(encryptColumn.getLikeQuery()).thenReturn(Optional.empty()); - EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE GroupName IS NOT NULL"); - String actual = passThroughSQL.rewrite(Collections.singletonList(encryptColumn)); - assertThat(actual, is("SELECT group_name_cipher FROM dbo.Department WHERE GroupName IS NOT NULL")); + void assertParseKeepsColumnNameInsideStringLiteral() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE note_col = 'foo_col'"); + assertThat(actual.getRemainder(), is(" WHERE note_col = 'foo_col'")); } @Test - void assertRewritePreservesStringLiteralContainingColumnNameInPredicate() { - EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); - when(encryptColumn.getName()).thenReturn("GroupName"); - when(encryptColumn.getCipher().getName()).thenReturn("group_name_cipher"); - when(encryptColumn.getAssistedQuery()).thenReturn(Optional.empty()); - when(encryptColumn.getLikeQuery()).thenReturn(Optional.empty()); - EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM dbo.Department WHERE Note = 'GroupName'"); - String actual = passThroughSQL.rewrite(Collections.singletonList(encryptColumn)); - assertThat(actual, is("SELECT group_name_cipher FROM dbo.Department WHERE Note = 'GroupName'")); + void assertParseDoesNotRejectCommaInsideInList() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE id_col IN (1, 2, 4)"); + assertThat(actual.getRemainder(), is(" WHERE id_col IN (1, 2, 4)")); } @Test - void assertFindTableNameWithThreePartTableName() { - Optional actual = EncryptOpenQueryPassThroughSQL.findTableName("SELECT GroupName FROM db.schema.Department WHERE DepartmentID = 4"); - assertTrue(actual.isPresent()); - assertThat(actual.get(), is("Department")); + void assertParseDoesNotRejectSetOperationKeywordInsideStringLiteral() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE note_col = 'UNION ALL'"); + assertThat(actual.getRemainder(), is(" WHERE note_col = 'UNION ALL'")); } - @Test - void assertFindTableNameWithJoin() { - Optional actual = EncryptOpenQueryPassThroughSQL.findTableName( - "SELECT GroupName FROM dbo.Department JOIN dbo.Employee ON Department.DepartmentID = Employee.DepartmentID"); + @ParameterizedTest(name = "{0}") + @MethodSource("findTableNameArguments") + void assertFindTableNameDiscoversTableWithoutShapeValidation(final String scenario, final String passThroughSQL, final String expectedTableName) { + Optional actual = EncryptOpenQueryPassThroughSQL.findTableName(passThroughSQL); assertTrue(actual.isPresent()); - assertThat(actual.get(), is("Department")); + assertThat(actual.get(), is(expectedTableName)); } - @Test - void assertParseWithSelectListLiteralExpectsException() { - assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT 'GroupName' FROM dbo.Department")); + private static Stream findTableNameArguments() { + return Stream.of( + Arguments.of("double quoted multipart", "SELECT foo_col FROM \"foo_schema\".\"foo_tbl\" WHERE id_col = 1", "foo_tbl"), + Arguments.of("escaped bracket", "SELECT foo_col FROM [foo_schema].[foo]]tbl]", "foo]tbl"), + Arguments.of("three-part table", "SELECT foo_col FROM foo_db.foo_schema.foo_tbl WHERE id_col = 4", "foo_tbl"), + Arguments.of("join", "SELECT foo_col FROM foo_schema.foo_tbl JOIN foo_schema.bar_tbl ON foo_tbl.id_col = bar_tbl.id_col", "foo_tbl")); + } + + @ParameterizedTest(name = "{0}") + @MethodSource("unsupportedShapeArguments") + void assertParseRejectsUnsupportedShape(final String scenario, final String passThroughSQL) { + assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse(passThroughSQL)); + } + + private static Stream unsupportedShapeArguments() { + return Stream.of( + Arguments.of("select list literal", "SELECT 'foo_col' FROM foo_schema.foo_tbl"), + Arguments.of("select list expression", "SELECT UPPER(foo_col) FROM foo_schema.foo_tbl"), + Arguments.of("space-delimited identifier", "SELECT foo_col FROM [Human Resources].[foo_tbl]"), + Arguments.of("three-part table", "SELECT foo_col FROM foo_db.foo_schema.foo_tbl"), + Arguments.of("join", "SELECT foo_col FROM foo_schema.foo_tbl JOIN foo_schema.bar_tbl ON foo_tbl.id_col = bar_tbl.id_col"), + Arguments.of("comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl, foo_schema.bar_tbl"), + Arguments.of("alias then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl AS t, foo_schema.bar_tbl"), + Arguments.of("cross apply", "SELECT foo_col FROM foo_schema.foo_tbl CROSS APPLY foo_schema.fn_bar(foo_tbl.id_col)"), + Arguments.of("outer apply", "SELECT foo_col FROM foo_schema.foo_tbl OUTER APPLY foo_schema.fn_bar(foo_tbl.id_col)"), + Arguments.of("union", "SELECT foo_col FROM foo_schema.foo_tbl UNION SELECT foo_col FROM foo_schema.bar_tbl"), + Arguments.of("union all", "SELECT foo_col FROM foo_schema.foo_tbl UNION ALL SELECT foo_col FROM foo_schema.bar_tbl"), + Arguments.of("except", "SELECT foo_col FROM foo_schema.foo_tbl EXCEPT SELECT foo_col FROM foo_schema.bar_tbl"), + Arguments.of("intersect", "SELECT foo_col FROM foo_schema.foo_tbl INTERSECT SELECT foo_col FROM foo_schema.bar_tbl")); } @Test - void assertParseWithSelectListExpressionExpectsException() { - assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT UPPER(GroupName) FROM dbo.Department")); + void assertRewriteWithDoubleQuotedMultipartTableName() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM \"foo_schema\".\"foo_tbl\" WHERE id_col = 1"); + String actual = passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo_col_cipher"))); + assertThat(actual, is("SELECT [foo_col_cipher] FROM \"foo_schema\".\"foo_tbl\" WHERE id_col = 1")); } @Test - void assertParseWithSpaceDelimitedIdentifierExpectsException() { - assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM [Human Resources].[Department]")); + void assertRewriteKeepsColumnNameInsideStringLiteral() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE note_col = 'foo_col'"); + String actual = passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo_col_cipher"))); + assertThat(actual, is("SELECT [foo_col_cipher] FROM foo_schema.foo_tbl WHERE note_col = 'foo_col'")); } @Test - void assertParseWithThreePartTableNameExpectsException() { - assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName FROM db.schema.Department")); + void assertRewriteWithEncryptedColumnInWhereExpectsException() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE foo_col IS NOT NULL"); + assertThrows(UnsupportedEncryptSQLException.class, () -> passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo_col_cipher")))); } @Test - void assertParseWithJoinExpectsException() { - String passThroughSQL = "SELECT GroupName FROM dbo.Department JOIN dbo.Employee ON Department.DepartmentID = Employee.DepartmentID"; - assertThrows(UnsupportedEncryptSQLException.class, () -> EncryptOpenQueryPassThroughSQL.parse(passThroughSQL)); + void assertRewriteWithDelimitedEncryptedColumnInWhereExpectsException() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE [foo_col] IS NOT NULL"); + assertThrows(UnsupportedEncryptSQLException.class, () -> passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo_col_cipher")))); } @Test void assertRewriteWithDerivedColumns() { - EncryptColumn remarkColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + EncryptColumn remarkColumn = createEncryptColumn("Remark", "remark_cipher"); AssistedQueryColumnItem assistedQuery = mock(AssistedQueryColumnItem.class); LikeQueryColumnItem likeQuery = mock(LikeQueryColumnItem.class); - when(remarkColumn.getName()).thenReturn("Remark"); - when(remarkColumn.getCipher().getName()).thenReturn("remark_cipher"); when(assistedQuery.getName()).thenReturn("assisted_query_remark"); when(likeQuery.getName()).thenReturn("like_query_remark"); when(remarkColumn.getAssistedQuery()).thenReturn(Optional.of(assistedQuery)); when(remarkColumn.getLikeQuery()).thenReturn(Optional.of(likeQuery)); - EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName, Remark FROM dbo.Department WHERE DepartmentID = 4"); + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT GroupName, Remark FROM foo_schema.foo_tbl WHERE id_col = 4"); String actual = passThroughSQL.rewrite(Collections.singletonList(remarkColumn)); - assertThat(actual, is("SELECT GroupName, remark_cipher, assisted_query_remark, like_query_remark FROM dbo.Department WHERE DepartmentID = 4")); + assertThat(actual, is("SELECT GroupName, [remark_cipher], [assisted_query_remark], [like_query_remark] FROM foo_schema.foo_tbl WHERE id_col = 4")); + } + + @ParameterizedTest(name = "{0}") + @MethodSource("quotedPhysicalColumnNameArguments") + void assertRewriteQuotesPhysicalColumnName(final String scenario, final String physicalColumnName, final String expectedQuotedColumn) { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE id_col = 1"); + String actual = passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", physicalColumnName))); + assertThat(actual, is("SELECT " + expectedQuotedColumn + " FROM foo_schema.foo_tbl WHERE id_col = 1")); + } + + private static Stream quotedPhysicalColumnNameArguments() { + return Stream.of( + Arguments.of("space", "cipher name", "[cipher name]"), + Arguments.of("reserved word", "order", "[order]"), + Arguments.of("asterisk", "foo*bar", "[foo*bar]"), + Arguments.of("slash", "foo/bar", "[foo/bar]"), + Arguments.of("backslash", "foo\\bar", "[foo\\bar]"), + Arguments.of("left bracket", "foo[bar", "[foo[bar]")); + } + + private EncryptColumn createEncryptColumn(final String logicColumnName, final String physicalColumnName) { + EncryptColumn result = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.getName()).thenReturn(logicColumnName); + when(result.getCipher().getName()).thenReturn(physicalColumnName); + when(result.getAssistedQuery()).thenReturn(Optional.empty()); + when(result.getLikeQuery()).thenReturn(Optional.empty()); + return result; } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java index 1eaa46f3a406e..827ba33c90afb 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.java @@ -74,6 +74,16 @@ void assertFindEncryptTableWithDelimitedMultipartTableName() { assertThat(actual.get(), is(encryptTable)); } + @Test + void assertFindEncryptTableWithDoubleQuotedMultipartTableName() { + EncryptRule rule = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + when(rule.findEncryptTable("foo_tbl")).thenReturn(Optional.of(encryptTable)); + Optional actual = EncryptOpenQueryUtils.findEncryptTable(rule, createOpenQueryTableSegment("SELECT foo_col FROM \"foo_schema\".\"foo_tbl\"")); + assertTrue(actual.isPresent()); + assertThat(actual.get(), is(encryptTable)); + } + @Test void assertFindEncryptTableWithThreePartUnrelatedTableReturnsEmpty() { EncryptRule rule = mock(EncryptRule.class); @@ -89,6 +99,17 @@ void assertFindEncryptTableWithJoinUnrelatedTableReturnsEmpty() { createOpenQueryTableSegment("SELECT col FROM dbo.Employee JOIN dbo.Department ON Employee.DepartmentID = Department.DepartmentID")).isPresent()); } + @Test + void assertFindEncryptTableWithCommaTableSourcesStillDiscoversEncryptTable() { + EncryptRule rule = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + when(rule.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + Optional actual = EncryptOpenQueryUtils.findEncryptTable(rule, + createOpenQueryTableSegment("SELECT GroupName FROM dbo.Department, dbo.Other")); + assertTrue(actual.isPresent()); + assertThat(actual.get(), is(encryptTable)); + } + private FunctionTableSegment createOpenQueryTableSegment(final String openQuerySQL) { FunctionSegment functionSegment = new FunctionSegment(0, 0, "OPENQUERY", "OPENQUERY (foo_server, '" + openQuerySQL + "')"); functionSegment.getParameters().add(new ColumnSegment(0, 0, new IdentifierValue("foo_server"))); diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml index eb10ebf9bb3ce..7f2b3c8771164 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml @@ -117,62 +117,92 @@ - + - + - + - + - + - + - + - - - - - - - - - - - + - + - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml index 3f7412d141a63..0bc50290cc399 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml @@ -178,6 +178,30 @@ rules: likeQuery: name: like_query_remark encryptorName: rewrite_like_query_fixture + SecureLabel: + cipher: + name: cipher name + encryptorName: rewrite_normal_fixture + OrderLabel: + cipher: + name: order + encryptorName: rewrite_normal_fixture + AsteriskLabel: + cipher: + name: foo*bar + encryptorName: rewrite_normal_fixture + SlashLabel: + cipher: + name: foo/bar + encryptorName: rewrite_normal_fixture + BackslashLabel: + cipher: + name: "foo\\bar" + encryptorName: rewrite_normal_fixture + LeftBracketLabel: + cipher: + name: foo[bar + encryptorName: rewrite_normal_fixture StateRegion: columns: CountryRegionName: From c8cfb85aeb07652b23ea773a01cedc99b5101f57 Mon Sep 17 00:00:00 2001 From: Claire Date: Tue, 21 Jul 2026 21:52:55 +0800 Subject: [PATCH 6/8] update --- .../features/encrypt/limitations.cn.md | 19 ++++- .../features/encrypt/limitations.en.md | 19 ++++- .../EncryptAssignmentTokenGenerator.java | 41 ++++++----- .../EncryptOpenQueryPassThroughSQL.java | 5 ++ .../encrypt/rule/table/EncryptTable.java | 10 +++ ...eneratorOpenQueryUnsupportedShapeTest.java | 72 +++++++++++++++++++ .../EncryptAssignmentTokenGeneratorTest.java | 49 +++++++++++++ .../EncryptOpenQueryPassThroughSQLTest.java | 7 ++ .../query-with-cipher/dml/update/update.xml | 10 +++ 9 files changed, 208 insertions(+), 24 deletions(-) diff --git a/docs/document/content/features/encrypt/limitations.cn.md b/docs/document/content/features/encrypt/limitations.cn.md index f0169a7518ebc..261287112cde2 100644 --- a/docs/document/content/features/encrypt/limitations.cn.md +++ b/docs/document/content/features/encrypt/limitations.cn.md @@ -13,8 +13,21 @@ weight = 2 ## SQL Server OPENQUERY 加密功能 -`OPENQUERY` 函数的加密改写不支持以下场景: +`OPENQUERY` 的加密改写仅支持如下窄形态透传查询: +```sql +UPDATE OPENQUERY (linked_server, 'SELECT FROM [.] [WHERE ...]') +SET = +``` + +不支持以下场景: + +- `SELECT` 列表中的字符串字面量或表达式; +- 括号标识符中包含空格,例如 `[Human Resources]`; +- 三部分表名,例如 `db.schema.table`; +- 逗号分隔的多表源; +- `JOIN`、`CROSS APPLY`、`OUTER APPLY`; +- `UNION`、`UNION ALL`、`EXCEPT`、`INTERSECT`; - `WHERE` 后引用加密列; -- 物理列名中间包含 `]` 或 `[]`; -- 透传查询中使用 `JOIN`、`CROSS APPLY`、`OUTER APPLY`、`UNION`、`UNION ALL`、`EXCEPT`、`INTERSECT`。 +- 非字面量、非参数的赋值表达式,例如 `SET col = UPPER('x')`; +- 物理列名包含 `]`。 diff --git a/docs/document/content/features/encrypt/limitations.en.md b/docs/document/content/features/encrypt/limitations.en.md index 10ce6f3a308d4..8472f118d3b88 100644 --- a/docs/document/content/features/encrypt/limitations.en.md +++ b/docs/document/content/features/encrypt/limitations.en.md @@ -13,8 +13,21 @@ weight = 2 ## SQL Server OPENQUERY encryption -The following are not supported for encrypt rewrite with the `OPENQUERY` function: +Encrypt rewrite for `OPENQUERY` only supports a narrow pass-through shape: +```sql +UPDATE OPENQUERY (linked_server, 'SELECT FROM [.]
[WHERE ...]') +SET = +``` + +The following are not supported: + +- `SELECT` list items that are string literals or expressions. +- Identifiers that contain spaces inside brackets, such as `[Human Resources]`. +- Three-part table names, such as `db.schema.table`. +- Comma-separated table sources. +- `JOIN`, `CROSS APPLY`, `OUTER APPLY`. +- `UNION`, `UNION ALL`, `EXCEPT`, `INTERSECT`. - Predicates after `WHERE` that reference encrypted columns. -- Physical column names that contain `]` or `[]` in the middle. -- `JOIN`, `CROSS APPLY`, `OUTER APPLY`, `UNION`, `UNION ALL`, `EXCEPT`, and `INTERSECT` in the pass-through query. +- Assignment expressions other than literals or parameter markers, such as `SET col = UPPER('x')`. +- Physical column names that contain `]`. diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java index f17dccfa9be5a..b2b1e1177a1e1 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java @@ -22,6 +22,7 @@ import org.apache.shardingsphere.database.connector.core.type.DatabaseType; import org.apache.shardingsphere.database.connector.core.type.DatabaseTypeRegistry; import org.apache.shardingsphere.encrypt.enums.EncryptDerivedColumnSuffix; +import org.apache.shardingsphere.encrypt.exception.syntax.UnsupportedEncryptSQLException; import org.apache.shardingsphere.encrypt.rewrite.token.pojo.EncryptAssignmentToken; import org.apache.shardingsphere.encrypt.rewrite.token.pojo.EncryptLiteralAssignmentToken; import org.apache.shardingsphere.encrypt.rewrite.token.pojo.EncryptOpenQuerySQLToken; @@ -55,6 +56,8 @@ @RequiredArgsConstructor public final class EncryptAssignmentTokenGenerator { + private static final String UNSUPPORTED_ASSIGNMENT_EXPRESSION = "OPENQUERY with unsupported assignment expression"; + private final EncryptRule rule; private final ShardingSphereDatabase database; @@ -100,27 +103,28 @@ private Collection generateNormalUpdateTokens(final TablesContext tabl } private Collection generateOpenQueryUpdateTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final TableSegment openQueryTable) { + Optional encryptTable = findOpenQueryEncryptTable(openQueryTable); + if (!encryptTable.isPresent()) { + return Collections.emptyList(); + } + EncryptTable table = encryptTable.get(); Collection result = new LinkedList<>(); - Collection openQueryEncryptColumns = new LinkedList<>(); + boolean hasEncryptAssignment = false; for (ColumnAssignmentSegment each : setAssignmentSegment.getAssignments()) { - ColumnSegment assignedColumn = getAssignedColumn(each); - String columnName = assignedColumn.getIdentifier().getValue(); - Optional encryptTable = findOpenQueryEncryptTable(openQueryTable, columnName); - if (!encryptTable.isPresent()) { + String columnName = getAssignedColumn(each).getIdentifier().getValue(); + if (!table.isEncryptColumn(columnName)) { continue; } - EncryptColumn encryptColumn = encryptTable.get().getEncryptColumn(columnName); - appendOpenQueryAssignmentTokens(result, tablesContext, openQueryTable, each, encryptTable.get(), encryptColumn); - openQueryEncryptColumns.add(encryptColumn); + appendOpenQueryAssignmentTokens(result, tablesContext, openQueryTable, each, table, table.getEncryptColumn(columnName)); + hasEncryptAssignment = true; + } + if (hasEncryptAssignment) { + appendComposedOpenQuerySQLToken(result, openQueryTable, table.getEncryptColumns()); } - appendComposedOpenQuerySQLToken(result, openQueryTable, openQueryEncryptColumns); return result; } private void appendComposedOpenQuerySQLToken(final Collection result, final TableSegment openQueryTable, final Collection encryptColumns) { - if (encryptColumns.isEmpty()) { - return; - } Optional openQuerySQL = EncryptOpenQueryUtils.findOpenQuerySQLLiteral(openQueryTable); if (!openQuerySQL.isPresent()) { return; @@ -142,14 +146,15 @@ private void appendOpenQueryAssignmentTokens(final Collection result, String schemaName = EncryptOpenQueryUtils.findSchemaName(openQueryTable) .orElseGet(() -> tablesContext.getSchemaName().orElseGet(() -> databaseTypeRegistry.getDefaultSchemaName(database.getName()))); QuoteCharacter quoteCharacter = databaseTypeRegistry.getDialectDatabaseMetaData().getQuoteCharacter(); - result.addAll(generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptColumn, assignmentSegment, quoteCharacter, true)); + Collection assignmentTokens = generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptColumn, assignmentSegment, quoteCharacter, true); + if (assignmentTokens.isEmpty()) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_ASSIGNMENT_EXPRESSION); + } + result.addAll(assignmentTokens); } - private Optional findOpenQueryEncryptTable(final TableSegment openQueryTable, final String columnName) { - if (!EncryptOpenQueryUtils.findOpenQuerySQLLiteral(openQueryTable).isPresent()) { - return Optional.empty(); - } - return EncryptOpenQueryUtils.findEncryptTable(rule, openQueryTable).filter(optional -> optional.isEncryptColumn(columnName)); + private Optional findOpenQueryEncryptTable(final TableSegment openQueryTable) { + return EncryptOpenQueryUtils.findEncryptTable(rule, openQueryTable); } private EncryptOpenQuerySQLToken generateOpenQuerySQLToken(final LiteralExpressionSegment openQuerySQL, final Collection encryptColumns) { diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java index 5a88a824cbeb4..c05fb7b8d6fe3 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java @@ -53,6 +53,8 @@ final class EncryptOpenQueryPassThroughSQL { private static final String UNSUPPORTED_ENCRYPTED_PREDICATE = "OPENQUERY with predicate on encrypted column"; + private static final String UNSUPPORTED_PHYSICAL_COLUMN_NAME = "OPENQUERY with physical column name containing ]"; + private final String selectList; private final String tableExpression; @@ -585,6 +587,9 @@ private static String getPhysicalColumnNames(final EncryptColumn encryptColumn) } private static String quotePhysicalColumnName(final String physicalColumnName) { + if (physicalColumnName.contains("]")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_PHYSICAL_COLUMN_NAME); + } return QuoteCharacter.BRACKETS.wrap(physicalColumnName); } diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rule/table/EncryptTable.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rule/table/EncryptTable.java index c4416d7aa0b84..3244e21f9797b 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rule/table/EncryptTable.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rule/table/EncryptTable.java @@ -32,6 +32,7 @@ import org.apache.shardingsphere.infra.annotation.HighFrequencyInvocation; import org.apache.shardingsphere.infra.exception.ShardingSpherePreconditions; +import java.util.Collection; import java.util.Map; import java.util.Map.Entry; import java.util.Optional; @@ -73,6 +74,15 @@ private EncryptColumn createEncryptColumn(final EncryptColumnRuleConfiguration c return result; } + /** + * Get encrypt columns. + * + * @return encrypt columns + */ + public Collection getEncryptColumns() { + return columns.values(); + } + /** * Find encryptor. * diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java index d0b1acfef726d..d194b7a4b9bb5 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java @@ -31,6 +31,7 @@ import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.column.ColumnSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.FunctionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment; +import org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.ParameterMarkerExpressionSegment; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.bound.ColumnSegmentBoundInfo; import org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.FunctionTableSegment; import org.apache.shardingsphere.sql.parser.statement.core.value.identifier.IdentifierValue; @@ -40,6 +41,7 @@ import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import java.util.Arrays; import java.util.Collections; import java.util.Optional; @@ -72,6 +74,33 @@ void assertGenerateSQLTokenWithCommaTableSourcesExpectsException() { () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithCommaTableSources())); } + @Test + void assertGenerateSQLTokenWithOpenQueryUnsetEncryptColumnInWhereExpectsException() { + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(mock(ParameterMarkerExpressionSegment.class)); + when(setAssignmentSegment.getAssignments()).thenReturn(Collections.singleton(assignmentSegment)); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + EncryptAssignmentTokenGenerator tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRuleWithUnsetColumn(), mock(ShardingSphereDatabase.class), + TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + assertThrows(UnsupportedEncryptSQLException.class, + () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithExtraColumnInWhere())); + } + + @Test + void assertGenerateSQLTokenWithOpenQueryUnsupportedAssignmentExpressionExpectsException() { + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(new FunctionSegment(124, 134, "UPPER", "UPPER('x')")); + when(setAssignmentSegment.getAssignments()).thenReturn(Collections.singleton(assignmentSegment)); + EncryptAssignmentTokenGenerator tokenGenerator = new EncryptAssignmentTokenGenerator( + mockOpenQueryEncryptRuleWithTable(), mock(ShardingSphereDatabase.class), TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + assertThrows(UnsupportedEncryptSQLException.class, + () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegment())); + } + private EncryptRule mockOpenQueryEncryptRule() { EncryptRule result = mock(EncryptRule.class); EncryptTable encryptTable = mock(EncryptTable.class); @@ -82,6 +111,22 @@ private EncryptRule mockOpenQueryEncryptRule() { return result; } + private EncryptRule mockOpenQueryEncryptRuleWithUnsetColumn() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn groupNameColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + EncryptColumn extraColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); + when(encryptTable.getTable()).thenReturn("Department"); + when(encryptTable.getEncryptColumn("GroupName")).thenReturn(groupNameColumn); + when(encryptTable.getEncryptColumns()).thenReturn(Arrays.asList(groupNameColumn, extraColumn)); + when(groupNameColumn.getName()).thenReturn("GroupName"); + when(groupNameColumn.getCipher().getName()).thenReturn("group_name_cipher"); + when(extraColumn.getName()).thenReturn("ExtraCol"); + return result; + } + private FunctionTableSegment createOpenQueryTableSegmentWithCommaTableSources() { FunctionSegment functionSegment = new FunctionSegment(7, 95, "OPENQUERY", "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department, dbo.Other')"); @@ -89,4 +134,31 @@ private FunctionTableSegment createOpenQueryTableSegmentWithCommaTableSources() functionSegment.getParameters().add(new LiteralExpressionSegment(34, 84, "SELECT GroupName FROM dbo.Department, dbo.Other")); return new FunctionTableSegment(7, 95, functionSegment); } + + private FunctionTableSegment createOpenQueryTableSegmentWithExtraColumnInWhere() { + FunctionSegment functionSegment = new FunctionSegment(7, 112, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE ExtraCol IS NOT NULL')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 101, "SELECT GroupName FROM dbo.Department WHERE ExtraCol IS NOT NULL")); + return new FunctionTableSegment(7, 112, functionSegment); + } + + private EncryptRule mockOpenQueryEncryptRuleWithTable() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); + when(encryptTable.getTable()).thenReturn("Department"); + when(encryptTable.getEncryptColumn("GroupName")).thenReturn(encryptColumn); + return result; + } + + private FunctionTableSegment createOpenQueryTableSegment() { + FunctionSegment functionSegment = new FunctionSegment(7, 106, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE DepartmentID = 4')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 95, "SELECT GroupName FROM dbo.Department WHERE DepartmentID = 4")); + return new FunctionTableSegment(7, 106, functionSegment); + } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java index e7127b70ab485..6a8d9e7ee40ba 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java @@ -241,6 +241,7 @@ private EncryptRule mockOpenQueryEncryptRule() { when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); when(encryptTable.getTable()).thenReturn("Department"); when(encryptTable.getEncryptColumn("GroupName")).thenReturn(encryptColumn); + when(encryptTable.getEncryptColumns()).thenReturn(Collections.singletonList(encryptColumn)); when(encryptColumn.getName()).thenReturn("GroupName"); when(encryptColumn.getCipher().getName()).thenReturn("group_name_cipher"); when(encryptColumn.getCipher().encrypt("foo_db", "dbo", "Department", "GroupName", Collections.singletonList("Sales and Marketing"))) @@ -256,6 +257,7 @@ private EncryptRule mockOpenQuerySpaceDelimitedPhysicalColumnEncryptRule() { when(encryptTable.isEncryptColumn("SecureLabel")).thenReturn(true); when(encryptTable.getTable()).thenReturn("Department"); when(encryptTable.getEncryptColumn("SecureLabel")).thenReturn(encryptColumn); + when(encryptTable.getEncryptColumns()).thenReturn(Collections.singletonList(encryptColumn)); when(encryptColumn.getName()).thenReturn("SecureLabel"); when(encryptColumn.getCipher().getName()).thenReturn("cipher name"); when(encryptColumn.getCipher().encrypt("foo_db", "dbo", "Department", "SecureLabel", Collections.singletonList("secret"))) @@ -297,6 +299,7 @@ private EncryptRule mockOpenQueryMultiColumnEncryptRule() { when(encryptTable.getTable()).thenReturn("Department"); when(encryptTable.getEncryptColumn("GroupName")).thenReturn(groupNameColumn); when(encryptTable.getEncryptColumn("DeptCode")).thenReturn(deptCodeColumn); + when(encryptTable.getEncryptColumns()).thenReturn(Arrays.asList(groupNameColumn, deptCodeColumn)); when(groupNameColumn.getName()).thenReturn("GroupName"); when(groupNameColumn.getCipher().getName()).thenReturn("group_name_cipher"); when(groupNameColumn.getCipher().encrypt("foo_db", "dbo", "Department", "GroupName", Collections.singletonList("Sales"))) @@ -349,6 +352,7 @@ private EncryptRule mockOpenQueryDerivedColumnsEncryptRule() { when(encryptTable.getTable()).thenReturn("Department"); when(encryptTable.getEncryptColumn("GroupName")).thenReturn(groupNameColumn); when(encryptTable.getEncryptColumn("Remark")).thenReturn(remarkColumn); + when(encryptTable.getEncryptColumns()).thenReturn(Arrays.asList(groupNameColumn, remarkColumn)); when(groupNameColumn.getName()).thenReturn("GroupName"); when(groupNameColumn.getCipher().getName()).thenReturn("group_name_cipher"); when(groupNameColumn.getCipher().encrypt("foo_db", "dbo", "Department", "GroupName", Collections.singletonList("Sales"))) @@ -389,4 +393,49 @@ private FunctionTableSegment createDerivedColumnsOpenQueryTableSegment() { functionSegment.getParameters().add(new LiteralExpressionSegment(34, 97, "SELECT GroupName, Remark FROM dbo.Department WHERE DepartmentID = 4")); return new FunctionTableSegment(7, 108, functionSegment); } + + @Test + void assertGenerateSQLTokenWithOpenQuerySelectExtraEncryptColumnRewritesBoth() { + ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); + when(database.getName()).thenReturn("foo_db"); + tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRuleWithTwoColumns(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(new LiteralExpressionSegment(124, 144, "Sales and Marketing")); + when(assignmentSegment.getStopIndex()).thenReturn(144); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + Collection actual = tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithExtraColumn()); + assertThat(actual.size(), is(2)); + Iterator iterator = actual.iterator(); + assertThat(iterator.next().toString(), is("group_name_cipher = 'groupEncryptValue'")); + assertThat(iterator.next().toString(), is("'SELECT [group_name_cipher], [extra_col_cipher] FROM dbo.Department WHERE DepartmentID = 4'")); + } + + private EncryptRule mockOpenQueryEncryptRuleWithTwoColumns() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn groupNameColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + EncryptColumn extraColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("GroupName")).thenReturn(true); + when(encryptTable.getTable()).thenReturn("Department"); + when(encryptTable.getEncryptColumn("GroupName")).thenReturn(groupNameColumn); + when(encryptTable.getEncryptColumns()).thenReturn(Arrays.asList(groupNameColumn, extraColumn)); + when(groupNameColumn.getName()).thenReturn("GroupName"); + when(groupNameColumn.getCipher().getName()).thenReturn("group_name_cipher"); + when(groupNameColumn.getCipher().encrypt("foo_db", "dbo", "Department", "GroupName", Collections.singletonList("Sales and Marketing"))) + .thenReturn(Collections.singletonList("groupEncryptValue")); + when(extraColumn.getName()).thenReturn("ExtraCol"); + when(extraColumn.getCipher().getName()).thenReturn("extra_col_cipher"); + return result; + } + + private FunctionTableSegment createOpenQueryTableSegmentWithExtraColumn() { + FunctionSegment functionSegment = new FunctionSegment(7, 112, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName, ExtraCol FROM dbo.Department WHERE DepartmentID = 4')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 101, "SELECT GroupName, ExtraCol FROM dbo.Department WHERE DepartmentID = 4")); + return new FunctionTableSegment(7, 112, functionSegment); + } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java index ba128b45d0f17..7be2c536238d2 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java @@ -174,6 +174,13 @@ void assertRewriteQuotesPhysicalColumnName(final String scenario, final String p assertThat(actual, is("SELECT " + expectedQuotedColumn + " FROM foo_schema.foo_tbl WHERE id_col = 1")); } + @Test + void assertRewriteRejectsPhysicalColumnNameContainingRightBracket() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE id_col = 1"); + assertThrows(UnsupportedEncryptSQLException.class, + () -> passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo]bar")))); + } + private static Stream quotedPhysicalColumnNameArguments() { return Stream.of( Arguments.of("space", "cipher name", "[cipher name]"), diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml index 7f2b3c8771164..44fe98307f91f 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml @@ -155,6 +155,16 @@ + + + + + + + + + + From 6e5ec36170022fbc1b53010fde73ecf439420141 Mon Sep 17 00:00:00 2001 From: Claire Date: Tue, 21 Jul 2026 22:18:37 +0800 Subject: [PATCH 7/8] fix validateNoCommaSeparatedTableSource --- .../EncryptOpenQueryPassThroughSQL.java | 70 +++++++++++++------ .../EncryptOpenQueryPassThroughSQLTest.java | 11 ++- 2 files changed, 57 insertions(+), 24 deletions(-) diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java index c05fb7b8d6fe3..1bbac5e653229 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java @@ -258,32 +258,56 @@ private static void validateRemainder(final String remainder) { } private static void validateNoCommaSeparatedTableSource(final String remainder) { - int index = skipWhitespace(remainder, 0); - if (index >= remainder.length()) { - return; - } - if (',' == remainder.charAt(index)) { - throw new UnsupportedEncryptSQLException(UNSUPPORTED_COMMA_TABLE_SOURCE); - } - if (isClauseKeywordAt(remainder, index)) { - return; - } - if (isKeywordAt(remainder, index, "AS")) { - index = skipWhitespace(remainder, index + "AS".length()); - Optional aliasPart = readIdentifierPartIfPresent(remainder, index); - if (!aliasPart.isPresent()) { - return; + int index = 0; + int parenDepth = 0; + boolean inString = false; + while (index < remainder.length()) { + char current = remainder.charAt(index); + if (!inString && '/' == current && index + 1 < remainder.length() && '*' == remainder.charAt(index + 1)) { + int closeIndex = remainder.indexOf("*/", index + 2); + index = closeIndex < 0 ? remainder.length() : closeIndex + 2; + continue; + } + if (!inString && '-' == current && index + 1 < remainder.length() && '-' == remainder.charAt(index + 1)) { + int newlineIndex = remainder.indexOf('\n', index + 2); + index = newlineIndex < 0 ? remainder.length() : newlineIndex + 1; + continue; + } + if ('\'' == current) { + if (!inString) { + inString = true; + } else if (index + 1 < remainder.length() && '\'' == remainder.charAt(index + 1)) { + index += 2; + continue; + } else { + inString = false; + } + index++; + continue; } - index = skipWhitespace(remainder, aliasPart.get().getStopIndex()); - } else { - Optional aliasPart = readIdentifierPartIfPresent(remainder, index); - if (!aliasPart.isPresent()) { + if (inString) { + index++; + continue; + } + if ('(' == current) { + parenDepth++; + index++; + continue; + } + if (')' == current) { + if (parenDepth > 0) { + parenDepth--; + } + index++; + continue; + } + if (',' == current && 0 == parenDepth) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_COMMA_TABLE_SOURCE); + } + if (0 == parenDepth && isClauseKeywordAt(remainder, index)) { return; } - index = skipWhitespace(remainder, aliasPart.get().getStopIndex()); - } - if (index < remainder.length() && ',' == remainder.charAt(index)) { - throw new UnsupportedEncryptSQLException(UNSUPPORTED_COMMA_TABLE_SOURCE); + index++; } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java index 7be2c536238d2..b10662719ad7d 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java @@ -123,7 +123,16 @@ private static Stream unsupportedShapeArguments() { Arguments.of("union", "SELECT foo_col FROM foo_schema.foo_tbl UNION SELECT foo_col FROM foo_schema.bar_tbl"), Arguments.of("union all", "SELECT foo_col FROM foo_schema.foo_tbl UNION ALL SELECT foo_col FROM foo_schema.bar_tbl"), Arguments.of("except", "SELECT foo_col FROM foo_schema.foo_tbl EXCEPT SELECT foo_col FROM foo_schema.bar_tbl"), - Arguments.of("intersect", "SELECT foo_col FROM foo_schema.foo_tbl INTERSECT SELECT foo_col FROM foo_schema.bar_tbl")); + Arguments.of("intersect", "SELECT foo_col FROM foo_schema.foo_tbl INTERSECT SELECT foo_col FROM foo_schema.bar_tbl"), + Arguments.of("with hint then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl WITH (NOLOCK), foo_schema.bar_tbl"), + Arguments.of("inline hint then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl(NOLOCK), foo_schema.bar_tbl"), + Arguments.of("block comment then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl/*hint*/, foo_schema.bar_tbl")); + } + + @Test + void assertParseDoesNotRejectSingleTableWithHint() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WITH (NOLOCK) WHERE id_col = 1"); + assertThat(actual.getRemainder(), is(" WITH (NOLOCK) WHERE id_col = 1")); } @Test From f1855a143bc6a142ada2041365f0e7c0d80d1bb4 Mon Sep 17 00:00:00 2001 From: Claire Date: Sat, 25 Jul 2026 17:38:24 +0800 Subject: [PATCH 8/8] update --- .../features/encrypt/limitations.cn.md | 4 +- .../features/encrypt/limitations.en.md | 4 +- .../EncryptAssignmentTokenGenerator.java | 6 +- .../EncryptOpenQueryPassThroughSQL.java | 69 ++++++++++++++++++- .../token/pojo/EncryptOpenQuerySQLToken.java | 2 +- ...eneratorOpenQueryUnsupportedShapeTest.java | 56 ++++++++++++++- .../EncryptAssignmentTokenGeneratorTest.java | 26 +++++++ .../EncryptOpenQueryPassThroughSQLTest.java | 40 ++++++++++- .../query-with-cipher/dml/update/update.xml | 40 +++++++++++ .../encrypt/config/query-with-cipher.yaml | 4 ++ 10 files changed, 237 insertions(+), 14 deletions(-) diff --git a/docs/document/content/features/encrypt/limitations.cn.md b/docs/document/content/features/encrypt/limitations.cn.md index 261287112cde2..1357f7ad3bd3a 100644 --- a/docs/document/content/features/encrypt/limitations.cn.md +++ b/docs/document/content/features/encrypt/limitations.cn.md @@ -22,12 +22,14 @@ SET = 不支持以下场景: -- `SELECT` 列表中的字符串字面量或表达式; +- `SELECT` 列表中的字符串字面量、数字字面量、关键字表达式(例如 `NULL`)或表达式; - 括号标识符中包含空格,例如 `[Human Resources]`; - 三部分表名,例如 `db.schema.table`; - 逗号分隔的多表源; - `JOIN`、`CROSS APPLY`、`OUTER APPLY`; - `UNION`、`UNION ALL`、`EXCEPT`、`INTERSECT`; +- 表引用后的 `ORDER BY`、`GROUP BY`、`HAVING` 等额外子句; +- 使用 `;` 分隔的多条语句; - `WHERE` 后引用加密列; - 非字面量、非参数的赋值表达式,例如 `SET col = UPPER('x')`; - 物理列名包含 `]`。 diff --git a/docs/document/content/features/encrypt/limitations.en.md b/docs/document/content/features/encrypt/limitations.en.md index 8472f118d3b88..f733175ff5059 100644 --- a/docs/document/content/features/encrypt/limitations.en.md +++ b/docs/document/content/features/encrypt/limitations.en.md @@ -22,12 +22,14 @@ SET = The following are not supported: -- `SELECT` list items that are string literals or expressions. +- `SELECT` list items that are string literals, numeric literals, keyword expressions such as `NULL`, or other expressions. - Identifiers that contain spaces inside brackets, such as `[Human Resources]`. - Three-part table names, such as `db.schema.table`. - Comma-separated table sources. - `JOIN`, `CROSS APPLY`, `OUTER APPLY`. - `UNION`, `UNION ALL`, `EXCEPT`, `INTERSECT`. +- Additional trailing clauses after the table reference, such as `ORDER BY`, `GROUP BY`, and `HAVING`. +- Multiple statements separated by `;`. - Predicates after `WHERE` that reference encrypted columns. - Assignment expressions other than literals or parameter markers, such as `SET col = UPPER('x')`. - Physical column names that contain `]`. diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java index b2b1e1177a1e1..d1d79c418c95f 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGenerator.java @@ -109,18 +109,14 @@ private Collection generateOpenQueryUpdateTokens(final TablesContext t } EncryptTable table = encryptTable.get(); Collection result = new LinkedList<>(); - boolean hasEncryptAssignment = false; for (ColumnAssignmentSegment each : setAssignmentSegment.getAssignments()) { String columnName = getAssignedColumn(each).getIdentifier().getValue(); if (!table.isEncryptColumn(columnName)) { continue; } appendOpenQueryAssignmentTokens(result, tablesContext, openQueryTable, each, table, table.getEncryptColumn(columnName)); - hasEncryptAssignment = true; - } - if (hasEncryptAssignment) { - appendComposedOpenQuerySQLToken(result, openQueryTable, table.getEncryptColumns()); } + appendComposedOpenQuerySQLToken(result, openQueryTable, table.getEncryptColumns()); return result; } diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java index 1bbac5e653229..e07da2d655c25 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java @@ -55,6 +55,10 @@ final class EncryptOpenQueryPassThroughSQL { private static final String UNSUPPORTED_PHYSICAL_COLUMN_NAME = "OPENQUERY with physical column name containing ]"; + private static final String UNSUPPORTED_TRAILING_CLAUSE = "OPENQUERY with unsupported trailing clause"; + + private static final String UNSUPPORTED_STATEMENT_TERMINATOR = "OPENQUERY with statement terminator"; + private final String selectList; private final String tableExpression; @@ -81,7 +85,7 @@ private EncryptOpenQueryPassThroughSQL(final String selectList, final String tab * @return table name */ static Optional findTableName(final String passThroughSQL) { - String trimmedSQL = passThroughSQL.trim(); + String trimmedSQL = decodeTSqlStringLiteralEscaping(passThroughSQL.trim()); if (!startsWithKeyword(trimmedSQL, "SELECT")) { return Optional.empty(); } @@ -100,7 +104,7 @@ static Optional findTableName(final String passThroughSQL) { * @throws UnsupportedEncryptSQLException if pass-through SQL shape is unsupported */ static EncryptOpenQueryPassThroughSQL parse(final String passThroughSQL) { - String trimmedSQL = passThroughSQL.trim(); + String trimmedSQL = decodeTSqlStringLiteralEscaping(passThroughSQL.trim()); if (!startsWithKeyword(trimmedSQL, "SELECT")) { throw new UnsupportedEncryptSQLException(UNSUPPORTED_SHAPE); } @@ -237,6 +241,15 @@ private static void validateColumnIdentifier(final String identifier) { throw new UnsupportedEncryptSQLException(UNSUPPORTED_SELECT_EXPRESSION); } } + if (Character.isDigit(identifier.charAt(0)) || isReservedKeywordIdentifier(identifier)) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SELECT_EXPRESSION); + } + } + + private static boolean isReservedKeywordIdentifier(final String identifier) { + return isKeywordAt(identifier, 0, "NULL") + || isKeywordAt(identifier, 0, "TRUE") + || isKeywordAt(identifier, 0, "FALSE"); } private static void validateRemainder(final String remainder) { @@ -244,6 +257,14 @@ private static void validateRemainder(final String remainder) { return; } validateNoCommaSeparatedTableSource(remainder); + if (containsStatementTerminatorOutsideString(remainder)) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_STATEMENT_TERMINATOR); + } + if (containsKeywordOutsideString(remainder, "ORDER") + || containsKeywordOutsideString(remainder, "GROUP") + || containsKeywordOutsideString(remainder, "HAVING")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_TRAILING_CLAUSE); + } if (containsKeywordOutsideString(remainder, "JOIN")) { throw new UnsupportedEncryptSQLException(UNSUPPORTED_JOIN); } @@ -362,6 +383,37 @@ private static boolean containsKeywordOutsideString(final String sqlFragment, fi return false; } + private static boolean containsStatementTerminatorOutsideString(final String sqlFragment) { + int index = 0; + boolean inString = false; + while (index < sqlFragment.length()) { + char current = sqlFragment.charAt(index); + if ('\'' == current) { + if (!inString) { + inString = true; + index++; + continue; + } + if (index + 1 < sqlFragment.length() && '\'' == sqlFragment.charAt(index + 1)) { + index += 2; + continue; + } + inString = false; + index++; + continue; + } + if (inString) { + index++; + continue; + } + if (';' == current) { + return true; + } + index++; + } + return false; + } + private static TableReference parseTableReference(final String passThroughSQL, final int startIndex) { int index = skipWhitespace(passThroughSQL, startIndex); IdentifierPart schemaPart = readIdentifierPart(passThroughSQL, index); @@ -432,6 +484,15 @@ private static Optional findFromKeywordIndexIfPresent(final String pass return Optional.of(index); } char current = passThroughSQL.charAt(index); + if (!inString && ('[' == current || '"' == current)) { + Optional delimitedPart = readDelimitedIdentifierPartIfPresent(passThroughSQL, index); + if (delimitedPart.isPresent()) { + index = delimitedPart.get().getStopIndex(); + continue; + } + index++; + continue; + } if ('\'' != current) { index++; continue; @@ -617,6 +678,10 @@ private static String quotePhysicalColumnName(final String physicalColumnName) { return QuoteCharacter.BRACKETS.wrap(physicalColumnName); } + private static String decodeTSqlStringLiteralEscaping(final String encoded) { + return encoded.replace("''", "'"); + } + @Getter private static final class TableReference { diff --git a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java index 11f2dc659295d..eb5e195a512f3 100644 --- a/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/pojo/EncryptOpenQuerySQLToken.java @@ -39,6 +39,6 @@ public EncryptOpenQuerySQLToken(final int startIndex, final int stopIndex, final @Override public String toString() { - return "'" + sql + "'"; + return "'" + sql.replace("'", "''") + "'"; } } diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java index d194b7a4b9bb5..4c65ff6afff96 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java @@ -96,11 +96,36 @@ void assertGenerateSQLTokenWithOpenQueryUnsupportedAssignmentExpressionExpectsEx when(assignmentSegment.getValue()).thenReturn(new FunctionSegment(124, 134, "UPPER", "UPPER('x')")); when(setAssignmentSegment.getAssignments()).thenReturn(Collections.singleton(assignmentSegment)); EncryptAssignmentTokenGenerator tokenGenerator = new EncryptAssignmentTokenGenerator( - mockOpenQueryEncryptRuleWithTable(), mock(ShardingSphereDatabase.class), TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + mockOpenQueryEncryptRuleForUnsupportedAssignmentExpression(), mock(ShardingSphereDatabase.class), TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); assertThrows(UnsupportedEncryptSQLException.class, () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegment())); } + @Test + void assertGenerateSQLTokenWithOpenQueryOrderByExpectsException() { + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(setAssignmentSegment.getAssignments()).thenReturn(Collections.singleton(assignmentSegment)); + EncryptAssignmentTokenGenerator tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRule(), mock(ShardingSphereDatabase.class), + TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + assertThrows(UnsupportedEncryptSQLException.class, + () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithOrderBy())); + } + + @Test + void assertGenerateSQLTokenWithOpenQueryNonEncryptAssignmentEncryptedPredicateExpectsException() { + ColumnSegment columnSegment = new ColumnSegment(112, 124, new IdentifierValue("DepartmentID")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("DepartmentID"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(new LiteralExpressionSegment(128, 129, 5)); + when(setAssignmentSegment.getAssignments()).thenReturn(Collections.singleton(assignmentSegment)); + EncryptAssignmentTokenGenerator tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRuleForEncryptedPredicateInWhere(), mock(ShardingSphereDatabase.class), + TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + assertThrows(UnsupportedEncryptSQLException.class, + () -> tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithNonEncryptAssignmentEncryptedPredicate())); + } + private EncryptRule mockOpenQueryEncryptRule() { EncryptRule result = mock(EncryptRule.class); EncryptTable encryptTable = mock(EncryptTable.class); @@ -135,6 +160,14 @@ private FunctionTableSegment createOpenQueryTableSegmentWithCommaTableSources() return new FunctionTableSegment(7, 95, functionSegment); } + private FunctionTableSegment createOpenQueryTableSegmentWithOrderBy() { + FunctionSegment functionSegment = new FunctionSegment(7, 108, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department ORDER BY DepartmentID')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 97, "SELECT GroupName FROM dbo.Department ORDER BY DepartmentID")); + return new FunctionTableSegment(7, 108, functionSegment); + } + private FunctionTableSegment createOpenQueryTableSegmentWithExtraColumnInWhere() { FunctionSegment functionSegment = new FunctionSegment(7, 112, "OPENQUERY", "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE ExtraCol IS NOT NULL')"); @@ -143,7 +176,15 @@ private FunctionTableSegment createOpenQueryTableSegmentWithExtraColumnInWhere() return new FunctionTableSegment(7, 112, functionSegment); } - private EncryptRule mockOpenQueryEncryptRuleWithTable() { + private FunctionTableSegment createOpenQueryTableSegmentWithNonEncryptAssignmentEncryptedPredicate() { + FunctionSegment functionSegment = new FunctionSegment(7, 127, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT GroupName, DepartmentID FROM dbo.Department WHERE GroupName IS NOT NULL')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 116, "SELECT GroupName, DepartmentID FROM dbo.Department WHERE GroupName IS NOT NULL")); + return new FunctionTableSegment(7, 127, functionSegment); + } + + private EncryptRule mockOpenQueryEncryptRuleForUnsupportedAssignmentExpression() { EncryptRule result = mock(EncryptRule.class); EncryptTable encryptTable = mock(EncryptTable.class); EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); @@ -154,6 +195,17 @@ private EncryptRule mockOpenQueryEncryptRuleWithTable() { return result; } + private EncryptRule mockOpenQueryEncryptRuleForEncryptedPredicateInWhere() { + EncryptRule result = mock(EncryptRule.class); + EncryptTable encryptTable = mock(EncryptTable.class); + EncryptColumn encryptColumn = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); + when(result.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + when(encryptTable.isEncryptColumn("DepartmentID")).thenReturn(false); + when(encryptTable.getEncryptColumns()).thenReturn(Collections.singletonList(encryptColumn)); + when(encryptColumn.getName()).thenReturn("GroupName"); + return result; + } + private FunctionTableSegment createOpenQueryTableSegment() { FunctionSegment functionSegment = new FunctionSegment(7, 106, "OPENQUERY", "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE DepartmentID = 4')"); diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java index 6a8d9e7ee40ba..930be87b7bc37 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorTest.java @@ -168,6 +168,24 @@ void assertGenerateSQLTokenWithOpenQueryLiteralExpressionSegment() { assertThat(iterator.next().toString(), is("'SELECT [group_name_cipher] FROM dbo.Department WHERE DepartmentID = 4'")); } + @Test + void assertGenerateSQLTokenWithOpenQueryBracketedFromInSelectList() { + ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); + when(database.getName()).thenReturn("foo_db"); + tokenGenerator = new EncryptAssignmentTokenGenerator(mockOpenQueryEncryptRule(), database, TypedSPILoader.getService(DatabaseType.class, "FIXTURE")); + ColumnSegment columnSegment = new ColumnSegment(112, 120, new IdentifierValue("GroupName")); + columnSegment.setColumnBoundInfo(new ColumnSegmentBoundInfo(null, null, new IdentifierValue("GroupName"), TableSourceType.TEMPORARY_TABLE)); + when(assignmentSegment.getColumns()).thenReturn(Collections.singletonList(columnSegment)); + when(assignmentSegment.getValue()).thenReturn(new LiteralExpressionSegment(124, 144, "Sales and Marketing")); + when(assignmentSegment.getStopIndex()).thenReturn(144); + when(tablesContext.getSchemaName()).thenReturn(Optional.of("dbo")); + Collection actual = tokenGenerator.generateSQLTokens(tablesContext, setAssignmentSegment, createOpenQueryTableSegmentWithBracketedFromInSelectList()); + Iterator iterator = actual.iterator(); + assertThat(actual.size(), is(2)); + assertThat(iterator.next().toString(), is("group_name_cipher = 'encryptValue'")); + assertThat(iterator.next().toString(), is("'SELECT [FROM], [group_name_cipher] FROM dbo.Department WHERE DepartmentID = 4'")); + } + @Test void assertGenerateSQLTokenWithOpenQueryMultipleAssignments() { ShardingSphereDatabase database = mock(ShardingSphereDatabase.class); @@ -272,6 +290,14 @@ private FunctionTableSegment createOpenQueryTableSegment() { return new FunctionTableSegment(7, 106, functionSegment); } + private FunctionTableSegment createOpenQueryTableSegmentWithBracketedFromInSelectList() { + FunctionSegment functionSegment = new FunctionSegment(7, 114, "OPENQUERY", + "OPENQUERY (MyLinkedServer, 'SELECT [FROM], GroupName FROM dbo.Department WHERE DepartmentID = 4')"); + functionSegment.getParameters().add(new ColumnSegment(18, 31, new IdentifierValue("MyLinkedServer"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(34, 103, "SELECT [FROM], GroupName FROM dbo.Department WHERE DepartmentID = 4")); + return new FunctionTableSegment(7, 114, functionSegment); + } + private FunctionTableSegment createOpenQueryTableSegmentWithColumnInWhere() { FunctionSegment functionSegment = new FunctionSegment(7, 115, "OPENQUERY", "OPENQUERY (MyLinkedServer, 'SELECT GroupName FROM dbo.Department WHERE GroupName IS NOT NULL')"); diff --git a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java index b10662719ad7d..864fd88d98b62 100644 --- a/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java @@ -100,7 +100,10 @@ private static Stream findTableNameArguments() { Arguments.of("double quoted multipart", "SELECT foo_col FROM \"foo_schema\".\"foo_tbl\" WHERE id_col = 1", "foo_tbl"), Arguments.of("escaped bracket", "SELECT foo_col FROM [foo_schema].[foo]]tbl]", "foo]tbl"), Arguments.of("three-part table", "SELECT foo_col FROM foo_db.foo_schema.foo_tbl WHERE id_col = 4", "foo_tbl"), - Arguments.of("join", "SELECT foo_col FROM foo_schema.foo_tbl JOIN foo_schema.bar_tbl ON foo_tbl.id_col = bar_tbl.id_col", "foo_tbl")); + Arguments.of("join", "SELECT foo_col FROM foo_schema.foo_tbl JOIN foo_schema.bar_tbl ON foo_tbl.id_col = bar_tbl.id_col", "foo_tbl"), + Arguments.of("bracketed from in select list", "SELECT [FROM], foo_col FROM foo_schema.foo_tbl WHERE id_col = 4", "foo_tbl"), + Arguments.of("double quoted from in select list", "SELECT \"FROM\", foo_col FROM foo_schema.foo_tbl WHERE id_col = 4", "foo_tbl"), + Arguments.of("escaped bracketed from in select list", "SELECT [FR]]OM], foo_col FROM foo_schema.foo_tbl", "foo_tbl")); } @ParameterizedTest(name = "{0}") @@ -126,7 +129,13 @@ private static Stream unsupportedShapeArguments() { Arguments.of("intersect", "SELECT foo_col FROM foo_schema.foo_tbl INTERSECT SELECT foo_col FROM foo_schema.bar_tbl"), Arguments.of("with hint then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl WITH (NOLOCK), foo_schema.bar_tbl"), Arguments.of("inline hint then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl(NOLOCK), foo_schema.bar_tbl"), - Arguments.of("block comment then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl/*hint*/, foo_schema.bar_tbl")); + Arguments.of("block comment then comma table sources", "SELECT foo_col FROM foo_schema.foo_tbl/*hint*/, foo_schema.bar_tbl"), + Arguments.of("select list numeric literal", "SELECT 1 FROM foo_schema.foo_tbl"), + Arguments.of("select list null keyword", "SELECT NULL FROM foo_schema.foo_tbl"), + Arguments.of("statement terminator", "SELECT foo_col FROM foo_schema.foo_tbl; DELETE FROM foo_schema.bar_tbl"), + Arguments.of("order by", "SELECT foo_col FROM foo_schema.foo_tbl ORDER BY id_col"), + Arguments.of("group by", "SELECT foo_col FROM foo_schema.foo_tbl GROUP BY id_col"), + Arguments.of("having", "SELECT foo_col FROM foo_schema.foo_tbl HAVING COUNT(1) > 0")); } @Test @@ -200,6 +209,20 @@ private static Stream quotedPhysicalColumnNameArguments() { Arguments.of("left bracket", "foo[bar", "[foo[bar]")); } + @Test + void assertRewriteWithBracketedFromInSelectList() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT [FROM], foo_col FROM foo_schema.foo_tbl WHERE id_col = 4"); + String actual = passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo_col_cipher"))); + assertThat(actual, is("SELECT [FROM], [foo_col_cipher] FROM foo_schema.foo_tbl WHERE id_col = 4")); + } + + @Test + void assertRewriteWithDoubleQuotedFromInSelectList() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT \"FROM\", foo_col FROM foo_schema.foo_tbl WHERE id_col = 4"); + String actual = passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo_col_cipher"))); + assertThat(actual, is("SELECT \"FROM\", [foo_col_cipher] FROM foo_schema.foo_tbl WHERE id_col = 4")); + } + private EncryptColumn createEncryptColumn(final String logicColumnName, final String physicalColumnName) { EncryptColumn result = mock(EncryptColumn.class, RETURNS_DEEP_STUBS); when(result.getName()).thenReturn(logicColumnName); @@ -208,4 +231,17 @@ private EncryptColumn createEncryptColumn(final String logicColumnName, final St when(result.getLikeQuery()).thenReturn(Optional.empty()); return result; } + + @Test + void assertParseDoesNotRejectSetKeywordInsideDoubledApostropheStringLiteral() { + EncryptOpenQueryPassThroughSQL actual = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE note_col = ''UNION ALL''"); + assertThat(actual.getRemainder(), is(" WHERE note_col = 'UNION ALL'")); + } + + @Test + void assertRewriteDoesNotRejectEncryptColumnNameInsideDoubledApostropheStringLiteral() { + EncryptOpenQueryPassThroughSQL passThroughSQL = EncryptOpenQueryPassThroughSQL.parse("SELECT foo_col FROM foo_schema.foo_tbl WHERE note_col = ''foo_col''"); + String actual = passThroughSQL.rewrite(Collections.singletonList(createEncryptColumn("foo_col", "foo_col_cipher"))); + assertThat(actual, is("SELECT [foo_col_cipher] FROM foo_schema.foo_tbl WHERE note_col = 'foo_col'")); + } } diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml index 44fe98307f91f..58a60095431cd 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/case/query-with-cipher/dml/update/update.xml @@ -215,6 +215,46 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml index 0bc50290cc399..a836b1e9a0893 100644 --- a/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml +++ b/test/it/rewriter/src/test/resources/scenario/encrypt/config/query-with-cipher.yaml @@ -202,6 +202,10 @@ rules: cipher: name: foo[bar encryptorName: rewrite_normal_fixture + ApostropheLabel: + cipher: + name: foo'bar + encryptorName: rewrite_normal_fixture StateRegion: columns: CountryRegionName: