Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
package ai.chat2db.plugin.oracle;

import ai.chat2db.plugin.oracle.identifier.OracleIdentifierProcessor;
import ai.chat2db.spi.IDbManager;
import ai.chat2db.spi.DefaultDBManager;
import ai.chat2db.community.domain.api.model.account.*;
Expand All @@ -20,7 +21,6 @@
import ai.chat2db.spi.model.request.TriggerMetadataRequest;
import ai.chat2db.spi.model.request.ViewMetadataRequest;
import ai.chat2db.spi.DefaultSQLExecutor;
import ai.chat2db.spi.util.SqlUtils;
import cn.hutool.core.date.DateUtil;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections4.CollectionUtils;
Expand Down Expand Up @@ -70,7 +70,7 @@ private void exportTables(Connection connection, String databaseName, String sch
public void exportTable(Connection connection, String databaseName, String schemaName, String tableName, AsyncContext asyncContext) throws SQLException {
String tableDDL = Chat2DBContext.getDbMetaData().tableDDL(connection,
new TableMetadataRequest(databaseName, schemaName, tableName));
String sqlBuilder = "DROP TABLE " + SqlUtils.quoteObjectName(tableName) + ";\n" + tableDDL + "\n";
String sqlBuilder = "DROP TABLE " + qualifiedName(schemaName, tableName, false) + ";\n" + tableDDL + "\n";
asyncContext.write(sqlBuilder);
if (asyncContext.isContainsData()) {
exportTableData(connection, databaseName, schemaName, tableName, asyncContext);
Expand Down Expand Up @@ -112,7 +112,7 @@ private void exportProcedure(Connection connection, String schemaName, String pr
}

private void exportTriggers(Connection connection, String schemaName, AsyncContext asyncContext) throws SQLException {
String sql = String.format(SQL_SELECT_TRIGGER_NAME_ALL_TRIGGERS, schemaName);
String sql = String.format(SQL_SELECT_TRIGGER_NAME_ALL_TRIGGERS, OracleIdentifierProcessor.INSTANCE.escapeString(schemaName));
try (PreparedStatement preparedStatement = connection.prepareStatement(sql); ResultSet resultSet = preparedStatement.executeQuery()) {
while (resultSet.next()) {
String triggerName = resultSet.getString("TRIGGER_NAME");
Expand Down Expand Up @@ -152,27 +152,35 @@ public void connectDatabase(Connection connection, String database) {
}
String schemaName = connectInfo.getSchemaName();
try {
DefaultSQLExecutor.getInstance().execute(connection, SQL_ALTER_SESSION_SET_CURRENT_SCHEMA + schemaName + "\"");
DefaultSQLExecutor.getInstance().execute(connection,
SQL_ALTER_SESSION_SET_CURRENT_SCHEMA
+ OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName));
} catch (SQLException e) {
log.error("connectDatabase error", e);
}
}

@Override
public void copyTable(Connection connection, String databaseName, String schemaName, String tableName, String newTableName, boolean copyData) throws SQLException {
String sql = "";
String source = qualifiedName(schemaName, tableName, true);
String target = qualifiedName(schemaName, newTableName, true);
String sql;
if (copyData) {
sql = "CREATE TABLE " + SqlUtils.quoteObjectName(newTableName) + " AS SELECT * FROM " + SqlUtils.quoteObjectName(tableName);
sql = "CREATE TABLE " + target + " AS SELECT * FROM " + source;
} else {
sql = "CREATE TABLE " + SqlUtils.quoteObjectName(newTableName) + " AS SELECT * FROM " + SqlUtils.quoteObjectName(tableName) + " WHERE 1=0";
sql = "CREATE TABLE " + target + " AS SELECT * FROM " + source + " WHERE 1=0";
}
DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> null);
}

@Override
public String dropTable(Connection connection, String databaseName, String schemaName, String tableName) {
String sql = "DROP TABLE " + SqlUtils.quoteObjectName(tableName);
return sql;
return "DROP TABLE " + qualifiedName(schemaName, tableName, false);
}

@Override
public String truncateTable(Connection connection, String databaseName, String schemaName, String tableName) {
return "TRUNCATE TABLE " + qualifiedName(schemaName, tableName, true);
}

@Override
Expand All @@ -182,7 +190,23 @@ public void exportTableData(Connection connection, String databaseName, String s

@Override
public void dropView(Connection connection, String databaseName, String schemaName, String viewName) {
String sql = "DROP VIEW " + SqlUtils.quoteObjectName(schemaName) + "." + SqlUtils.quoteObjectName(viewName);
String sql = "DROP VIEW " + qualifiedName(schemaName, viewName, false);
DefaultSQLExecutor.getInstance().execute(connection, sql, (resultSet) -> null);
}

private static String qualifiedName(String schemaName, String objectName, boolean normalizeQuotedObject) {
String normalizedObject = normalizeQuotedObject ? normalizeQuotedIdentifier(objectName) : objectName;
String quotedObject = OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways(normalizedObject);
if (StringUtils.isBlank(schemaName)) {
return quotedObject;
}
return OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName) + "." + quotedObject;
}

private static String normalizeQuotedIdentifier(String identifier) {
if (OracleIdentifierProcessor.INSTANCE.isQuoteIdentifier(identifier)) {
return OracleIdentifierProcessor.INSTANCE.removeIdentifierQuote(identifier);
}
return identifier;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
import ai.chat2db.plugin.oracle.enums.type.OracleDefaultValueEnum;
import ai.chat2db.plugin.oracle.enums.type.OracleIndexTypeEnum;
import ai.chat2db.plugin.oracle.value.OracleValueProcessor;
import ai.chat2db.community.tools.util.EasyStringUtils;
import ai.chat2db.community.tools.util.I18nUtils;
import ai.chat2db.spi.IDbMetaData;
import ai.chat2db.spi.ISQLIdentifierProcessor;
Expand Down Expand Up @@ -45,11 +44,11 @@
public class OracleMetaData extends DefaultMetaService implements IDbMetaData {


public static final ISQLIdentifierProcessor ORACLE_SQL_IDENTIFIER_PROCESSOR = new OracleIdentifierProcessor();
public static final ISQLIdentifierProcessor ORACLE_SQL_IDENTIFIER_PROCESSOR = OracleIdentifierProcessor.INSTANCE;

@Override
public List<Procedure> procedures(Connection connection, String databaseName, String schemaName) {
String sql = String.format(PROCEDURE_LIST_DDL, schemaName);
String sql = String.format(PROCEDURE_LIST_DDL, escapeSqlLiteral(schemaName));
ArrayList<Procedure> procedures = new ArrayList<>();
DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
Expand All @@ -69,11 +68,11 @@ public List<Schema> schemas(Connection connection, String databaseName) {

@Override
public String tableDDL(Connection connection, String databaseName, String schemaName, String tableName) {
String sql = String.format(TABLE_DDL_SQL, tableName, schemaName);
String tableCommentSql = String.format(TABLE_COMMENT_SQL, schemaName, tableName);
String tableColumnCommentSql = String.format(TABLE_COLUMN_COMMENT_SQL, schemaName, tableName);
String tableIndexSql = String.format(TABLE_INDEX_DDL_SQL, schemaName, tableName);
String PUIndexSql = String.format(PU_INDEX_NAME_SQL, schemaName, tableName);
String sql = String.format(TABLE_DDL_SQL, escapeSqlLiteral(tableName), escapeSqlLiteral(schemaName));
String tableCommentSql = String.format(TABLE_COMMENT_SQL, escapeSqlLiteral(schemaName), escapeSqlLiteral(tableName));
String tableColumnCommentSql = String.format(TABLE_COLUMN_COMMENT_SQL, escapeSqlLiteral(schemaName), escapeSqlLiteral(tableName));
String tableIndexSql = String.format(TABLE_INDEX_DDL_SQL, escapeSqlLiteral(schemaName), escapeSqlLiteral(tableName));
String PUIndexSql = String.format(PU_INDEX_NAME_SQL, escapeSqlLiteral(schemaName), escapeSqlLiteral(tableName));
StringBuilder ddlBuilder = new StringBuilder();
DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
try {
Expand All @@ -88,8 +87,8 @@ public String tableDDL(Connection connection, String databaseName, String schema
if (resultSet.next()) {
String tableComment = resultSet.getString("comments");
if (StringUtils.isNotBlank(tableComment)) {
ddlBuilder.append("\nCOMMENT ON TABLE ").append(SqlUtils.quoteObjectName(tableName)).append(" IS ")
.append(EasyStringUtils.escapeAndQuoteString(tableComment)).append(";");
ddlBuilder.append("\nCOMMENT ON TABLE ").append(qualifiedName(schemaName, tableName)).append(" IS ")
.append(quoteStringLiteral(tableComment)).append(";");
}
}
});
Expand All @@ -99,9 +98,9 @@ public String tableDDL(Connection connection, String databaseName, String schema
String columnComment = resultSet.getString("comments");
if (StringUtils.isNotBlank(columnComment)) {
ddlBuilder.append("\nCOMMENT ON COLUMN ")
.append(SqlUtils.quoteObjectName(tableName)).append(".")
.append(SqlUtils.quoteObjectName(columnName)).append(" IS ")
.append(EasyStringUtils.escapeAndQuoteString(columnComment)).append(";");
.append(qualifiedName(schemaName, tableName)).append(".")
.append(OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways(columnName)).append(" IS ")
.append(quoteStringLiteral(columnComment)).append(";");
}
}
});
Expand Down Expand Up @@ -159,7 +158,19 @@ static String buildTablesSql(String schemaName, String tableName) {
}

private static String escapeSqlLiteral(String value) {
return StringUtils.replace(value, "'", "''");
return OracleIdentifierProcessor.INSTANCE.escapeString(value);
}

private static String quoteStringLiteral(String value) {
return OracleIdentifierProcessor.INSTANCE.quoteStringLiteral(value);
}

private static String qualifiedName(String schemaName, String objectName) {
String quotedObject = OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways(objectName);
if (StringUtils.isBlank(schemaName)) {
return quotedObject;
}
return OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName) + "." + quotedObject;
}


Expand Down Expand Up @@ -206,7 +217,7 @@ public List<TableColumn> columns(Connection connection, String databaseName, Str

private Map<String, TableColumn> getTableColumns(Connection connection, String databaseName, String schemaName, String tableName) {
Map<String, TableColumn> tableColumns = new HashMap<>();
String sql = String.format(SELECT_TAB_COLS, schemaName, tableName);
String sql = String.format(SELECT_TAB_COLS, escapeSqlLiteral(schemaName), escapeSqlLiteral(tableName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
TableColumn tableColumn = new TableColumn();
Expand Down Expand Up @@ -261,7 +272,7 @@ private Map<String, TableColumn> getTableColumns(Connection connection, String d
@Override
public Function function(Connection connection, @NotEmpty String databaseName, String schemaName,
String functionName) {
String sql = String.format(ROUTINES_SQL, "FUNCTION", schemaName, functionName);
String sql = String.format(ROUTINES_SQL, "FUNCTION", escapeSqlLiteral(schemaName), escapeSqlLiteral(functionName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Function function = new Function();
function.setDatabaseName(databaseName);
Expand Down Expand Up @@ -291,7 +302,7 @@ public Function function(Connection connection, @NotEmpty String databaseName, S

@Override
public List<TableIndex> indexes(Connection connection, String databaseName, String schemaName, String tableName) {
String pkSql = String.format(SELECT_PK_SQL, schemaName, tableName);
String pkSql = String.format(SELECT_PK_SQL, escapeSqlLiteral(schemaName), escapeSqlLiteral(tableName));
Set<String> pkSet = new HashSet<>();
DefaultSQLExecutor.getInstance().execute(connection, pkSql, resultSet -> {
while (resultSet.next()) {
Expand All @@ -301,7 +312,7 @@ public List<TableIndex> indexes(Connection connection, String databaseName, Stri
}
);

String sql = String.format(SELECT_TABLE_INDEX, schemaName, tableName);
String sql = String.format(SELECT_TABLE_INDEX, escapeSqlLiteral(schemaName), escapeSqlLiteral(tableName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
LinkedHashMap<String, TableIndex> map = new LinkedHashMap();
while (resultSet.next()) {
Expand Down Expand Up @@ -360,7 +371,7 @@ private TableIndexColumn getTableIndexColumn(ResultSet resultSet) throws SQLExce
@Override
public List<Trigger> triggers(Connection connection, String databaseName, String schemaName) {
List<Trigger> triggers = new ArrayList<>();
return DefaultSQLExecutor.getInstance().execute(connection, String.format(TRIGGER_SQL_LIST, schemaName),
return DefaultSQLExecutor.getInstance().execute(connection, String.format(TRIGGER_SQL_LIST, escapeSqlLiteral(schemaName)),
resultSet -> {
while (resultSet.next()) {
String triggerName = resultSet.getString("TRIGGER_NAME");
Expand All @@ -378,7 +389,7 @@ public List<Trigger> triggers(Connection connection, String databaseName, String
@Override
public Trigger trigger(Connection connection, @NotEmpty String databaseName, String schemaName,
String triggerName) {
String sql = String.format(TRIGGER_DDL_SQL, triggerName, schemaName);
String sql = String.format(TRIGGER_DDL_SQL, escapeSqlLiteral(triggerName), escapeSqlLiteral(schemaName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Trigger trigger = new Trigger();
trigger.setDatabaseName(databaseName);
Expand All @@ -394,7 +405,7 @@ public Trigger trigger(Connection connection, @NotEmpty String databaseName, Str
@Override
public Procedure procedure(Connection connection, @NotEmpty String databaseName, String schemaName,
String procedureName) {
String sql = String.format(ROUTINES_SQL, "PROCEDURE", schemaName, procedureName);
String sql = String.format(ROUTINES_SQL, "PROCEDURE", escapeSqlLiteral(schemaName), escapeSqlLiteral(procedureName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Procedure procedure = new Procedure();
procedure.setDatabaseName(databaseName);
Expand Down Expand Up @@ -428,7 +439,7 @@ static void appendRoutineSourceText(StringBuilder bodyBuilder, String sourceText

@Override
public Table view(Connection connection, String databaseName, String schemaName, String viewName) {
String sql = String.format(VIEW_DDL_SQL, schemaName, viewName);
String sql = String.format(VIEW_DDL_SQL, escapeSqlLiteral(schemaName), escapeSqlLiteral(viewName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Table table = new Table();
table.setDatabaseName(databaseName);
Expand Down Expand Up @@ -459,7 +470,13 @@ public TableMeta getTableMeta(String databaseName, String schemaName, String tab

@Override
public String getMetaDataName(String... names) {
return Arrays.stream(names).filter(name -> StringUtils.isNotBlank(name)).map(name -> "\"" + name + "\"").collect(Collectors.joining("."));
List<String> identifiers = Arrays.stream(names)
.filter(StringUtils::isNotBlank)
.toList();
int first = Math.max(0, identifiers.size() - 2);
return identifiers.subList(first, identifiers.size()).stream()
.map(OracleIdentifierProcessor.INSTANCE::quoteIdentifierAlways)
.collect(Collectors.joining("."));
}


Expand Down Expand Up @@ -525,9 +542,9 @@ public ModifyViewConfiguration viewMeta(String databaseName, String schemaName)
StringBuilder sqlBuilder = new StringBuilder(100);
sqlBuilder.append(SQL_CREATE).append("view ");
if (StringUtils.isNotBlank(schemaName)) {
sqlBuilder.append("\"").append(schemaName).append("\"").append(".");
sqlBuilder.append(OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName)).append(".");
}
sqlBuilder.append("\"").append("undefined").append("\"");
sqlBuilder.append(OracleIdentifierProcessor.INSTANCE.quoteIdentifierAlways("undefined"));
sqlBuilder.append(" AS \n").append(sql).append(";");
configuration.setPreviewSql(sqlBuilder.toString());
configuration.setSql(sql);
Expand Down
Loading
Loading