feature: support sql tpye for at mode (#4832)
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/AbstractConnectionProxy.java b/rm-datasource/src/main/java/io/seata/rm/datasource/AbstractConnectionProxy.java index ef8c810..df3185c 100644 --- a/rm-datasource/src/main/java/io/seata/rm/datasource/AbstractConnectionProxy.java +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/AbstractConnectionProxy.java
@@ -22,6 +22,8 @@ import io.seata.rm.datasource.sql.struct.TableMetaCacheFactory; import io.seata.sqlparser.SQLRecognizer; import io.seata.sqlparser.SQLType; +import io.seata.sqlparser.util.JdbcConstants; + import java.sql.Array; import java.sql.Blob; import java.sql.CallableStatement; @@ -111,7 +113,8 @@ List<SQLRecognizer> sqlRecognizers = SQLVisitorFactory.get(sql, dbType); if (sqlRecognizers != null && sqlRecognizers.size() == 1) { SQLRecognizer sqlRecognizer = sqlRecognizers.get(0); - if (sqlRecognizer != null && sqlRecognizer.getSQLType() == SQLType.INSERT) { + if (sqlRecognizer != null && (sqlRecognizer.getSQLType() == SQLType.INSERT || sqlRecognizer.getSQLType() == SQLType.INSERT_IGNORE + || sqlRecognizer.getSQLType() == SQLType.INSERT_SELECT && JdbcConstants.MYSQL.equals(getDbType()))) { TableMeta tableMeta = TableMetaCacheFactory.getTableMetaCache(dbType).getTableMeta(getTargetConnection(), sqlRecognizer.getTableName(), getDataSourceProxy().getResourceId()); String[] pkNameArray = new String[tableMeta.getPrimaryKeyOnlyName().size()];
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/BaseInsertExecutor.java b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/BaseInsertExecutor.java index 63dbcbf..78fd8b5 100644 --- a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/BaseInsertExecutor.java +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/BaseInsertExecutor.java
@@ -20,6 +20,7 @@ import java.sql.SQLException; import java.sql.Statement; import java.util.ArrayList; +import java.util.Collection; import java.util.HashMap; import java.util.List; import java.util.Map; @@ -145,8 +146,7 @@ boolean ps = true; if (statementProxy instanceof PreparedStatementProxy) { PreparedStatementProxy preparedStatementProxy = (PreparedStatementProxy) statementProxy; - - List<List<Object>> insertRows = recognizer.getInsertRows(pkIndexMap.values()); + List<List<Object>> insertRows = getInsertRows(pkIndexMap.values()); if (insertRows != null && !insertRows.isEmpty()) { Map<Integer, ArrayList<Object>> parameters = preparedStatementProxy.getParameters(); final int rowSize = insertRows.size(); @@ -197,7 +197,7 @@ } } else { ps = false; - List<List<Object>> insertRows = recognizer.getInsertRows(pkIndexMap.values()); + List<List<Object>> insertRows = getInsertRows(pkIndexMap.values()); for (List<Object> row : insertRows) { pkIndexMap.forEach((pkKey, pkIndex) -> { List<Object> pkValues = pkValuesMap.get(pkKey); @@ -220,6 +220,17 @@ } /** + * user for insert select + * + * @param primaryKeyIndex the primary key index + * @return the insert rows + */ + protected List<List<Object>> getInsertRows(Collection<Integer> primaryKeyIndex) { + SQLInsertRecognizer recognizer = (SQLInsertRecognizer) sqlRecognizer; + return recognizer.getInsertRows(primaryKeyIndex); + } + + /** * default get generated keys. * @return the generate keys * @throws SQLException the sql exception
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/ExecuteTemplate.java b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/ExecuteTemplate.java index d317cd2..5683694 100644 --- a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/ExecuteTemplate.java +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/ExecuteTemplate.java
@@ -25,7 +25,9 @@ import io.seata.core.context.RootContext; import io.seata.core.model.BranchType; import io.seata.rm.datasource.StatementProxy; +import io.seata.rm.datasource.exec.mysql.MySQLInsertIgnoreExecutor; import io.seata.rm.datasource.exec.mysql.MySQLInsertOnDuplicateUpdateExecutor; +import io.seata.rm.datasource.exec.mysql.MySQLInsertSelectExecutor; import io.seata.rm.datasource.exec.mysql.MySQLUpdateJoinExecutor; import io.seata.rm.datasource.sql.SQLVisitorFactory; import io.seata.sqlparser.SQLRecognizer; @@ -124,6 +126,26 @@ throw new NotSupportYetException(dbType + " not support to " + SQLType.UPDATE_JOIN.name()); } break; + case INSERT_IGNORE: + switch (dbType) { + case JdbcConstants.MYSQL: + case JdbcConstants.ORACLE: + executor = new MySQLInsertIgnoreExecutor(statementProxy, statementCallback, sqlRecognizer); + break; + default: + throw new NotSupportYetException(dbType + " not support to insert ignore"); + } + break; + case INSERT_SELECT: + switch (dbType) { + case JdbcConstants.MYSQL: + case JdbcConstants.ORACLE: + executor = new MySQLInsertSelectExecutor(statementProxy, statementCallback, sqlRecognizer); + break; + default: + throw new NotSupportYetException(dbType + " not support to insert select"); + } + break; default: executor = new PlainExecutor<>(statementProxy, statementCallback); break;
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertIgnoreExecutor.java b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertIgnoreExecutor.java new file mode 100644 index 0000000..54bc54b --- /dev/null +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertIgnoreExecutor.java
@@ -0,0 +1,52 @@ +/* + * Copyright 1999-2019 Seata.io Group. + * + * Licensed 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 io.seata.rm.datasource.exec.mysql; + +import io.seata.common.util.CollectionUtils; +import io.seata.rm.datasource.ConnectionProxy; +import io.seata.rm.datasource.StatementProxy; +import io.seata.rm.datasource.exec.StatementCallback; +import io.seata.rm.datasource.sql.struct.Row; +import io.seata.rm.datasource.sql.struct.TableRecords; +import io.seata.sqlparser.SQLRecognizer; +import io.seata.sqlparser.SQLType; +import io.seata.sqlparser.struct.Defaultable; + +import java.util.Map; +import java.util.List; + +/** + * @author: lyx + */ +public class MySQLInsertIgnoreExecutor extends MySQLInsertOnDuplicateUpdateExecutor implements Defaultable { + + public MySQLInsertIgnoreExecutor(StatementProxy statementProxy, StatementCallback statementCallback, SQLRecognizer sqlRecognizer) { + super(statementProxy, statementCallback, sqlRecognizer); + } + + @Override + protected void buildUndoItemAll(ConnectionProxy connectionProxy, TableRecords beforeImage, TableRecords afterImage) { + Map<SQLType, List<Row>> updateAndInsertRow = getUpdateAndInsertRow(beforeImage, afterImage); + List<Row> insertRows = updateAndInsertRow.get(SQLType.INSERT); + if (CollectionUtils.isNotEmpty(insertRows)) { + TableRecords partAfterImage = new TableRecords(afterImage.getTableMeta()); + partAfterImage.setTableName(afterImage.getTableName()); + partAfterImage.setRows(insertRows); + connectionProxy.appendUndoLog(buildUndoItem(SQLType.INSERT, TableRecords.empty(getTableMeta()), partAfterImage)); + } + } + +}
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertOnDuplicateUpdateExecutor.java b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertOnDuplicateUpdateExecutor.java index b1fd711..811c0b3 100644 --- a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertOnDuplicateUpdateExecutor.java +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertOnDuplicateUpdateExecutor.java
@@ -23,7 +23,10 @@ import java.util.ArrayList; import java.util.Map; import java.util.Collections; +import java.util.Objects; +import java.util.Optional; import java.util.StringJoiner; +import java.util.concurrent.atomic.AtomicReference; import java.util.stream.Collectors; import com.google.common.base.Joiner; @@ -41,6 +44,7 @@ import io.seata.rm.datasource.exec.StatementCallback; import io.seata.rm.datasource.sql.struct.ColumnMeta; import io.seata.rm.datasource.sql.struct.Field; +import io.seata.rm.datasource.sql.struct.IndexType; import io.seata.rm.datasource.sql.struct.Row; import io.seata.rm.datasource.sql.struct.TableMeta; import io.seata.rm.datasource.sql.struct.TableRecords; @@ -50,6 +54,7 @@ import io.seata.sqlparser.SQLType; import io.seata.sqlparser.struct.Defaultable; import io.seata.sqlparser.struct.Null; +import io.seata.sqlparser.util.ColumnUtils; import io.seata.sqlparser.util.JdbcConstants; /** @@ -58,31 +63,43 @@ @LoadLevel(name = JdbcConstants.MYSQL, scope = Scope.PROTOTYPE) public class MySQLInsertOnDuplicateUpdateExecutor extends MySQLInsertExecutor implements Defaultable { - private static final String COLUMN_SEPARATOR = "|"; - /** - * is updated or not - */ - private boolean isUpdateFlag = false; - public String getSelectSQL() { return selectSQL; } /** + * just for test + * + * @param selectSQL select sql + */ + public void setSelectSQL(String selectSQL) { + this.selectSQL = selectSQL; + } + + /** * before image sql and after image sql,condition is unique index */ - private String selectSQL; + protected String selectSQL; - public ArrayList<List<Object>> getParamAppenderList() { - return paramAppenderList; + public HashMap<List<String>, List<Object>> getParamAppenderMap() { + return paramAppenderMap; } /** * the params of selectSQL, value is the unique index */ - private ArrayList<List<Object>> paramAppenderList; + public HashMap<List<String>, List<Object>> paramAppenderMap; + + /** + * for test + * + * @param paramAppenderMap paramAppenderMap + */ + public void setParamAppenderMap(HashMap<List<String>, List<Object>> paramAppenderMap) { + this.paramAppenderMap = paramAppenderMap; + } public MySQLInsertOnDuplicateUpdateExecutor(StatementProxy statementProxy, StatementCallback statementCallback, SQLRecognizer sqlRecognizer) { super(statementProxy, statementCallback, sqlRecognizer); @@ -101,9 +118,7 @@ throw new NotSupportYetException("multi pk only support mysql!"); } TableRecords beforeImage = beforeImage(); - if (CollectionUtils.isNotEmpty(beforeImage.getRows())) { - isUpdateFlag = true; - } else { + if (CollectionUtils.isEmpty(beforeImage.getRows())) { beforeImage = TableRecords.empty(getTableMeta()); } Object result = statementCallback.execute(statementProxy.getTargetStatement(), args); @@ -138,11 +153,36 @@ * @param afterImage the after image */ protected void buildUndoItemAll(ConnectionProxy connectionProxy, TableRecords beforeImage, TableRecords afterImage) { - if (!isUpdateFlag) { - SQLUndoLog sqlUndoLog = buildUndoItem(SQLType.INSERT, TableRecords.empty(getTableMeta()), afterImage); - connectionProxy.appendUndoLog(sqlUndoLog); - return; + Map<SQLType, List<Row>> updateAndInsertRow = getUpdateAndInsertRow(beforeImage, afterImage); + List<Row> insertRows = updateAndInsertRow.get(SQLType.INSERT); + List<Row> updateRows = updateAndInsertRow.get(SQLType.UPDATE); + if (CollectionUtils.isNotEmpty(updateRows)) { + TableRecords partAfterImage = new TableRecords(afterImage.getTableMeta()); + partAfterImage.setTableName(afterImage.getTableName()); + partAfterImage.setRows(updateRows); + if (beforeImage.getRows().size() != partAfterImage.getRows().size()) { + throw new ShouldNeverHappenException("Before image size is not equaled to after image size, probably because you updated the primary keys."); + } + connectionProxy.appendUndoLog(buildUndoItem(SQLType.UPDATE, beforeImage, partAfterImage)); } + if (CollectionUtils.isNotEmpty(insertRows)) { + TableRecords partAfterImage = new TableRecords(afterImage.getTableMeta()); + partAfterImage.setTableName(afterImage.getTableName()); + partAfterImage.setRows(insertRows); + connectionProxy.appendUndoLog(buildUndoItem(SQLType.INSERT, TableRecords.empty(getTableMeta()), partAfterImage)); + } + } + + /** + * if beforeImage and afterImage both have,then sql type is update + * if afterImage have but beforeImage have not, then sql type is insert + * + * @param beforeImage before image + * @param afterImage after image + * @return map + */ + protected Map<SQLType, List<Row>> getUpdateAndInsertRow(TableRecords beforeImage, TableRecords afterImage) { + Map<SQLType, List<Row>> result = new HashMap<>(2, 1.001f); List<Row> beforeImageRows = beforeImage.getRows(); List<String> beforePrimaryValues = new ArrayList<>(); for (Row r : beforeImageRows) { @@ -166,23 +206,12 @@ insertRows.add(r); } } - if (CollectionUtils.isNotEmpty(updateRows)) { - TableRecords partAfterImage = new TableRecords(afterImage.getTableMeta()); - partAfterImage.setTableName(afterImage.getTableName()); - partAfterImage.setRows(updateRows); - if (beforeImage.getRows().size() != partAfterImage.getRows().size()) { - throw new ShouldNeverHappenException("Before image size is not equaled to after image size, probably because you updated the primary keys."); - } - connectionProxy.appendUndoLog(buildUndoItem(SQLType.UPDATE, beforeImage, partAfterImage)); - } - if (CollectionUtils.isNotEmpty(insertRows)) { - TableRecords partAfterImage = new TableRecords(afterImage.getTableMeta()); - partAfterImage.setTableName(afterImage.getTableName()); - partAfterImage.setRows(insertRows); - connectionProxy.appendUndoLog(buildUndoItem(SQLType.INSERT, TableRecords.empty(getTableMeta()), partAfterImage)); - } + result.put(SQLType.INSERT, insertRows); + result.put(SQLType.UPDATE, updateRows); + return result; } + /** * build a SQLUndoLog * @@ -203,17 +232,18 @@ @Override - protected TableRecords afterImage(TableRecords beforeImage) throws SQLException { + public TableRecords afterImage(TableRecords beforeImage) throws SQLException { TableMeta tableMeta = getTableMeta(); List<Row> rows = beforeImage.getRows(); - Map<String, ArrayList<Object>> primaryValueMap = new HashMap<>(); + Map<List<String>, ArrayList<Object>> primaryValueMap = new HashMap<>(); + AtomicReference<List<String>> nameList = new AtomicReference<>(); rows.forEach(m -> { List<Field> fields = m.primaryKeys(); - fields.forEach(f -> { - ArrayList<Object> values = primaryValueMap.computeIfAbsent(f.getName(), v -> new ArrayList<>()); - values.add(f.getValue()); - }); + nameList.set(fields.stream().map(Field::getName).collect(Collectors.toList())); + ArrayList<Object> tempList = new ArrayList<>(); + fields.forEach(f -> tempList.add(f.getValue())); + primaryValueMap.computeIfAbsent(nameList.get(), v -> new ArrayList<>()).addAll(tempList); }); // The origin select sql contains the unique keys sql @@ -221,16 +251,25 @@ List<Object> primaryValues = new ArrayList<>(); // Appends the pk when the origin select sql not contains - for (int i = 0; i < rows.size(); i++) { - List<String> wherePrimaryList = new ArrayList<>(); - primaryValueMap.forEach((k, v) -> { - wherePrimaryList.add(k + " = ? "); - primaryValues.add(v); + if (CollectionUtils.isNotEmpty(primaryValueMap)) { + primaryValueMap.forEach((columnsName, columnsValue) -> { + afterImageSql.append("OR ("); + afterImageSql.append(Joiner.on(",").join(columnsName)); + afterImageSql.append(") in("); + for (int i = 0; i < columnsValue.size() / columnsName.size(); i++) { + afterImageSql.append("("); + for (int j = 0; j < columnsName.size(); j++) { + afterImageSql.append("?,"); + } + afterImageSql.insert(afterImageSql.length() - 1, ")"); + } + afterImageSql.deleteCharAt(afterImageSql.length() - 1); + afterImageSql.append(")"); }); - afterImageSql.append(" OR (").append(Joiner.on(" and ").join(wherePrimaryList)).append(") "); } - - return buildTableRecords2(tableMeta, afterImageSql.toString(), paramAppenderList, primaryValues); + ArrayList<Object> pkList = new ArrayList<>(); + primaryValueMap.values().forEach(pkList::addAll); + return buildTableRecords2(tableMeta, afterImageSql.toString(), new ArrayList<>(paramAppenderMap.values()), pkList); } @Override @@ -238,10 +277,13 @@ TableMeta tableMeta = getTableMeta(); // After image sql the same of before image if (StringUtils.isBlank(selectSQL)) { - paramAppenderList = new ArrayList<>(); selectSQL = buildImageSQL(tableMeta); } - return buildTableRecords2(tableMeta, selectSQL, paramAppenderList, Collections.emptyList()); + if (CollectionUtils.isEmpty(paramAppenderMap)) { + throw new ShouldNeverHappenException("can not find unique param,may be you should add unique key" + + " when use the sqlType of " + sqlRecognizer.getSQLType().getName()); + } + return buildTableRecords2(tableMeta, selectSQL, new ArrayList<>(paramAppenderMap.values()), Collections.emptyList()); } /** @@ -284,55 +326,77 @@ * @return image sql */ public String buildImageSQL(TableMeta tableMeta) { - if (CollectionUtils.isEmpty(paramAppenderList)) { - paramAppenderList = new ArrayList<>(); - } SQLInsertRecognizer recognizer = (SQLInsertRecognizer) sqlRecognizer; - int insertNum = recognizer.getInsertParamsValue().size(); + int insertNum = getInsertParamsValue().size(); Map<String, ArrayList<Object>> imageParameterMap = buildImageParameters(recognizer); + if (Objects.isNull(paramAppenderMap)) { + paramAppenderMap = new HashMap<>(); + } + List<Object> nullList = new ArrayList<>(); + List<String> nullColumn = new ArrayList<>(); String prefix = "SELECT * "; StringBuilder suffix = new StringBuilder(" FROM ").append(getFromTableInSQL()); - boolean[] isContainWhere = {false}; for (int i = 0; i < insertNum; i++) { int finalI = i; - List<Object> paramAppenderTempList = new ArrayList<>(); tableMeta.getAllIndexes().forEach((k, v) -> { if (!v.isNonUnique()) { - boolean columnIsNull = true; - List<String> uniqueList = new ArrayList<>(); + List<String> columnList = new ArrayList<>(v.getValues().size()); + List<Object> columnValue = new ArrayList<>(v.getValues().size()); for (ColumnMeta m : v.getValues()) { String columnName = m.getColumnName(); + if (JdbcConstants.ORACLE.equals(getDbType()) && recognizer.isIgnore() + && !columnName.equals(ColumnUtils.delEscape(recognizer.getHintColumnName(), getDbType()))) { + break; + } List<Object> imageParameters = imageParameterMap.get(columnName); if (imageParameters == null && m.getColumnDef() != null) { - uniqueList.add(columnName + " = DEFAULT(" + columnName + ") "); - columnIsNull = false; + columnList.add(columnName); + columnValue.add("DEFAULT(" + columnName + ")"); continue; } if ((imageParameters == null && m.getColumnDef() == null) || imageParameters.get(finalI) == null || imageParameters.get(finalI) instanceof Null) { - if (!"PRIMARY".equalsIgnoreCase(k)) { - columnIsNull = false; - uniqueList.add(columnName + " is ? "); - paramAppenderTempList.add("NULL"); + if (!"PRIMARY".equalsIgnoreCase(k) && !IndexType.PRIMARY.equals(v.getIndextype())) { + nullColumn.add("(" + columnName + " is null or " + columnName + " = ?)"); + nullList.add("NULL"); continue; } + // break for the situation of composite primary key break; } - columnIsNull = false; - uniqueList.add(columnName + " = ? "); - paramAppenderTempList.add(imageParameters.get(finalI)); + columnList.add(columnName); + columnValue.add(imageParameterMap.get(columnName).get(finalI)); } - if (!columnIsNull) { - if (isContainWhere[0]) { - suffix.append(" OR (").append(Joiner.on(" and ").join(uniqueList)).append(") "); - } else { - suffix.append(" WHERE (").append(Joiner.on(" and ").join(uniqueList)).append(") "); - isContainWhere[0] = true; - } + if (CollectionUtils.isNotEmpty(columnList)) { + CollectionUtils.computeIfAbsent(paramAppenderMap, columnList, e -> new ArrayList<>()) + .addAll(columnValue); } } }); - paramAppenderList.add(paramAppenderTempList); } + suffix.append(" WHERE "); + if (CollectionUtils.isNotEmpty(nullColumn)) { + paramAppenderMap.put(nullColumn, nullList); + } + paramAppenderMap.forEach((columnsName, columnsValue) -> { + if (columnsName.equals(nullColumn)) { + suffix.append(Joiner.on(" OR ").join(nullColumn)); + suffix.append(" OR "); + return; + } + suffix.append("("); + suffix.append(Joiner.on(",").join(columnsName)); + suffix.append(") in("); + for (int i = 0; i < columnsValue.size() / columnsName.size(); i++) { + suffix.append("("); + for (int j = 0; j < columnsName.size(); j++) { + suffix.append("?,"); + } + suffix.insert(suffix.length() - 1, ")"); + } + suffix.deleteCharAt(suffix.length() - 1); + suffix.append(") OR "); + }); + suffix.delete(suffix.length() - 4, suffix.length() - 1); StringJoiner selectSQLJoin = new StringJoiner(", ", prefix, suffix.toString()); return selectSQLJoin.toString(); } @@ -350,7 +414,7 @@ List<String> duplicateKeyUpdateLowerCaseColumns = duplicateKeyUpdateColumns.parallelStream().map(String::toLowerCase).collect(Collectors.toList()); getTableMeta().getAllIndexes().forEach((k, v) -> { - if ("PRIMARY".equalsIgnoreCase(k)) { + if ("PRIMARY".equalsIgnoreCase(k) && !IndexType.PRIMARY.equals(v.getIndextype())) { for (ColumnMeta m : v.getValues()) { if (duplicateKeyUpdateLowerCaseColumns.contains(m.getColumnName().toLowerCase())) { throw new ShouldNeverHappenException("update pk value is not supported!"); @@ -362,14 +426,17 @@ Map<String, ArrayList<Object>> imageParameterMap = new LowerCaseLinkHashMap<>(); Map<Integer, ArrayList<Object>> parameters = ((PreparedStatementProxy) statementProxy).getParameters(); // VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) - List<String> insertParamsList = recognizer.getInsertParamsValue(); - List<String> sqlRecognizerColumns = recognizer.getInsertColumns(); - List<String> insertColumns = CollectionUtils.isEmpty(sqlRecognizerColumns) ? new ArrayList<>(getTableMeta().getAllColumns().keySet()) : sqlRecognizerColumns; + List<String> insertParamsList = getInsertParamsValue(); + List<String> insertColumns = Optional.ofNullable(recognizer.getInsertColumns()).map(list -> list.stream() + .map(column -> ColumnUtils.delEscape(column, getDbType())).collect(Collectors.toList())).orElse(null); + if (CollectionUtils.isEmpty(insertColumns)) { + insertColumns = new ArrayList<>(getTableMeta().getAllColumns().keySet()); + } int paramsindex = 1; for (String insertParams : insertParamsList) { String[] insertParamsArray = insertParams.split(","); for (int i = 0; i < insertColumns.size(); i++) { - String m = insertColumns.get(i); + String m = ColumnUtils.delEscape(insertColumns.get(i), getDbType()); String params = insertParamsArray[i]; ArrayList<Object> imageListTemp = imageParameterMap.computeIfAbsent(m, k -> new ArrayList<>()); if ("?".equals(params.trim())) { @@ -390,4 +457,14 @@ return imageParameterMap; } + /** + * just for the different recognize or sql + * normal to see {@link SQLInsertRecognizer#getInsertParamsValue} + * + * @return + */ + protected List<String> getInsertParamsValue() { + SQLInsertRecognizer recognizer = (SQLInsertRecognizer) sqlRecognizer; + return recognizer.getInsertParamsValue(); + } }
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertSelectExecutor.java b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertSelectExecutor.java new file mode 100644 index 0000000..6b570eb --- /dev/null +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/exec/mysql/MySQLInsertSelectExecutor.java
@@ -0,0 +1,242 @@ +/* + * Copyright 1999-2019 Seata.io Group. + * + * Licensed 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 io.seata.rm.datasource.exec.mysql; + +import io.seata.common.exception.NotSupportYetException; +import io.seata.common.exception.ShouldNeverHappenException; +import io.seata.common.util.CollectionUtils; +import io.seata.common.util.LowerCaseLinkHashMap; +import io.seata.rm.datasource.ConnectionProxy; +import io.seata.rm.datasource.PreparedStatementProxy; +import io.seata.rm.datasource.StatementProxy; +import io.seata.rm.datasource.exec.StatementCallback; +import io.seata.rm.datasource.sql.SQLVisitorFactory; +import io.seata.rm.datasource.sql.struct.Field; +import io.seata.rm.datasource.sql.struct.Row; +import io.seata.rm.datasource.sql.struct.TableMeta; +import io.seata.rm.datasource.sql.struct.TableMetaCacheFactory; +import io.seata.rm.datasource.sql.struct.TableRecords; +import io.seata.sqlparser.SQLInsertRecognizer; +import io.seata.sqlparser.SQLRecognizer; +import io.seata.sqlparser.SQLType; +import io.seata.sqlparser.struct.Defaultable; +import io.seata.sqlparser.util.ColumnUtils; +import io.seata.sqlparser.util.JdbcConstants; + +import java.sql.SQLException; +import java.util.ArrayList; +import java.util.Collection; +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.stream.Collectors; + +/** + * @author: lyx + */ +public class MySQLInsertSelectExecutor extends MySQLInsertOnDuplicateUpdateExecutor implements Defaultable { + + private static final String SELECT = "SELECT"; + + private static final String FOR_UPDATE = " FOR UPDATE"; + + /** + * insert recognizer from sql + */ + private SQLInsertRecognizer insertRecognizer; + + public MySQLInsertSelectExecutor(StatementProxy statementProxy, StatementCallback statementCallback, SQLRecognizer sqlRecognizer) throws SQLException { + super(statementProxy, statementCallback, sqlRecognizer); + createInsertRecognizer(); + } + + public void createInsertRecognizer() throws SQLException { + SQLInsertRecognizer recognizer = (SQLInsertRecognizer) sqlRecognizer; + // get the sql after insert + String querySQL = recognizer.getQuerySQL(); + insertRecognizer = doCreateInsertRecognizer(querySQL); + } + + @Override + public TableRecords beforeImage() throws SQLException { + TableMeta tableMeta = getTableMeta(); + SQLInsertRecognizer recognizer = (SQLInsertRecognizer) sqlRecognizer; + // when insertRecognizer is null, the select sql have not value + if (Objects.isNull(insertRecognizer)) { + return TableRecords.empty(tableMeta); + } + // oracle insert select can not get pk from drive,so should get uk when build image sql + if (recognizer.isIgnore() || CollectionUtils.isNotEmpty(recognizer.getDuplicateKeyUpdate()) + || JdbcConstants.ORACLE.equals(getDbType())) { + if (io.seata.common.util.StringUtils.isBlank(selectSQL)) { + selectSQL = buildImageSQL(tableMeta); + } + if (CollectionUtils.isEmpty(paramAppenderMap)) { + throw new NotSupportYetException("can not find unique param,may be you should add unique key when use the sqlType of" + + " on duplicate key update or insert select"); + } + return buildTableRecords2(tableMeta, selectSQL, new ArrayList<>(paramAppenderMap.values()), Collections.emptyList()); + } + return TableRecords.empty(tableMeta); + } + + @Override + public TableRecords afterImage(TableRecords beforeImage) throws SQLException { + TableMeta tableMeta = getTableMeta(); + SQLInsertRecognizer recognizer = (SQLInsertRecognizer) sqlRecognizer; + if (Objects.isNull(insertRecognizer)) { + return TableRecords.empty(tableMeta); + } + if (recognizer.isIgnore() || CollectionUtils.isNotEmpty(recognizer.getDuplicateKeyUpdate()) + || JdbcConstants.ORACLE.equals(getDbType())) { + return super.afterImage(beforeImage); + } + Map<String, List<Object>> pkValues = getPkValues(); + TableRecords afterImage = buildTableRecords(pkValues); + if (afterImage == null) { + throw new SQLException("Failed to build after-image for insert"); + } + return afterImage; + } + + @Override + protected void buildUndoItemAll(ConnectionProxy connectionProxy, TableRecords beforeImage, TableRecords afterImage) { + SQLInsertRecognizer recognizer = (SQLInsertRecognizer) sqlRecognizer; + if (CollectionUtils.isNotEmpty(recognizer.getDuplicateKeyUpdate())) { + super.buildUndoItemAll(connectionProxy, beforeImage, afterImage); + } else { + Map<SQLType, List<Row>> updateAndInsertRow = getUpdateAndInsertRow(beforeImage, afterImage); + List<Row> insertRows = updateAndInsertRow.get(SQLType.INSERT); + if (CollectionUtils.isNotEmpty(insertRows)) { + TableRecords partAfterImage = new TableRecords(afterImage.getTableMeta()); + partAfterImage.setTableName(afterImage.getTableName()); + partAfterImage.setRows(insertRows); + connectionProxy.appendUndoLog(buildUndoItem(SQLType.INSERT, TableRecords.empty(getTableMeta()), partAfterImage)); + } + } + } + + /** + * from the sql type to get insert rows + * unless select insert,other in {@link SQLInsertRecognizer#getInsertRows(Collection)} + * + * @param primaryKeyIndex the primary key index + * @return the insert rows + */ + @Override + public List<List<Object>> getInsertRows(Collection primaryKeyIndex) { + return Objects.nonNull(insertRecognizer) ? insertRecognizer.getInsertRows(primaryKeyIndex) : Collections.emptyList(); + } + + /** + * from the sql type to get insert params values + * unless select insert,other in {@link SQLInsertRecognizer#getInsertParamsValue} + * + * @return the insert params values + */ + @Override + public List<String> getInsertParamsValue() { + return Objects.nonNull(insertRecognizer) ? insertRecognizer.getInsertParamsValue() : Collections.emptyList(); + } + + @Override + public Map<String, ArrayList<Object>> buildImageParameters(SQLInsertRecognizer recognizer) { + List<String> insertParamsList = getInsertParamsValue(); + List<String> insertColumns = Optional.ofNullable(recognizer.getInsertColumns()).map(list -> list.stream() + .map(column -> ColumnUtils.delEscape(column, getDbType())).collect(Collectors.toList())).orElse(null); + if (CollectionUtils.isEmpty(insertColumns)) { + insertColumns = new ArrayList<>(getTableMeta().getAllColumns().keySet()); + } + Map<String, ArrayList<Object>> imageParameterMap = new LowerCaseLinkHashMap<>(insertColumns.size(), 1); + + for (String insertParams : insertParamsList) { + String[] insertParamsArray = insertParams.split(","); + for (int i = 0; i < insertColumns.size(); i++) { + String m = ColumnUtils.delEscape(insertColumns.get(i), getDbType()); + String params = insertParamsArray[i]; + ArrayList<Object> imageListTemp = imageParameterMap.computeIfAbsent(m, k -> new ArrayList<>()); + imageListTemp.add(params.trim()); + imageParameterMap.put(m, imageListTemp); + } + } + return imageParameterMap; + } + + /** + * create the real insert recognizer + * + * @param querySQL the sql after insert + * @throws SQLException + */ + protected SQLInsertRecognizer doCreateInsertRecognizer(String querySQL) throws SQLException { + Map<Integer, ArrayList<Object>> parameters = ((PreparedStatementProxy) statementProxy).getParameters(); + List<SQLRecognizer> sqlRecognizers = SQLVisitorFactory.get(querySQL + FOR_UPDATE, getDbType()); + SQLRecognizer selectRecognizer = sqlRecognizers.get(0); + ConnectionProxy connectionProxy = statementProxy.getConnectionProxy(); + TableMeta selectTableMeta = TableMetaCacheFactory.getTableMetaCache(connectionProxy.getDbType()) + .getTableMeta(connectionProxy.getTargetConnection(), selectRecognizer.getTableName(), connectionProxy.getDataSourceProxy().getResourceId()); + // use query SQL to get values from database + TableRecords tableRecords = buildTableRecords2(selectTableMeta, querySQL, new ArrayList<>(parameters.values()), Collections.emptyList()); + if (CollectionUtils.isNotEmpty(tableRecords.getRows())) { + StringBuilder valuesSQL = new StringBuilder(); + // build values sql + valuesSQL.append(" VALUES"); + tableRecords.getRows().forEach(row -> { + List<Object> values = row.getFields().stream().map(Field::getValue) + .map(value -> Objects.isNull(value) ? null : value).collect(Collectors.toList()); + valuesSQL.append("("); + for (Object value : values) { + if (Objects.isNull(value)) { + valuesSQL.append((String) null); + } else { + valuesSQL.append(value); + } + valuesSQL.append(","); + } + valuesSQL.insert(valuesSQL.length() - 1, ")"); + }); + valuesSQL.deleteCharAt(valuesSQL.length() - 1); + List<SQLRecognizer> insertSQLRecognizers = SQLVisitorFactory.get(formatOriginSQL(valuesSQL.toString()), getDbType()); + if (CollectionUtils.isEmpty(insertSQLRecognizers)) { + throw new NotSupportYetException("can not support the sql type together with select"); + } + return (SQLInsertRecognizer) insertSQLRecognizers.get(0); + } + return null; + } + + /** + * format origin sql + * + * @param valueSQL the value after insert sql + * @return eg: insert into test values(1,1) + */ + private String formatOriginSQL(String valueSQL) { + String tableName = this.sqlRecognizer.getTableName().toUpperCase(); + String originalSQL = this.sqlRecognizer.getOriginalSQL().toUpperCase(); + int index = originalSQL.indexOf(SELECT); + if (tableName.equalsIgnoreCase(SELECT)) { + // choose the next select + index = originalSQL.indexOf(SELECT, index + SELECT.length()); + } + if (index == -1) { + throw new ShouldNeverHappenException("may be the query sql is not a select SQL"); + } + return this.sqlRecognizer.getOriginalSQL().substring(0, index) + valueSQL; + } +}
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Field.java b/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Field.java index 4c6c2bc..cfcaaf6 100755 --- a/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Field.java +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Field.java
@@ -15,6 +15,7 @@ */ package io.seata.rm.datasource.sql.struct; + /** * Field *
diff --git a/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Row.java b/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Row.java index 28a59e6..67b130d 100755 --- a/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Row.java +++ b/rm-datasource/src/main/java/io/seata/rm/datasource/sql/struct/Row.java
@@ -18,7 +18,6 @@ import java.util.ArrayList; import java.util.List; - /** * The type Row. *
diff --git a/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertIgnoreExecutorTest.java b/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertIgnoreExecutorTest.java new file mode 100644 index 0000000..a477bb8 --- /dev/null +++ b/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertIgnoreExecutorTest.java
@@ -0,0 +1,288 @@ +/* + * Copyright 1999-2019 Seata.io Group. + * + * Licensed 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 io.seata.rm.datasource.exec; + +import java.sql.SQLException; +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import com.google.common.collect.Lists; +import io.seata.rm.datasource.ConnectionProxy; +import io.seata.rm.datasource.PreparedStatementProxy; +import io.seata.rm.datasource.StatementProxy; +import io.seata.rm.datasource.exec.mysql.MySQLInsertIgnoreExecutor; +import io.seata.rm.datasource.sql.struct.ColumnMeta; +import io.seata.rm.datasource.sql.struct.IndexMeta; +import io.seata.rm.datasource.sql.struct.IndexType; +import io.seata.rm.datasource.sql.struct.Row; +import io.seata.rm.datasource.sql.struct.TableMeta; +import io.seata.rm.datasource.sql.struct.TableRecords; +import io.seata.sqlparser.SQLInsertRecognizer; +import io.seata.sqlparser.util.JdbcConstants; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mockito; + +import static org.mockito.Mockito.doReturn; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +/** + * @author: lyx + */ +public class MySQLInsertIgnoreExecutorTest { + + private static final String ID_COLUMN = "id"; + private static final String USER_ID_COLUMN = "user_id"; + private static final String USER_NAME_COLUMN = "user_name"; + private static final String USER_STATUS_COLUMN = "user_status"; + private static final Integer PK_VALUE = 100; + + private StatementProxy statementProxy; + + private SQLInsertRecognizer sqlInsertRecognizer; + + private TableMeta tableMeta; + + private MySQLInsertIgnoreExecutor insertIgnoreExecutor; + + private final int pkIndex = 0; + private HashMap<String, Integer> pkIndexMap; + + @BeforeEach + public void init() throws SQLException { + ConnectionProxy connectionProxy = mock(ConnectionProxy.class); + when(connectionProxy.getDbType()).thenReturn(JdbcConstants.MYSQL); + + statementProxy = mock(PreparedStatementProxy.class); + when(statementProxy.getConnectionProxy()).thenReturn(connectionProxy); + when(statementProxy.getConnection()).thenReturn(connectionProxy); + StatementCallback statementCallback = mock(StatementCallback.class); + sqlInsertRecognizer = mock(SQLInsertRecognizer.class); + tableMeta = mock(TableMeta.class); + insertIgnoreExecutor = Mockito.spy(new MySQLInsertIgnoreExecutor(statementProxy, statementCallback, sqlInsertRecognizer)); + + pkIndexMap = new HashMap<String, Integer>() { + { + put(ID_COLUMN, pkIndex); + } + }; + } + + @Test + public void TestBuildImageParamperters() { + mockParameters(); + List<String> insertParamsList = new ArrayList<>(); + insertParamsList.add("?,?,?,?"); + insertParamsList.add("?,?,?,?"); + when(sqlInsertRecognizer.getInsertParamsValue()).thenReturn(insertParamsList); + mockInsertColumns(); + Map<String, ArrayList<Object>> imageParamperterMap = insertIgnoreExecutor.buildImageParameters(sqlInsertRecognizer); + Assertions.assertEquals(imageParamperterMap.toString(), mockImageParamperterMap().toString()); + } + + @Test + public void TestBuildImageParamperters_contain_constant() { + mockImageParamperterMap_contain_constant(); + List<String> insertParamsList = new ArrayList<>(); + insertParamsList.add("?,?,?,userStatus1"); + insertParamsList.add("?,?,?,userStatus2"); + when(sqlInsertRecognizer.getInsertParamsValue()).thenReturn(insertParamsList); + mockInsertColumns(); + Map<String, ArrayList<Object>> imageParamperterMap = insertIgnoreExecutor.buildImageParameters(sqlInsertRecognizer); + Assertions.assertEquals(imageParamperterMap.toString(), mockImageParamperterMap().toString()); + } + + @Test + public void testBuildImageSQL() { + String selectSQLStr = "SELECT * FROM null WHERE (user_id) in((?),(?)) OR (id) in((?),(?)) "; + String paramAppenderListStr = "{[user_id]=[userId1, userId2], [id]=[100, 101]}"; + mockImageParamperterMap_contain_constant(); + List<String> insertParamsList = new ArrayList<>(); + insertParamsList.add("?,?,?,userStatus1"); + insertParamsList.add("?,?,?,userStatus2"); + when(sqlInsertRecognizer.getInsertParamsValue()).thenReturn(insertParamsList); + mockInsertColumns(); + mockAllIndexes(); + String selectSQL = insertIgnoreExecutor.buildImageSQL(tableMeta); + Assertions.assertEquals(selectSQLStr, selectSQL); + Assertions.assertEquals(paramAppenderListStr, insertIgnoreExecutor.getParamAppenderMap().toString()); + } + + @Test + public void testBeforeImages() throws SQLException { + mockImageParamperterMap_contain_constant(); + List<String> insertParamsList = new ArrayList<>(); + insertParamsList.add("?,?,?,userStatus1"); + insertParamsList.add("?,?,?,userStatus2"); + when(sqlInsertRecognizer.getInsertParamsValue()).thenReturn(insertParamsList); + mockInsertColumns(); + mockAllIndexes(); + String selectSQL = insertIgnoreExecutor.buildImageSQL(tableMeta); + insertIgnoreExecutor.setSelectSQL(selectSQL); + HashMap<List<String>, List<Object>> paramAppenderMap = insertIgnoreExecutor.getParamAppenderMap(); + doReturn(tableMeta).when(insertIgnoreExecutor).getTableMeta(); + TableRecords tableRecords = new TableRecords(); + doReturn(tableRecords).when(insertIgnoreExecutor).buildTableRecords2(tableMeta, selectSQL, new ArrayList<>(paramAppenderMap.values()), Collections.emptyList()); + TableRecords beforeImage = insertIgnoreExecutor.beforeImage(); + Assertions.assertEquals(beforeImage, tableRecords); + } + + @Test + public void testAfterImage() throws SQLException { + String selectSQL = "SELECT * FROM null WHERE (id) in((?),(?)) "; + TableRecords afterImage = new TableRecords(); + doReturn(tableMeta).when(insertIgnoreExecutor).getTableMeta(); + TableRecords beforeImage = new TableRecords(); + HashMap<List<String>, List<Object>> hashMap = new HashMap<>(); + insertIgnoreExecutor.setParamAppenderMap(hashMap); + hashMap.put(Collections.singletonList("id"), Collections.singletonList("id1")); + insertIgnoreExecutor.setSelectSQL(selectSQL); + doReturn(afterImage).when(insertIgnoreExecutor).buildTableRecords2(tableMeta, selectSQL, new ArrayList<>(hashMap.values()), Collections.emptyList()); + TableRecords resultRecords = insertIgnoreExecutor.afterImage(beforeImage); + Assertions.assertEquals(resultRecords, afterImage); + } + + private void getMockTableRecords(TableRecords tableRecords) { + ArrayList<Row> rows = new ArrayList<>(); + Row beforeRow01 = new Row(); + rows.add(beforeRow01); + Row beforeRow02 = new Row(); + rows.add(beforeRow02); + tableRecords.setRows(rows); + } + + + private void mockAllIndexes() { + Map<String, IndexMeta> allIndex = new LinkedHashMap<>(); + List<String> primaryKeyOnlyNameList = new ArrayList<>(); + primaryKeyOnlyNameList.add("id"); + + IndexMeta primary = new IndexMeta(); + primary.setIndextype(IndexType.PRIMARY); + ColumnMeta columnMeta = new ColumnMeta(); + columnMeta.setColumnName("id"); + primary.setValues(Lists.newArrayList(columnMeta)); + allIndex.put("id", primary); + + IndexMeta unique = new IndexMeta(); + unique.setIndextype(IndexType.PRIMARY); + ColumnMeta columnMetaUnique = new ColumnMeta(); + columnMetaUnique.setColumnName("user_id"); + unique.setValues(Lists.newArrayList(columnMetaUnique)); + allIndex.put("user_id", unique); + when(tableMeta.getAllIndexes()).thenReturn(allIndex); + when(tableMeta.getPrimaryKeyOnlyName()).thenReturn(primaryKeyOnlyNameList); + } + + private List<String> mockInsertColumns() { + List<String> columns = new ArrayList<>(); + columns.add(ID_COLUMN); + columns.add(USER_ID_COLUMN); + columns.add(USER_NAME_COLUMN); + columns.add(USER_STATUS_COLUMN); + when(sqlInsertRecognizer.getInsertColumns()).thenReturn(columns); + return columns; + } + + /** + * all insert params is variable + * {1=[100], 2=[userId1], 3=[userName1], 4=[userStatus1], 5=[101], 6=[userId2], 7=[userName2], 8=[userStatus2]} + */ + private void mockParameters() { + Map<Integer, ArrayList<Object>> paramters = new HashMap<>(4); + ArrayList arrayList10 = new ArrayList<>(); + arrayList10.add(PK_VALUE); + ArrayList arrayList11 = new ArrayList<>(); + arrayList11.add("userId1"); + ArrayList arrayList12 = new ArrayList<>(); + arrayList12.add("userName1"); + ArrayList arrayList13 = new ArrayList<>(); + arrayList13.add("userStatus1"); + paramters.put(1, arrayList10); + paramters.put(2, arrayList11); + paramters.put(3, arrayList12); + paramters.put(4, arrayList13); + ArrayList arrayList20 = new ArrayList<>(); + arrayList20.add(PK_VALUE + 1); + ArrayList arrayList21 = new ArrayList<>(); + arrayList21.add("userId2"); + ArrayList arrayList22 = new ArrayList<>(); + arrayList22.add("userName2"); + ArrayList arrayList23 = new ArrayList<>(); + arrayList23.add("userStatus2"); + paramters.put(5, arrayList20); + paramters.put(6, arrayList21); + paramters.put(7, arrayList22); + paramters.put(8, arrayList23); + PreparedStatementProxy psp = (PreparedStatementProxy) this.statementProxy; + when(psp.getParameters()).thenReturn(paramters); + } + + /** + * exist insert parms is constant + * {1=[100], 2=[userId1], 3=[userName1], 4=[101], 5=[userId2], 6=[userName2]} + */ + private void mockImageParamperterMap_contain_constant() { + Map<Integer, ArrayList<Object>> paramters = new HashMap<>(4); + ArrayList arrayList10 = new ArrayList<>(); + arrayList10.add(PK_VALUE); + ArrayList arrayList11 = new ArrayList<>(); + arrayList11.add("userId1"); + ArrayList arrayList12 = new ArrayList<>(); + arrayList12.add("userName1"); + paramters.put(1, arrayList10); + paramters.put(2, arrayList11); + paramters.put(3, arrayList12); + ArrayList arrayList20 = new ArrayList<>(); + arrayList20.add(PK_VALUE + 1); + ArrayList arrayList21 = new ArrayList<>(); + arrayList21.add("userId2"); + ArrayList arrayList22 = new ArrayList<>(); + arrayList22.add("userName2"); + paramters.put(4, arrayList20); + paramters.put(5, arrayList21); + paramters.put(6, arrayList22); + PreparedStatementProxy psp = (PreparedStatementProxy) this.statementProxy; + when(psp.getParameters()).thenReturn(paramters); + } + + private Map<String, ArrayList<Object>> mockImageParamperterMap() { + Map<String, ArrayList<Object>> imageParamperterMap = new LinkedHashMap<>(); + ArrayList<Object> idList = new ArrayList<>(); + idList.add("100"); + idList.add("101"); + imageParamperterMap.put("id", idList); + ArrayList<Object> user_idList = new ArrayList<>(); + user_idList.add("userId1"); + user_idList.add("userId2"); + imageParamperterMap.put("user_id", user_idList); + ArrayList<Object> user_nameList = new ArrayList<>(); + user_nameList.add("userName1"); + user_nameList.add("userName2"); + imageParamperterMap.put("user_name", user_nameList); + ArrayList<Object> user_statusList = new ArrayList<>(); + user_statusList.add("userStatus1"); + user_statusList.add("userStatus2"); + imageParamperterMap.put("user_status", user_statusList); + return imageParamperterMap; + } +} \ No newline at end of file
diff --git a/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertOnDuplicateUpdateExecutorTest.java b/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertOnDuplicateUpdateExecutorTest.java index e54bb79..8a15119 100644 --- a/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertOnDuplicateUpdateExecutorTest.java +++ b/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertOnDuplicateUpdateExecutorTest.java
@@ -18,11 +18,13 @@ import java.sql.SQLException; import java.util.ArrayList; import java.util.Arrays; +import java.util.Collections; import java.util.HashMap; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; -import java.util.Collections; +import java.util.concurrent.atomic.AtomicReference; +import java.util.stream.Collectors; import com.google.common.collect.Lists; import io.seata.rm.datasource.ConnectionProxy; @@ -30,8 +32,11 @@ import io.seata.rm.datasource.StatementProxy; import io.seata.rm.datasource.exec.mysql.MySQLInsertOnDuplicateUpdateExecutor; import io.seata.rm.datasource.sql.struct.ColumnMeta; +import io.seata.rm.datasource.sql.struct.Field; import io.seata.rm.datasource.sql.struct.IndexMeta; import io.seata.rm.datasource.sql.struct.IndexType; +import io.seata.rm.datasource.sql.struct.KeyType; +import io.seata.rm.datasource.sql.struct.Row; import io.seata.rm.datasource.sql.struct.TableMeta; import io.seata.rm.datasource.sql.struct.TableRecords; import io.seata.sqlparser.SQLInsertRecognizer; @@ -66,7 +71,7 @@ private MySQLInsertOnDuplicateUpdateExecutor insertOrUpdateExecutor; private final int pkIndex = 0; - private HashMap<String,Integer> pkIndexMap; + private HashMap<String, Integer> pkIndexMap; @BeforeEach public void init() { @@ -79,9 +84,10 @@ StatementCallback statementCallback = mock(StatementCallback.class); sqlInsertRecognizer = mock(SQLInsertRecognizer.class); tableMeta = mock(TableMeta.class); + when(tableMeta.getPrimaryKeyOnlyName()).thenReturn(Collections.singletonList(ID_COLUMN)); insertOrUpdateExecutor = Mockito.spy(new MySQLInsertOnDuplicateUpdateExecutor(statementProxy, statementCallback, sqlInsertRecognizer)); - pkIndexMap = new HashMap<String,Integer>(){ + pkIndexMap = new HashMap<String, Integer>() { { put(ID_COLUMN, pkIndex); } @@ -89,7 +95,7 @@ } @Test - public void TestBuildImageParameters(){ + public void TestBuildImageParameters() { mockParameters(); List<String> insertParamsList = new ArrayList<>(); insertParamsList.add("?,?,?,?"); @@ -97,11 +103,11 @@ when(sqlInsertRecognizer.getInsertParamsValue()).thenReturn(insertParamsList); mockInsertColumns(); Map<String, ArrayList<Object>> imageParameterMap = insertOrUpdateExecutor.buildImageParameters(sqlInsertRecognizer); - Assertions.assertEquals(imageParameterMap.toString(),mockImageParameterMap().toString()); + Assertions.assertEquals(imageParameterMap.toString(), mockImageParameterMap().toString()); } @Test - public void TestBuildImageParameters_contain_constant(){ + public void TestBuildImageParameters_contain_constant() { mockImageParameterMap_contain_constant(); List<String> insertParamsList = new ArrayList<>(); insertParamsList.add("?,?,?,userStatus1"); @@ -109,13 +115,13 @@ when(sqlInsertRecognizer.getInsertParamsValue()).thenReturn(insertParamsList); mockInsertColumns(); Map<String, ArrayList<Object>> imageParameterMap = insertOrUpdateExecutor.buildImageParameters(sqlInsertRecognizer); - Assertions.assertEquals(imageParameterMap.toString(),mockImageParameterMap().toString()); + Assertions.assertEquals(imageParameterMap.toString(), mockImageParameterMap().toString()); } @Test - public void testBuildImageSQL(){ - String selectSQLStr = "SELECT * FROM null WHERE (user_id = ? ) OR (id = ? ) OR (user_id = ? ) OR (id = ? ) "; - String paramAppenderListStr = "[[userId1, 100], [userId2, 101]]"; + public void testBuildImageSQL() { + String selectSQLStr = "SELECT * FROM null WHERE (user_id) in((?),(?)) OR (id) in((?),(?)) "; + String paramAppenderListStr = "{[user_id]=[userId1, userId2], [id]=[100, 101]}"; mockImageParameterMap_contain_constant(); List<String> insertParamsList = new ArrayList<>(); insertParamsList.add("?,?,?,userStatus1"); @@ -124,12 +130,12 @@ mockInsertColumns(); mockAllIndexes(); String selectSQL = insertOrUpdateExecutor.buildImageSQL(tableMeta); - Assertions.assertEquals(selectSQLStr,selectSQL); - Assertions.assertEquals(paramAppenderListStr,insertOrUpdateExecutor.getParamAppenderList().toString()); + Assertions.assertEquals(selectSQLStr, selectSQL); + Assertions.assertEquals(paramAppenderListStr, insertOrUpdateExecutor.getParamAppenderMap().toString()); } @Test - public void testBeforeImage(){ + public void testBeforeImage() { mockImageParameterMap_contain_constant(); List<String> insertParamsList = new ArrayList<>(); insertParamsList.add("?,?,?,userStatus1"); @@ -141,16 +147,109 @@ try { TableRecords tableRecords = new TableRecords(); String selectSQL = insertOrUpdateExecutor.buildImageSQL(tableMeta); - ArrayList<List<Object>> paramAppenderList = insertOrUpdateExecutor.getParamAppenderList(); - doReturn(tableRecords).when(insertOrUpdateExecutor).buildTableRecords2(tableMeta,selectSQL,paramAppenderList, Collections.emptyList()); + HashMap<List<String>, List<Object>> paramAppenderMap = insertOrUpdateExecutor.getParamAppenderMap(); + doReturn(tableRecords).when(insertOrUpdateExecutor).buildTableRecords2(tableMeta, selectSQL, new ArrayList<>(paramAppenderMap.values()), Collections.emptyList()); + insertOrUpdateExecutor.setSelectSQL(selectSQL); TableRecords tableRecordsResult = insertOrUpdateExecutor.beforeImage(); - Assertions.assertEquals(tableRecords,tableRecordsResult); + Assertions.assertEquals(tableRecords, tableRecordsResult); } catch (SQLException throwables) { throwables.printStackTrace(); } } - private void mockAllIndexes(){ + @Test + public void testAfterImage() throws SQLException { + String selectSQL = "SELECT * FROM null WHERE (id) in((?),(?)) "; + String expectAppend = "OR (id) in((?),(?),(?))"; + Map<Integer, ArrayList<Object>> primaryIndexValueMap = new HashMap<>(); + ArrayList<Object> list = new ArrayList<>(); + TableRecords afterImage = new TableRecords(); + doReturn(tableMeta).when(insertOrUpdateExecutor).getTableMeta(); + TableRecords beforeImage = getOnePkBeforeImage(3); + beforeImage.getRows().forEach(row -> row.getFields().forEach(field -> { + if (KeyType.PRIMARY_KEY.equals(field.getKeyType())) { + list.add(field.getValue()); + } + })); + primaryIndexValueMap.put(0, list); + HashMap<List<String>, List<Object>> paramAppenderMap = new HashMap<>(); + insertOrUpdateExecutor.setParamAppenderMap(paramAppenderMap); + insertOrUpdateExecutor.setSelectSQL(selectSQL); + ArrayList<Object> pkList = new ArrayList<>(); + primaryIndexValueMap.values().forEach(pkList::addAll); + doReturn(afterImage).when(insertOrUpdateExecutor).buildTableRecords2(tableMeta, selectSQL + expectAppend, new ArrayList<>(paramAppenderMap.values()), pkList); + TableRecords resultRecords = insertOrUpdateExecutor.afterImage(beforeImage); + Assertions.assertEquals(resultRecords, afterImage); + + TableRecords beforeImage01 = getTwoPkBeforeImage(3); + String expectAppend01 = "OR (id,user_id) in((?,?),(?,?),(?,?))"; + List<Row> rows = beforeImage01.getRows(); + Map<List<String>, ArrayList<Object>> primaryValueMap = new HashMap<>(1, 1.001f); + AtomicReference<List<String>> nameList = new AtomicReference<>(); + rows.forEach(m -> { + List<Field> fields = m.primaryKeys(); + nameList.set(fields.stream().map(Field::getName).collect(Collectors.toList())); + ArrayList<Object> tempList = new ArrayList<>(); + fields.forEach(f -> tempList.add(f.getValue())); + primaryValueMap.computeIfAbsent(nameList.get(), v -> new ArrayList<>()).addAll(tempList); + }); + ArrayList<Object> pkList1 = new ArrayList<>(); + primaryValueMap.values().forEach(pkList1::addAll); + doReturn(afterImage).when(insertOrUpdateExecutor).buildTableRecords2(tableMeta, selectSQL + expectAppend01, new ArrayList<>(paramAppenderMap.values()), pkList1); + TableRecords resultRecords01 = insertOrUpdateExecutor.afterImage(beforeImage01); + Assertions.assertEquals(resultRecords01, afterImage); + } + + private TableRecords getOnePkBeforeImage(int rowCount) { + TableRecords beforeImage = new TableRecords(); + ArrayList<Row> rows = new ArrayList<>(); + for (int i = 0; i < rowCount; i++) { + ArrayList<Field> fields = new ArrayList<>(); + Row row = new Row(); + Field field = new Field(); + field.setKeyType(KeyType.PRIMARY_KEY); + field.setValue("id" + i); + field.setName("id"); + fields.add(field); + Field field01 = new Field(); + field01.setKeyType(KeyType.NULL); + field01.setValue("userName" + i); + field01.setName("name"); + fields.add(field01); + row.setFields(fields); + rows.add(row); + } + beforeImage.setRows(rows); + return beforeImage; + } + + private TableRecords getTwoPkBeforeImage(int rowCount) { + TableRecords beforeImage = new TableRecords(); + ArrayList<Row> rows = new ArrayList<>(); + for (int i = 0; i < rowCount; i++) { + ArrayList<Field> fields = new ArrayList<>(); + Row row = new Row(); + Field field = new Field(); + field.setKeyType(KeyType.PRIMARY_KEY); + field.setValue("id" + i); + field.setName("id"); + fields.add(field); + Field field01 = new Field(); + field01.setKeyType(KeyType.PRIMARY_KEY); + field01.setValue("userId" + i); + field01.setName("user_id"); + fields.add(field01); + row.setFields(fields); + rows.add(row); + } + beforeImage.setRows(rows); + return beforeImage; + } + + + private void mockAllIndexes() { + List<String> primaryKeyOnlyNameList = new ArrayList<>(); + primaryKeyOnlyNameList.add("id"); Map<String, IndexMeta> allIndex = new HashMap<>(); IndexMeta primary = new IndexMeta(); primary.setIndextype(IndexType.PRIMARY); @@ -166,10 +265,10 @@ unique.setValues(Lists.newArrayList(columnMetaUnique)); allIndex.put("user_id", unique); when(tableMeta.getAllIndexes()).thenReturn(allIndex); + when(tableMeta.getPrimaryKeyOnlyName()).thenReturn(primaryKeyOnlyNameList); } - private List<String> mockInsertColumns() { List<String> columns = new ArrayList<>(); columns.add(ID_COLUMN); @@ -185,7 +284,7 @@ * {1=[100], 2=[userId1], 3=[userName1], 4=[userStatus1], 5=[101], 6=[userId2], 7=[userName2], 8=[userStatus2]} */ private void mockParameters() { - Map<Integer,ArrayList<Object>> paramters = new HashMap<>(4); + Map<Integer, ArrayList<Object>> paramters = new HashMap<>(4); ArrayList arrayList10 = new ArrayList<>(); arrayList10.add(PK_VALUE); ArrayList arrayList11 = new ArrayList<>(); @@ -199,7 +298,7 @@ paramters.put(3, arrayList12); paramters.put(4, arrayList13); ArrayList arrayList20 = new ArrayList<>(); - arrayList20.add(PK_VALUE+1); + arrayList20.add(PK_VALUE + 1); ArrayList arrayList21 = new ArrayList<>(); arrayList21.add("userId2"); ArrayList arrayList22 = new ArrayList<>(); @@ -219,7 +318,7 @@ * {1=[100], 2=[userId1], 3=[userName1], 4=[101], 5=[userId2], 6=[userName2]} */ private void mockImageParameterMap_contain_constant() { - Map<Integer,ArrayList<Object>> paramters = new HashMap<>(4); + Map<Integer, ArrayList<Object>> paramters = new HashMap<>(4); ArrayList arrayList10 = new ArrayList<>(); arrayList10.add(PK_VALUE); ArrayList arrayList11 = new ArrayList<>(); @@ -230,7 +329,7 @@ paramters.put(2, arrayList11); paramters.put(3, arrayList12); ArrayList arrayList20 = new ArrayList<>(); - arrayList20.add(PK_VALUE+1); + arrayList20.add(PK_VALUE + 1); ArrayList arrayList21 = new ArrayList<>(); arrayList21.add("userId2"); ArrayList arrayList22 = new ArrayList<>(); @@ -242,29 +341,29 @@ when(psp.getParameters()).thenReturn(paramters); } - private Map<String, ArrayList<Object>> mockImageParameterMap(){ + private Map<String, ArrayList<Object>> mockImageParameterMap() { Map<String, ArrayList<Object>> imageParameterMap = new LinkedHashMap<>(); ArrayList<Object> idList = new ArrayList<>(); idList.add("100"); idList.add("101"); - imageParameterMap.put("id",idList); + imageParameterMap.put("id", idList); ArrayList<Object> user_idList = new ArrayList<>(); user_idList.add("userId1"); user_idList.add("userId2"); - imageParameterMap.put("user_id",user_idList); + imageParameterMap.put("user_id", user_idList); ArrayList<Object> user_nameList = new ArrayList<>(); user_nameList.add("userName1"); user_nameList.add("userName2"); - imageParameterMap.put("user_name",user_nameList); + imageParameterMap.put("user_name", user_nameList); ArrayList<Object> user_statusList = new ArrayList<>(); user_statusList.add("userStatus1"); user_statusList.add("userStatus2"); - imageParameterMap.put("user_status",user_statusList); + imageParameterMap.put("user_status", user_statusList); return imageParameterMap; } private void mockParametersOfOnePk() { - Map<Integer,ArrayList<Object>> paramters = new HashMap<>(4); + Map<Integer, ArrayList<Object>> paramters = new HashMap<>(4); ArrayList arrayList1 = new ArrayList<>(); arrayList1.add(PK_VALUE); paramters.put(1, arrayList1);
diff --git a/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertSelectExecutorTest.java b/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertSelectExecutorTest.java new file mode 100644 index 0000000..741b112 --- /dev/null +++ b/rm-datasource/src/test/java/io/seata/rm/datasource/exec/MySQLInsertSelectExecutorTest.java
@@ -0,0 +1,319 @@ +/* + * Copyright 1999-2019 Seata.io Group. + * + * Licensed 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 io.seata.rm.datasource.exec; + +import com.google.common.collect.Lists; +import io.seata.rm.datasource.ConnectionProxy; +import io.seata.rm.datasource.DataSourceProxy; +import io.seata.rm.datasource.PreparedStatementProxy; +import io.seata.rm.datasource.StatementProxy; +import io.seata.rm.datasource.exec.mysql.MySQLInsertSelectExecutor; +import io.seata.rm.datasource.mock.MockConnection; +import io.seata.rm.datasource.mock.MockDataSource; +import io.seata.rm.datasource.mock.MockDriver; +import io.seata.rm.datasource.sql.struct.ColumnMeta; +import io.seata.rm.datasource.sql.struct.Field; +import io.seata.rm.datasource.sql.struct.IndexMeta; +import io.seata.rm.datasource.sql.struct.IndexType; +import io.seata.rm.datasource.sql.struct.Row; +import io.seata.rm.datasource.sql.struct.TableMeta; +import io.seata.rm.datasource.sql.struct.TableMetaCacheFactory; +import io.seata.rm.datasource.sql.struct.TableRecords; +import io.seata.sqlparser.SQLInsertRecognizer; +import io.seata.sqlparser.SQLRecognizer; +import io.seata.sqlparser.util.JdbcConstants; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mockito; + +import java.sql.SQLException; +import java.sql.Types; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static org.mockito.Mockito.doReturn; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +/** + * @author: lyx + */ +class MySQLInsertSelectExecutorTest { + + private static final String ID_COLUMN = "id"; + private static final String USER_ID_COLUMN = "user_id"; + private static final String USER_NAME_COLUMN = "user_name"; + private static final String USER_STATUS_COLUMN = "user_status"; + private static final List<Integer> PK_VALUE_LIST = Arrays.asList(100, 200); + + private StatementProxy statementProxy; + + private SQLInsertRecognizer sqlInsertRecognizer; + + private TableMeta tableMeta; + + private MySQLInsertSelectExecutorMock insertSelectExecutorMock; + + private final int pkIndex = 0; + private HashMap<String, Integer> pkIndexMap; + + private String selectTableName = "test02"; + + private String selectSQL = "select *\n" + "from " + selectTableName + "\n" + "where id = ?"; + + private String tableName = "test01"; + + private String originSQL = "insert ignore into " + tableName + "\n" + + "select * from test02\n where id in (?,?);"; + + private TableMeta selectTableMeta; + + @BeforeEach + public void init() throws SQLException { + List<String> returnValueColumnLabels = Lists.newArrayList("id"); + Object[][] returnValue = new Object[][]{ + new Object[]{1}, + }; + Object[][] columnMetas = new Object[][]{ + new Object[]{"", "", selectTableName, "id", Types.INTEGER, "INTEGER", 64, 0, 10, 1, "", "", 0, 0, 64, 1, "NO", "YES"}, + }; + Object[][] indexMetas = new Object[][]{ + new Object[]{"PRIMARY", "id", false, "", 3, 1, "A", 34}, + }; + MockDriver mockDriver = new MockDriver(returnValueColumnLabels, returnValue, columnMetas, indexMetas); + ConnectionProxy connectionProxy = mock(ConnectionProxy.class); + when(connectionProxy.getDbType()).thenReturn(JdbcConstants.MYSQL); + DataSourceProxy dataSourceProxy = new DataSourceProxy(new MockDataSource()); + when(connectionProxy.getDataSourceProxy()).thenReturn(dataSourceProxy); + MockConnection mockConnection = new MockConnection(mockDriver, "", null); + when(connectionProxy.getTargetConnection()).thenReturn(mockConnection); + selectTableMeta = TableMetaCacheFactory.getTableMetaCache(connectionProxy.getDbType()) + .getTableMeta(connectionProxy.getTargetConnection(), selectTableName, connectionProxy.getDataSourceProxy().getResourceId()); + + statementProxy = mock(PreparedStatementProxy.class); + when(statementProxy.getConnectionProxy()).thenReturn(connectionProxy); + when(statementProxy.getConnection()).thenReturn(connectionProxy); + + StatementCallback statementCallback = mock(StatementCallback.class); + sqlInsertRecognizer = mock(SQLInsertRecognizer.class); + when(sqlInsertRecognizer.getQuerySQL()).thenReturn(selectSQL); + when(sqlInsertRecognizer.getOriginalSQL()).thenReturn(originSQL); + when(sqlInsertRecognizer.getTableName()).thenReturn(tableName); + when(sqlInsertRecognizer.isIgnore()).thenReturn(true); + + this.tableMeta = mock(TableMeta.class); + insertSelectExecutorMock = Mockito.spy(new MySQLInsertSelectExecutorMock(statementProxy, statementCallback, sqlInsertRecognizer)); + doReturn(tableMeta).when(insertSelectExecutorMock).getTableMeta(tableName); + + pkIndexMap = new HashMap<String, Integer>() { + { + put(ID_COLUMN, pkIndex); + } + }; + } + + @Test + void getInsertRows() throws SQLException { + String expectResult = "[[100, NOT_PLACEHOLDER, NOT_PLACEHOLDER, NOT_PLACEHOLDER], [200, NOT_PLACEHOLDER, NULL, NOT_PLACEHOLDER]]"; + mockEmptyRecognize(); + List<List<Object>> insertRows = insertSelectExecutorMock.getInsertRows(pkIndexMap.values()); + Assertions.assertEquals(expectResult, insertRows.toString()); + } + + @Test + void getInsertParamsValue() throws SQLException { + String expectResult = "[100, userId1, userName1, userStatus1, 200, userId2, NULL, userStatus2]"; + mockEmptyRecognize(); + List<String> insertParamsValue = insertSelectExecutorMock.getInsertParamsValue(); + Assertions.assertEquals(expectResult, insertParamsValue.toString()); + } + + @Test + public void testGetPkValuesByColumn() throws SQLException { + mockEmptyRecognize(); + mockInsertColumns(); + when(tableMeta.getPrimaryKeyOnlyName()).thenReturn(Arrays.asList(new String[]{ID_COLUMN})); + doReturn(pkIndexMap).when(insertSelectExecutorMock).getPkIndex(); + Map<String, List<Object>> pkValuesList = insertSelectExecutorMock.getPkValuesByColumn(); + Assertions.assertIterableEquals(pkValuesList.get(ID_COLUMN), PK_VALUE_LIST); + } + + @Test + public void buildImageSQLTest() throws SQLException { + String expectSelectSQL = "SELECT * FROM " + tableName + " WHERE (user_id) in((?),(?)) OR (id) in((?),(?)) "; + HashMap<List<String>, List<Object>> paramAppenderMap = new HashMap<>(); + paramAppenderMap.put(Collections.singletonList("id"), Arrays.asList(PK_VALUE_LIST.get(0), PK_VALUE_LIST.get(1))); + paramAppenderMap.put(Collections.singletonList("user_id"), Arrays.asList("userId1", "userId2")); + mockGetInsertParamsValue(); + mockParametersOfOnePk(); + mockInsertColumns(); + mockAllIndexes(); + + TableRecords tableRecords = new TableRecords(); + mockRecordsRow(tableRecords); + + ArrayList<Object> list = new ArrayList<>(); + list.add(PK_VALUE_LIST.get(0)); + ArrayList<Object> list1 = new ArrayList<>(); + list1.add(PK_VALUE_LIST.get(1)); + ArrayList<List<Object>> paramFromSql = new ArrayList<>(); + paramFromSql.add(list); + paramFromSql.add(list1); + doReturn(tableRecords).when(insertSelectExecutorMock).buildTableRecords2(selectTableMeta, selectSQL, paramFromSql, Collections.emptyList()); + insertSelectExecutorMock.superCreateInsertRecognizer(); + + String imagesSQL = insertSelectExecutorMock.buildImageSQL(tableMeta); + Assertions.assertEquals(expectSelectSQL, imagesSQL); + Assertions.assertEquals(paramAppenderMap.toString(), insertSelectExecutorMock.getParamAppenderMap().toString()); + } + + @Test + public void beforeImageTest() throws SQLException { + mockInsertColumns(); + mockAllIndexes(); + mockEmptyRecognize(); + String imagesSQL = insertSelectExecutorMock.buildImageSQL(tableMeta); + HashMap<List<String>, List<Object>> paramAppenderMap = insertSelectExecutorMock.getParamAppenderMap(); + insertSelectExecutorMock.setSelectSQL(imagesSQL); + TableRecords tableRecords = new TableRecords(); + doReturn(tableRecords).when(insertSelectExecutorMock).buildTableRecords2(tableMeta, imagesSQL, new ArrayList<>(paramAppenderMap.values()), Collections.emptyList()); + TableRecords resultRecords = insertSelectExecutorMock.beforeImage(); + Assertions.assertEquals(resultRecords, tableRecords); + } + + private void mockEmptyRecognize() throws SQLException { + TableRecords tableRecords = new TableRecords(); + mockRecordsRow(tableRecords); + doReturn(tableRecords).when(insertSelectExecutorMock).buildTableRecords2(selectTableMeta, selectSQL, new ArrayList<>(Collections.EMPTY_LIST), Collections.emptyList()); + insertSelectExecutorMock.superCreateInsertRecognizer(); + } + + private List<String> mockInsertColumns() { + List<String> columns = new ArrayList<>(); + columns.add(ID_COLUMN); + columns.add(USER_ID_COLUMN); + columns.add(USER_NAME_COLUMN); + columns.add(USER_STATUS_COLUMN); + when(sqlInsertRecognizer.getInsertColumns()).thenReturn(columns); + return columns; + } + + private void mockGetInsertParamsValue() { + List<String> insertParamsList = new ArrayList<>(); + insertParamsList.add("?,?,?,userStatus1"); + insertParamsList.add("?,?,?,userStatus2"); + when(sqlInsertRecognizer.getInsertParamsValue()).thenReturn(insertParamsList); + } + + private void mockParametersOfOnePk() { + Map<Integer, ArrayList<Object>> paramters = new HashMap<>(4); + ArrayList arrayList1 = new ArrayList<>(); + arrayList1.add(PK_VALUE_LIST.get(0)); + paramters.put(1, arrayList1); + ArrayList arrayList2 = new ArrayList<>(); + arrayList2.add(PK_VALUE_LIST.get(1)); + paramters.put(2, arrayList2); + PreparedStatementProxy psp = (PreparedStatementProxy) this.statementProxy; + when(psp.getParameters()).thenReturn(paramters); + } + + private void mockAllIndexes() { + List<String> primaryKeyOnlyNameList = new ArrayList<>(); + primaryKeyOnlyNameList.add("id"); + + Map<String, IndexMeta> allIndex = new LinkedHashMap<>(); + IndexMeta primary = new IndexMeta(); + primary.setIndextype(IndexType.PRIMARY); + ColumnMeta columnMeta = new ColumnMeta(); + columnMeta.setColumnName("id"); + primary.setValues(Lists.newArrayList(columnMeta)); + allIndex.put("id", primary); + + IndexMeta unique = new IndexMeta(); + unique.setIndextype(IndexType.PRIMARY); + ColumnMeta columnMetaUnique = new ColumnMeta(); + columnMetaUnique.setColumnName("user_id"); + unique.setValues(Lists.newArrayList(columnMetaUnique)); + allIndex.put("user_id", unique); + when(tableMeta.getAllIndexes()).thenReturn(allIndex); + when(tableMeta.getPrimaryKeyOnlyName()).thenReturn(primaryKeyOnlyNameList); + } + + private void mockRecordsRow(TableRecords tableRecords) { + List<Row> rows = new ArrayList<>(); + Row row01 = new Row(); + Field field01 = new Field(); + field01.setValue(PK_VALUE_LIST.get(0)); + Field field02 = new Field(); + field02.setValue("userId1"); + Field field03 = new Field(); + field03.setValue("userName1"); + Field field04 = new Field(); + field04.setValue("userStatus1"); + List<Field> fields = new ArrayList<>(); + fields.add(field01); + fields.add(field02); + fields.add(field03); + fields.add(field04); + row01.setFields(fields); + rows.add(row01); + + Row row02 = new Row(); + Field field001 = new Field(); + field001.setValue(PK_VALUE_LIST.get(1)); + Field field002 = new Field(); + field002.setValue("userId2"); + Field field003 = new Field(); + field003.setValue(null); + Field field004 = new Field(); + field004.setValue("userStatus2"); + List<Field> fields01 = new ArrayList<>(); + fields01.add(field001); + fields01.add(field002); + fields01.add(field003); + fields01.add(field004); + row02.setFields(fields01); + rows.add(row02); + + tableRecords.setRows(rows); + } + + /** + * the class for mock + */ + class MySQLInsertSelectExecutorMock extends MySQLInsertSelectExecutor { + + public MySQLInsertSelectExecutorMock(StatementProxy statementProxy, StatementCallback statementCallback, SQLRecognizer sqlRecognizer) throws SQLException { + super(statementProxy, statementCallback, sqlRecognizer); + } + + // just for mock + @Override + public void createInsertRecognizer() throws SQLException { + } + + // just for test + public void superCreateInsertRecognizer() throws SQLException { + super.createInsertRecognizer(); + } + } +} \ No newline at end of file
diff --git a/sqlparser/seata-sqlparser-antlr/src/main/java/io/seata/sqlparser/antlr/mysql/AntlrMySQLInsertRecognizer.java b/sqlparser/seata-sqlparser-antlr/src/main/java/io/seata/sqlparser/antlr/mysql/AntlrMySQLInsertRecognizer.java index 894e673..6ebd819 100644 --- a/sqlparser/seata-sqlparser-antlr/src/main/java/io/seata/sqlparser/antlr/mysql/AntlrMySQLInsertRecognizer.java +++ b/sqlparser/seata-sqlparser-antlr/src/main/java/io/seata/sqlparser/antlr/mysql/AntlrMySQLInsertRecognizer.java
@@ -113,4 +113,19 @@ List<String> insertColumns = getInsertColumns(); return ColumnUtils.delEscape(insertColumns, JdbcConstants.MYSQL); } + + @Override + public String getQuerySQL() { + return null; + } + + @Override + public String getHintColumnName() { + return null; + } + + @Override + public boolean isIgnore() { + return false; + } } \ No newline at end of file
diff --git a/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLInsertRecognizer.java b/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLInsertRecognizer.java index c9a6d43..5ce9634 100644 --- a/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLInsertRecognizer.java +++ b/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLInsertRecognizer.java
@@ -66,4 +66,25 @@ * @return (`a`, `b`, `c`) -> (a, b, c) */ List<String> getInsertColumnsIsSimplified(); + + /** + * Gets query sql + * + * @return the select sql after insert ; return null if not present + */ + String getQuerySQL(); + + /** + * Gets hint column name + * + * @return the hint column name ; return null if not present + */ + String getHintColumnName(); + + /** + * Gets if the sql is ignore + * + * @return true when the sql is ignore + */ + boolean isIgnore(); }
diff --git a/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLType.java b/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLType.java index 520fc87..b81ff46 100644 --- a/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLType.java +++ b/sqlparser/seata-sqlparser-core/src/main/java/io/seata/sqlparser/SQLType.java
@@ -215,7 +215,11 @@ /** * update join sql type */ - UPDATE_JOIN(103); + UPDATE_JOIN(103), + /** + * Insert select sql type. + */ + INSERT_SELECT(104); private int i; @@ -246,4 +250,9 @@ } throw new IllegalArgumentException("Invalid SQLType:" + i); } + + public String getName() { + return this.name(); + } + }
diff --git a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/BaseRecognizer.java b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/BaseRecognizer.java index ed4365c..8566773 100644 --- a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/BaseRecognizer.java +++ b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/BaseRecognizer.java
@@ -25,7 +25,6 @@ import com.alibaba.druid.sql.ast.expr.SQLInListExpr; import com.alibaba.druid.sql.ast.expr.SQLInSubQueryExpr; import com.alibaba.druid.sql.ast.expr.SQLMethodInvokeExpr; -import com.alibaba.druid.sql.ast.statement.SQLInsertStatement; import com.alibaba.druid.sql.ast.statement.SQLMergeStatement; import com.alibaba.druid.sql.ast.statement.SQLReplaceStatement; import com.alibaba.druid.sql.ast.statement.SQLSubqueryTableSource; @@ -146,16 +145,6 @@ throw new NotSupportYetException("not support the sql syntax with MergeStatement:" + x + "\nplease see the doc about SQL restrictions https://seata.io/zh-cn/docs/user/sqlreference/dml.html"); } - - @Override - public boolean visit(SQLInsertStatement x) { - if (null != x.getQuery()) { - //just like: insert into t select * from t1 - throw new NotSupportYetException("not support the sql syntax insert with query:" + x - + "\nplease see the doc about SQL restrictions https://seata.io/zh-cn/docs/user/sqlreference/dml.html"); - } - return true; - } }; getAst().accept(visitor); return true;
diff --git a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java index 13d0ff8..fc97ad0 100644 --- a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java +++ b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java
@@ -18,12 +18,16 @@ import java.util.ArrayList; import java.util.Collection; import java.util.List; +import java.util.Optional; + import com.alibaba.druid.sql.ast.SQLExpr; +import com.alibaba.druid.sql.ast.SQLObjectImpl; import com.alibaba.druid.sql.ast.SQLStatement; import com.alibaba.druid.sql.ast.expr.SQLBinaryOpExpr; import com.alibaba.druid.sql.ast.expr.SQLIdentifierExpr; import com.alibaba.druid.sql.ast.expr.SQLMethodInvokeExpr; import com.alibaba.druid.sql.ast.expr.SQLNullExpr; +import com.alibaba.druid.sql.ast.expr.SQLPropertyExpr; import com.alibaba.druid.sql.ast.expr.SQLValuableExpr; import com.alibaba.druid.sql.ast.expr.SQLVariantRefExpr; import com.alibaba.druid.sql.ast.statement.SQLExprTableSource; @@ -60,7 +64,9 @@ @Override public SQLType getSQLType() { - return CollectionUtils.isNotEmpty(ast.getDuplicateKeyUpdate()) ? SQLType.INSERT_ON_DUPLICATE_UPDATE : SQLType.INSERT; + return ast.getQuery() != null ? SQLType.INSERT_SELECT + : CollectionUtils.isNotEmpty(ast.getDuplicateKeyUpdate()) ? SQLType.INSERT_ON_DUPLICATE_UPDATE + : ast.isIgnore() ? SQLType.INSERT_IGNORE : SQLType.INSERT; } @Override @@ -161,6 +167,9 @@ SQLExpr expr = ((SQLBinaryOpExpr)exprLeft).getLeft(); if (expr instanceof SQLIdentifierExpr) { list.add(((SQLIdentifierExpr)expr).getName()); + } + else if (expr instanceof SQLPropertyExpr) { + list.add(((SQLPropertyExpr) expr).getName()); } else { wrapSQLParsingException(expr); } @@ -175,6 +184,21 @@ } @Override + public String getQuerySQL() { + return Optional.ofNullable(ast.getQuery()).map(SQLObjectImpl::toString).orElse(null); + } + + @Override + public String getHintColumnName() { + return null; + } + + @Override + public boolean isIgnore() { + return ast.isIgnore(); + } + + @Override protected SQLStatement getAst() { return ast; }
diff --git a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/BaseOracleRecognizer.java b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/BaseOracleRecognizer.java index da41ef4..b4ca515 100644 --- a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/BaseOracleRecognizer.java +++ b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/BaseOracleRecognizer.java
@@ -26,7 +26,6 @@ import com.alibaba.druid.sql.ast.SQLOrderBy; import com.alibaba.druid.sql.ast.expr.SQLInSubQueryExpr; import com.alibaba.druid.sql.ast.expr.SQLVariantRefExpr; -import com.alibaba.druid.sql.ast.statement.SQLInsertStatement; import com.alibaba.druid.sql.ast.statement.SQLMergeStatement; import com.alibaba.druid.sql.ast.statement.SQLReplaceStatement; import com.alibaba.druid.sql.dialect.oracle.ast.stmt.OracleSelectJoin; @@ -175,15 +174,6 @@ + "\nplease see the doc about SQL restrictions https://seata.io/zh-cn/docs/user/sqlreference/dml.html"); } - @Override - public boolean visit(SQLInsertStatement x) { - if (null != x.getQuery()) { - //just like: insert into t select * from t1 - throw new NotSupportYetException("not support the sql syntax insert with query:" + x - + "\nplease see the doc about SQL restrictions https://seata.io/zh-cn/docs/user/sqlreference/dml.html"); - } - return true; - } }; getAst().accept(visitor); return true;
diff --git a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/OracleInsertRecognizer.java b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/OracleInsertRecognizer.java index 424ee26..b1bed17 100644 --- a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/OracleInsertRecognizer.java +++ b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/oracle/OracleInsertRecognizer.java
@@ -18,7 +18,11 @@ import java.util.ArrayList; import java.util.Collection; import java.util.List; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicReference; + import com.alibaba.druid.sql.ast.SQLExpr; +import com.alibaba.druid.sql.ast.SQLObjectImpl; import com.alibaba.druid.sql.ast.SQLStatement; import com.alibaba.druid.sql.ast.expr.SQLIdentifierExpr; import com.alibaba.druid.sql.ast.expr.SQLMethodInvokeExpr; @@ -30,8 +34,10 @@ import com.alibaba.druid.sql.ast.statement.SQLInsertStatement; import com.alibaba.druid.sql.dialect.oracle.ast.stmt.OracleInsertStatement; import com.alibaba.druid.sql.dialect.oracle.visitor.OracleOutputVisitor; +import io.seata.common.exception.ShouldNeverHappenException; import io.seata.common.util.CollectionUtils; import io.seata.sqlparser.util.ColumnUtils; +import io.seata.common.util.StringUtils; import io.seata.sqlparser.SQLInsertRecognizer; import io.seata.sqlparser.SQLType; import io.seata.sqlparser.struct.NotPlaceholderExpr; @@ -48,6 +54,16 @@ private final OracleInsertStatement ast; + private static final String PREFIX = "/*+"; + + private static final String SUFFIX = "*/"; + + private static final String IGNORE_HINT = "IGNORE_ROW_ON_DUPKEY_INDEX("; + + private static final char ESCAPE = '"'; + + private String hintColumnName; + /** * Instantiates a new My sql insert recognizer. * @@ -57,11 +73,13 @@ public OracleInsertRecognizer(String originalSQL, SQLStatement ast) { super(originalSQL); this.ast = (OracleInsertStatement)ast; + this.hintColumnName = getHintColumn(); } @Override public SQLType getSQLType() { - return SQLType.INSERT; + return ast.getQuery() != null ? SQLType.INSERT_SELECT : + StringUtils.isNotBlank(hintColumnName) ? SQLType.INSERT_IGNORE : SQLType.INSERT; } @Override @@ -143,7 +161,17 @@ @Override public List<String> getInsertParamsValue() { - return null; + List<SQLInsertStatement.ValuesClause> valuesList = ast.getValuesList(); + List<String> list = new ArrayList<>(); + for (SQLInsertStatement.ValuesClause m : valuesList) { + String values = m.toString().replace("VALUES", "").trim(); + // when all params is constant, the length of values less than 1 + if (values.length() > 1) { + values = values.substring(1, values.length() - 1); + } + list.add(values); + } + return list; } @Override @@ -158,7 +186,51 @@ } @Override + public String getQuerySQL() { + return Optional.ofNullable(ast.getQuery()).map(SQLObjectImpl::toString).orElse(null); + } + + @Override + public String getHintColumnName() { + return hintColumnName; + } + + @Override + public boolean isIgnore() { + return StringUtils.isNotBlank(hintColumnName); + } + + @Override protected SQLStatement getAst() { return ast; } + + /** + * get hint column name + * + * @return column name + */ + private String getHintColumn() { + AtomicReference<String> columnName = new AtomicReference<>(); + ast.getHints().forEach(sqlHint -> { + String hint = sqlHint.toString(); + if (hint.startsWith(PREFIX) && hint.endsWith(SUFFIX)) { + hint = hint.replaceAll(" ", ""); + StringBuilder matchHint = new StringBuilder(IGNORE_HINT); + int startIndex = hint.indexOf(matchHint.toString()); + if (startIndex != -1) { + int nextStartIndex = hint.indexOf("(", startIndex + matchHint.length()); + String tableName = hint.substring(startIndex + matchHint.length(), nextStartIndex); + if (!getTableName().equals(tableName) && !(ESCAPE + getTableName() + ESCAPE).equals(tableName)) { + throw new ShouldNeverHappenException("in IGNORE_ROW_ON_DUPKEY_INDEX hint,the table name should as same as what you insert"); + } + int endIndex = hint.indexOf(")", nextStartIndex + 1); + if (endIndex != -1) { + columnName.set(hint.substring(nextStartIndex + 1, endIndex)); + } + } + } + }); + return columnName.get(); + } }
diff --git a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/BasePostgresqlRecognizer.java b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/BasePostgresqlRecognizer.java index f44307a..b2da653 100644 --- a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/BasePostgresqlRecognizer.java +++ b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/BasePostgresqlRecognizer.java
@@ -23,7 +23,6 @@ import com.alibaba.druid.sql.ast.statement.SQLMergeStatement; import com.alibaba.druid.sql.ast.statement.SQLReplaceStatement; import com.alibaba.druid.sql.ast.statement.SQLSubqueryTableSource; -import com.alibaba.druid.sql.dialect.postgresql.ast.stmt.PGInsertStatement; import com.alibaba.druid.sql.dialect.postgresql.ast.stmt.PGUpdateStatement; import com.alibaba.druid.sql.dialect.postgresql.visitor.PGASTVisitor; import com.alibaba.druid.sql.dialect.postgresql.visitor.PGASTVisitorAdapter; @@ -118,15 +117,6 @@ + "\nplease see the doc about SQL restrictions https://seata.io/zh-cn/docs/user/sqlreference/dml.html"); } - @Override - public boolean visit(PGInsertStatement x) { - if (null != x.getQuery()) { - //just like: insert into t select * from t1 - throw new NotSupportYetException("not support the sql syntax insert with query:" + x - + "\nplease see the doc about SQL restrictions https://seata.io/zh-cn/docs/user/sqlreference/dml.html"); - } - return true; - } }; getAst().accept(visitor); return true;
diff --git a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/PostgresqlInsertRecognizer.java b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/PostgresqlInsertRecognizer.java index 79e6835..022477a 100644 --- a/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/PostgresqlInsertRecognizer.java +++ b/sqlparser/seata-sqlparser-druid/src/main/java/io/seata/sqlparser/druid/postgresql/PostgresqlInsertRecognizer.java
@@ -162,6 +162,21 @@ } @Override + public String getQuerySQL() { + return null; + } + + @Override + public String getHintColumnName() { + return null; + } + + @Override + public boolean isIgnore() { + return false; + } + + @Override public List<String> getInsertColumnsIsSimplified() { List<String> insertColumns = getInsertColumns(); return ColumnUtils.delEscape(insertColumns, getDbType());
diff --git a/sqlparser/seata-sqlparser-druid/src/test/java/io/seata/sqlparser/druid/DruidSQLRecognizerFactoryTest.java b/sqlparser/seata-sqlparser-druid/src/test/java/io/seata/sqlparser/druid/DruidSQLRecognizerFactoryTest.java index b4397fb..bd5e2bc 100644 --- a/sqlparser/seata-sqlparser-druid/src/test/java/io/seata/sqlparser/druid/DruidSQLRecognizerFactoryTest.java +++ b/sqlparser/seata-sqlparser-druid/src/test/java/io/seata/sqlparser/druid/DruidSQLRecognizerFactoryTest.java
@@ -66,9 +66,9 @@ Assertions.assertNotNull(recognizerFactory.create(sql6, JdbcConstants.POSTGRESQL)); String sql7 = "insert into a select * from b"; - Assertions.assertThrows(NotSupportYetException.class, () -> recognizerFactory.create(sql7, JdbcConstants.MYSQL)); - Assertions.assertThrows(NotSupportYetException.class, () -> recognizerFactory.create(sql7, JdbcConstants.ORACLE)); - Assertions.assertThrows(NotSupportYetException.class, () -> recognizerFactory.create(sql7, JdbcConstants.POSTGRESQL)); + Assertions.assertNotNull(recognizerFactory.create(sql7, JdbcConstants.MYSQL)); + Assertions.assertNotNull(recognizerFactory.create(sql7, JdbcConstants.ORACLE)); + Assertions.assertNotNull(recognizerFactory.create(sql7, JdbcConstants.POSTGRESQL)); String sql8 = "delete from t where id = ?"; Assertions.assertNotNull(recognizerFactory.create(sql8, JdbcConstants.MYSQL));
diff --git a/sqlparser/seata-sqlparser-druid/src/test/java/io/seata/sqlparser/druid/OracleInsertRecognizerTest.java b/sqlparser/seata-sqlparser-druid/src/test/java/io/seata/sqlparser/druid/OracleInsertRecognizerTest.java new file mode 100644 index 0000000..402a6c7 --- /dev/null +++ b/sqlparser/seata-sqlparser-druid/src/test/java/io/seata/sqlparser/druid/OracleInsertRecognizerTest.java
@@ -0,0 +1,51 @@ +/* + * Copyright 1999-2019 Seata.io Group. + * + * Licensed 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 io.seata.sqlparser.druid; + +import com.alibaba.druid.sql.ast.SQLStatement; +import io.seata.common.exception.ShouldNeverHappenException; +import io.seata.sqlparser.druid.oracle.OracleInsertRecognizer; +import io.seata.sqlparser.util.JdbcConstants; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +/** + * @author: lyx + */ +public class OracleInsertRecognizerTest extends AbstractRecognizerTest { + + @Test + public void testGetHintColumnName() { + String sql = "INSERT /*+ IGNORE_ROW_ON_DUPKEY_INDEX(TEST_TABLE(TEST_COLUMN)) */ INTO TEST_TABLE(TEST_COLUMN) VALUES ('TEST_COLUMN1')"; + SQLStatement statement = getSQLStatement(sql); + OracleInsertRecognizer oracleInsertRecognizer = new OracleInsertRecognizer(sql, statement); + Assertions.assertEquals("TEST_COLUMN", oracleInsertRecognizer.getHintColumnName()); + + String sql01 = "INSERT /*+ IGNORE_ROW_ON_DUPKEY_INDEX(\"TEST_TABLE\"(\"TEST_COLUMN\")) */ INTO TEST_TABLE(TEST_COLUMN) VALUES ('TEST_COLUMN1')"; + SQLStatement statement01 = getSQLStatement(sql01); + OracleInsertRecognizer oracleInsertRecognizer01 = new OracleInsertRecognizer(sql01, statement01); + Assertions.assertEquals("\"TEST_COLUMN\"", oracleInsertRecognizer01.getHintColumnName()); + + String sql02 = "INSERT /*+ IGNORE_ROW_ON_DUPKEY_INDEX(\"ERROR\"(\"TEST_COLUMN\")) */ INTO TEST_TABLE(TEST_COLUMN) VALUES ('TEST_COLUMN1')"; + SQLStatement statement02 = getSQLStatement(sql02); + Assertions.assertThrows(ShouldNeverHappenException.class, () -> new OracleInsertRecognizer(sql01, statement02)); + } + + @Override + public String getDbType() { + return JdbcConstants.ORACLE; + } +}