diff --git a/RELEASE-NOTES.md b/RELEASE-NOTES.md index a5af0bc22b30b..3612eebcbffbf 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 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..1357f7ad3bd3a 100644 --- a/docs/document/content/features/encrypt/limitations.cn.md +++ b/docs/document/content/features/encrypt/limitations.cn.md @@ -10,3 +10,26 @@ weight = 2 - 加密字段无法支持计算操作,如:AVG、SUM 以及计算表达式; - 不支持使用 `;` 分隔的多条 SQL 同时执行; - 当投影子查询中包含加密字段时,必须使用别名。 + +## SQL Server OPENQUERY 加密功能 + +`OPENQUERY` 的加密改写仅支持如下窄形态透传查询: + +```sql +UPDATE OPENQUERY (linked_server, 'SELECT FROM [.] [WHERE ...]') +SET = +``` + +不支持以下场景: + +- `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 ce6d3a831564e..f733175ff5059 100644 --- a/docs/document/content/features/encrypt/limitations.en.md +++ b/docs/document/content/features/encrypt/limitations.en.md @@ -10,3 +10,26 @@ 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 + +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, 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/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/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 25e949454a423..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 @@ -17,14 +17,15 @@ package org.apache.shardingsphere.encrypt.rewrite.token.generator.assignment; -import lombok.AllArgsConstructor; -import lombok.extern.slf4j.Slf4j; +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.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; 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,6 +41,7 @@ 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; @@ -50,11 +52,12 @@ /** * Assignment generator for encrypt. */ -@Slf4j @HighFrequencyInvocation -@AllArgsConstructor +@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; @@ -69,53 +72,125 @@ public final class EncryptAssignmentTokenGenerator { * @return generated SQL tokens */ public Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment) { + return generateNormalUpdateTokens(tablesContext, setAssignmentSegment); + } + + /** + * Generate SQL tokens. + * + * @param tablesContext SQL statement context + * @param setAssignmentSegment set assignment segment + * @param openQueryTable OPENQUERY function table segment + * @return generated SQL tokens + */ + Collection generateSQLTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment, final TableSegment openQueryTable) { + return generateOpenQueryUpdateTokens(tablesContext, setAssignmentSegment, openQueryTable); + } + + private Collection generateNormalUpdateTokens(final TablesContext tablesContext, final SetAssignmentSegment setAssignmentSegment) { 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)); - } - }); + 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 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<>(); + 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)); + } + appendComposedOpenQuerySQLToken(result, openQueryTable, table.getEncryptColumns()); + return result; + } + + private void appendComposedOpenQuerySQLToken(final Collection result, final TableSegment openQueryTable, final Collection encryptColumns) { + 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 = tablesContext.getSchemaName().orElseGet(() -> databaseTypeRegistry.getDefaultSchemaName(database.getName())); + QuoteCharacter quoteCharacter = databaseTypeRegistry.getDialectDatabaseMetaData().getQuoteCharacter(); + result.addAll(generateAssignmentSQLTokens(schemaName, encryptTable.getTable(), encryptColumn, assignmentSegment, quoteCharacter, false)); + } + + 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(); + 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) { + return EncryptOpenQueryUtils.findEncryptTable(rule, openQueryTable); + } + + 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, - 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 +219,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) { 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..e07da2d655c25 --- /dev/null +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQL.java @@ -0,0 +1,719 @@ +/* + * 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.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.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 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 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; + + 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 = decodeTSqlStringLiteralEscaping(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 = decodeTSqlStringLiteralEscaping(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 + * @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(); + 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 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(); + 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; + } + 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) { + 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) { + if (remainder.isEmpty()) { + 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); + } + 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 = 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; + } + 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++; + } + } + + 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 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); + 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); + } + Optional delimitedPart = readDelimitedIdentifierPartIfPresent(passThroughSQL, index); + if (delimitedPart.isPresent()) { + if ('[' == passThroughSQL.charAt(index) && delimitedPart.get().getValue().contains(" ")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_SPACE_DELIMITED_IDENTIFIER); + } + return delimitedPart.get(); + } + 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 (!inString && ('[' == current || '"' == current)) { + Optional delimitedPart = readDelimitedIdentifierPartIfPresent(passThroughSQL, index); + if (delimitedPart.isPresent()) { + index = delimitedPart.get().getStopIndex(); + continue; + } + index++; + continue; + } + 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(); + } + Optional delimitedPart = readDelimitedIdentifierPartIfPresent(passThroughSQL, index); + if (delimitedPart.isPresent()) { + return delimitedPart; + } + 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 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()) { + 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 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)) { + return Optional.of(each); + } + } + return Optional.empty(); + } + + private static String getPhysicalColumnNames(final EncryptColumn encryptColumn) { + 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) { + if (physicalColumnName.contains("]")) { + throw new UnsupportedEncryptSQLException(UNSUPPORTED_PHYSICAL_COLUMN_NAME); + } + return QuoteCharacter.BRACKETS.wrap(physicalColumnName); + } + + private static String decodeTSqlStringLiteralEscaping(final String encoded) { + return encoded.replace("''", "'"); + } + + @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 new file mode 100644 index 0000000000000..057659fdbad17 --- /dev/null +++ b/features/encrypt/core/src/main/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtils.java @@ -0,0 +1,98 @@ +/* + * 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.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; + +/** + * Encrypt OPENQUERY utility. + */ +@NoArgsConstructor(access = AccessLevel.PRIVATE) +public final class EncryptOpenQueryUtils { + + /** + * 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(); + } + Optional tableName = EncryptOpenQueryPassThroughSQL.findTableName(openQuerySQL.get().getText()); + return tableName.isPresent() ? rule.findEncryptTable(tableName.get()) : 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(); + } + 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 e070296811892..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; @@ -44,7 +45,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) { @@ -58,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()); + 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/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..eb5e195a512f3 --- /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.replace("'", "''") + "'"; + } +} 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/context/EncryptSQLRewriteContextDecoratorTest.java b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/context/EncryptSQLRewriteContextDecoratorTest.java index b51721d653d51..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 @@ -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,62 @@ 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()); + } + + @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)); 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); + } + + 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/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..4c65ff6afff96 --- /dev/null +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptAssignmentTokenGeneratorOpenQueryUnsupportedShapeTest.java @@ -0,0 +1,216 @@ +/* + * 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.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; +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.Arrays; +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())); + } + + @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( + 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); + 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 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')"); + 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); + } + + 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')"); + 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 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); + 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 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')"); + 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 4ded88a893a42..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 @@ -21,20 +21,27 @@ 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; +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; 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,15 +53,21 @@ 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; import java.util.Optional; 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.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) @@ -136,4 +149,319 @@ 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'")); + } + + @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); + 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 assertGenerateSQLTokenWithOpenQueryEncryptedColumnInWhereExpectsException() { + 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")); + 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(); + assertThat(iterator.next().toString(), is("cipher name = 'encryptValue'")); + assertThat(iterator.next().toString(), is("'SELECT [cipher name] FROM dbo.Department WHERE DepartmentID = 4'")); + } + + @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.findEncryptTable("Department")).thenReturn(Optional.of(encryptTable)); + 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"))) + .thenReturn(Collections.singletonList("encryptValue")); + 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(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"))) + .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); + } + + 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')"); + 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 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); + 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(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"))) + .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(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"))) + .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); + } + + @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 new file mode 100644 index 0000000000000..864fd88d98b62 --- /dev/null +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryPassThroughSQLTest.java @@ -0,0 +1,247 @@ +/* + * 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 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; +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 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 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 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 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 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'")); + } + + @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(expectedTableName)); + } + + 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"), + 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}") + @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"), + 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("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 + 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 + 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 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 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 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 = createEncryptColumn("Remark", "remark_cipher"); + AssistedQueryColumnItem assistedQuery = mock(AssistedQueryColumnItem.class); + LikeQueryColumnItem likeQuery = mock(LikeQueryColumnItem.class); + 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 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 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")); + } + + @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]"), + 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]")); + } + + @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); + when(result.getCipher().getName()).thenReturn(physicalColumnName); + when(result.getAssistedQuery()).thenReturn(Optional.empty()); + 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/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..827ba33c90afb --- /dev/null +++ b/features/encrypt/core/src/test/java/org/apache/shardingsphere/encrypt/rewrite/token/generator/assignment/EncryptOpenQueryUtilsTest.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 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.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.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()); + } + + @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 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); + 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()); + } + + @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"))); + functionSegment.getParameters().add(new LiteralExpressionSegment(0, 0, openQuerySQL)); + return new FunctionTableSegment(0, 0, 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..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 @@ -83,6 +83,11 @@ 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), + 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())); @@ -107,6 +112,13 @@ 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), + 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; } @@ -121,11 +133,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..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 @@ -115,6 +115,146 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + 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..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 @@ -158,6 +158,54 @@ rules: cipher: name: remark_cipher encryptorName: rewrite_normal_fixture + Department: + columns: + GroupName: + 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 + 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 + ApostropheLabel: + cipher: + name: foo'bar + encryptorName: rewrite_normal_fixture StateRegion: columns: CountryRegionName: