Skip to content
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
package ai.chat2db.plugin.xugudb;

import ai.chat2db.plugin.xugudb.identifier.XugudbIdentifierProcessor;
import ai.chat2db.spi.IDbManager;
import ai.chat2db.spi.DefaultDBManager;
import ai.chat2db.community.domain.api.model.async.AsyncContext;
Expand Down Expand Up @@ -27,7 +28,7 @@ public class XUGUDBManager extends DefaultDBManager implements IDbManager {


private String format(String tableName) {
return "\"" + tableName + "\"";
return XugudbIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableName);
}


Expand All @@ -54,7 +55,7 @@ public void exportDatabase(Connection connection, String databaseName, String sc
}

private void exportTables(Connection connection, String schemaName, AsyncContext asyncContext) throws SQLException {
String sql = String.format(SQL_SELECT_TABLE_NAME_ALL_TABLES, schemaName);
String sql = String.format(SQL_SELECT_TABLE_NAME_ALL_TABLES, XugudbIdentifierProcessor.INSTANCE.escapeString(schemaName));
try (PreparedStatement preparedStatement = connection.prepareStatement(sql); ResultSet resultSet = preparedStatement.executeQuery()) {
while (resultSet.next()) {
String tableName = resultSet.getString("TABLE_NAME");
Expand All @@ -71,7 +72,7 @@ private void exportTable(Connection connection, String tableName, String schemaN
(SELECT dbms_metadata.get_ddl('TABLE', '%s', '%s') FROM dual) AS ddl
FROM dual;
""";
try (PreparedStatement statement = connection.prepareStatement(String.format(sql, tableName, tableName, schemaName)); ResultSet resultSet = statement.executeQuery()) {
try (PreparedStatement statement = connection.prepareStatement(String.format(sql, XugudbIdentifierProcessor.INSTANCE.escapeString(tableName), XugudbIdentifierProcessor.INSTANCE.escapeString(tableName), XugudbIdentifierProcessor.INSTANCE.escapeString(schemaName))); ResultSet resultSet = statement.executeQuery()) {
String formatSchemaName = format(schemaName);
String formatTableName = format(tableName);
if (resultSet.next()) {
Expand All @@ -82,7 +83,7 @@ private void exportTable(Connection connection, String tableName, String schemaN
String comment = resultSet.getString("comments");
if (StringUtils.isNotBlank(comment)) {
sqlBuilder.append(SQL_COMMENT_TABLE).append(formatSchemaName).append(".").append(formatTableName)
.append(" IS ").append("'").append(comment).append("';");
.append(" IS ").append("'").append(XugudbIdentifierProcessor.INSTANCE.escapeString(comment)).append("';");
}
asyncContext.write(sqlBuilder.toString());
exportTableColumnComment(connection, schemaName, tableName, asyncContext);
Expand All @@ -95,14 +96,14 @@ private void exportTable(Connection connection, String tableName, String schemaN

private void exportTableColumnComment(Connection connection, String schemaName, String tableName, AsyncContext asyncContext) throws SQLException {
String sql = String.format(SQL_SELECT_COLNAME_COMMENT_SYS_SYSCOLUMNCOMMENTS +
"where SCHNAME = '%s' and TVNAME = '%s'and TABLE_TYPE = 'TABLE';", schemaName, tableName);
"where SCHNAME = '%s' and TVNAME = '%s'and TABLE_TYPE = 'TABLE';", XugudbIdentifierProcessor.INSTANCE.escapeString(schemaName), XugudbIdentifierProcessor.INSTANCE.escapeString(tableName));
try (PreparedStatement preparedStatement = connection.prepareStatement(sql); ResultSet resultSet = preparedStatement.executeQuery()) {
while (resultSet.next()) {
String columnName = resultSet.getString("COLNAME");
String comment = resultSet.getString("COMMENT$");
StringBuilder sqlBuilder = new StringBuilder();
sqlBuilder.append(SQL_COMMENT_COLUMN).append(format(schemaName)).append(".").append(format(tableName))
.append(".").append(format(columnName)).append(" IS ").append("'").append(comment).append("';").append("\n");
.append(".").append(format(columnName)).append(" IS ").append("'").append(XugudbIdentifierProcessor.INSTANCE.escapeString(comment)).append("';").append("\n");
asyncContext.write(sqlBuilder.toString());
}
}
Expand All @@ -119,7 +120,7 @@ private void exportViews(Connection connection, String schemaName, AsyncContext
}

private void exportView(Connection connection, String viewName, String schemaName, AsyncContext asyncContext) throws SQLException {
String sql = String.format(SQL_SELECT_DBMS_METADATA_GET_DDL, viewName, schemaName);
String sql = String.format(SQL_SELECT_DBMS_METADATA_GET_DDL, XugudbIdentifierProcessor.INSTANCE.escapeString(viewName), XugudbIdentifierProcessor.INSTANCE.escapeString(schemaName));
try (PreparedStatement preparedStatement = connection.prepareStatement(sql); ResultSet resultSet = preparedStatement.executeQuery()) {
if (resultSet.next()) {
StringBuilder sqlBuilder = new StringBuilder();
Expand All @@ -139,7 +140,7 @@ private void exportProcedures(Connection connection, String schemaName, AsyncCon
}

private void exportProcedure(Connection connection, String schemaName, String procedureName, AsyncContext asyncContext) throws SQLException {
String sql = String.format(ROUTINES_SQL, "PROC", schemaName, procedureName);
String sql = String.format(ROUTINES_SQL, "PROC", XugudbIdentifierProcessor.INSTANCE.escapeString(schemaName), XugudbIdentifierProcessor.INSTANCE.escapeString(procedureName));
try (PreparedStatement statement = connection.prepareStatement(sql); ResultSet resultSet = statement.executeQuery()) {
if (resultSet.next()) {
StringBuilder sqlBuilder = new StringBuilder();
Expand All @@ -150,7 +151,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(TRIGGER_SQL_LIST, schemaName);
String sql = String.format(TRIGGER_SQL_LIST, XugudbIdentifierProcessor.INSTANCE.escapeString(schemaName));
try (PreparedStatement preparedStatement = connection.prepareStatement(sql); ResultSet resultSet = preparedStatement.executeQuery()) {
while (resultSet.next()) {
String triggerName = resultSet.getString("TRIGGER_NAME");
Expand All @@ -160,7 +161,7 @@ private void exportTriggers(Connection connection, String schemaName, AsyncConte
}

private void exportTrigger(Connection connection, String schemaName, String triggerName, AsyncContext asyncContext) throws SQLException {
String sql = String.format(TRIGGER_SQL, schemaName, triggerName);
String sql = String.format(TRIGGER_SQL, XugudbIdentifierProcessor.INSTANCE.escapeString(schemaName), XugudbIdentifierProcessor.INSTANCE.escapeString(triggerName));
try (PreparedStatement preparedStatement = connection.prepareStatement(sql); ResultSet resultSet = preparedStatement.executeQuery()) {
if (resultSet.next()) {
StringBuilder sqlBuilder = new StringBuilder();
Expand All @@ -179,14 +180,14 @@ public void connectDatabase(Connection connection, String database) {
}
String schemaName = connectInfo.getSchemaName();
try {
DefaultSQLExecutor.getInstance().execute(connection, String.format(SQL_SET_SCHEMA, schemaName));
DefaultSQLExecutor.getInstance().execute(connection, String.format(SQL_SET_SCHEMA, format(schemaName)));
} catch (SQLException e) {
throw new RuntimeException(e);
}
}

@Override
public String dropTable(Connection connection, String databaseName, String schemaName, String tableName) {
return String.format(SQL_DROP_TABLE_EXISTS, tableName);
return String.format(SQL_DROP_TABLE_EXISTS, format(tableName));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,9 @@
import ai.chat2db.plugin.xugudb.enums.type.XUGUDBColumnTypeEnum;
import ai.chat2db.plugin.xugudb.enums.type.XUGUDBDefaultValueEnum;
import ai.chat2db.plugin.xugudb.enums.type.XUGUDBIndexTypeEnum;
import ai.chat2db.plugin.xugudb.identifier.XugudbIdentifierProcessor;
import ai.chat2db.spi.IDbMetaData;
import ai.chat2db.spi.ISQLIdentifierProcessor;
import ai.chat2db.spi.ISqlBuilder;
import ai.chat2db.spi.DefaultMetaService;
import ai.chat2db.community.domain.api.model.account.*;
Expand Down Expand Up @@ -37,6 +39,11 @@

public class XUGUDBMetaData extends DefaultMetaService implements IDbMetaData {

@Override
public ISQLIdentifierProcessor getSQLIdentifierProcessor() {
return XugudbIdentifierProcessor.INSTANCE;
}

@Override
public List<Database> databases(Connection connection) {
List<Database> databases = DefaultSQLExecutor.getInstance().databases(connection);
Expand All @@ -45,7 +52,7 @@ public List<Database> databases(Connection connection) {

@Override
public List<Schema> schemas(Connection connection, String databaseName) {
String sql = "select s.schema_name, db.db_name from all_schemas s left join all_databases db on db.db_id = s.db_id where db.db_name = '" + databaseName + "'";
String sql = "select s.schema_name, db.db_name from all_schemas s left join all_databases db on db.db_id = s.db_id where db.db_name = '" + getSQLIdentifierProcessor().escapeString(databaseName) + "'";
List<Schema> schemas = DefaultSQLExecutor.getInstance().execute(connection,
sql, resultSet -> {
List<Schema> databases = new ArrayList<>();
Expand All @@ -63,7 +70,7 @@ public List<Schema> schemas(Connection connection, String databaseName) {
}

private String format(String tableName) {
return "\"" + tableName + "\"";
return XugudbIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableName);
}

@Override
Expand All @@ -74,7 +81,7 @@ public String tableDDL(Connection connection, String databaseName, String schema
FROM dual;
""";
StringBuilder ddlBuilder = new StringBuilder();
String tableDDLSql = String.format(sql, schemaName, tableName);
String tableDDLSql = String.format(sql, getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(tableName));
DefaultSQLExecutor.getInstance().execute(connection, tableDDLSql, resultSet -> {
if (resultSet.next()) {
String ddl = resultSet.getString("ddl");
Expand All @@ -90,7 +97,7 @@ public String tableDDL(Connection connection, String databaseName, String schema
@Override
public List<Function> functions(Connection connection, String databaseName, String schemaName) {
List<Function> functions = new ArrayList<>();
String sql = String.format(FUNCTIONS_SQL, databaseName, schemaName);
String sql = String.format(FUNCTIONS_SQL, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
Function function = new Function();
Expand All @@ -109,7 +116,7 @@ public List<Function> functions(Connection connection, String databaseName, Stri
public Function function(Connection connection, @NotEmpty String databaseName, String schemaName,
String functionName) {

String sql = String.format(ROUTINES_SQL, databaseName, schemaName, functionName);
String sql = String.format(ROUTINES_SQL, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(functionName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
StringBuilder sb = new StringBuilder();
while (resultSet.next()) {
Expand All @@ -131,7 +138,7 @@ public Function function(Connection connection, @NotEmpty String databaseName, S
@Override
public Procedure procedure(Connection connection, @NotEmpty String databaseName, String schemaName,
String procedureName) {
String sql = String.format(PROCEDURE_SQL, databaseName, schemaName, procedureName);
String sql = String.format(PROCEDURE_SQL, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(procedureName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
StringBuilder sb = new StringBuilder();
while (resultSet.next()) {
Expand All @@ -152,7 +159,7 @@ public Procedure procedure(Connection connection, @NotEmpty String databaseName,
@Override
public List<Procedure> procedures(Connection connection, String databaseName, String schemaName) {
List<Procedure> procedures = new ArrayList<>();
String sql = String.format(PROCEDURES_SQL, databaseName, schemaName);
String sql = String.format(PROCEDURES_SQL, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
Procedure procedure = new Procedure();
Expand All @@ -173,7 +180,7 @@ public List<Procedure> procedures(Connection connection, String databaseName, St
@Override
public List<Trigger> triggers(Connection connection, String databaseName, String schemaName) {
List<Trigger> triggers = new ArrayList<>();
String sql = String.format(TRIGGER_SQL_LIST, databaseName, schemaName);
String sql = String.format(TRIGGER_SQL_LIST, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
Trigger trigger = new Trigger();
Expand All @@ -190,7 +197,7 @@ public List<Trigger> triggers(Connection connection, String databaseName, String
public Trigger trigger(Connection connection, @NotEmpty String databaseName, String schemaName,
String triggerName) {

String sql = String.format(TRIGGER_SQL, databaseName, schemaName, triggerName);
String sql = String.format(TRIGGER_SQL, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(triggerName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Trigger trigger = new Trigger();
trigger.setDatabaseName(databaseName);
Expand All @@ -207,7 +214,7 @@ public Trigger trigger(Connection connection, @NotEmpty String databaseName, Str

@Override
public List<Table> views(Connection connection, String databaseName, String schemaName) {
String sql = String.format(VIEW_SQL_LIST, databaseName, schemaName);
String sql = String.format(VIEW_SQL_LIST, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName));
List<Table> tables = new ArrayList<>();
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Table table = new Table();
Expand All @@ -227,7 +234,7 @@ public List<Table> views(Connection connection, String databaseName, String sche

@Override
public Table view(Connection connection, String databaseName, String schemaName, String viewName) {
String sql = String.format(VIEW_SQL, databaseName, schemaName, viewName);
String sql = String.format(VIEW_SQL, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(viewName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Table table = new Table();
table.setDatabaseName(databaseName);
Expand All @@ -244,7 +251,7 @@ public Table view(Connection connection, String databaseName, String schemaName,

@Override
public List<TableIndex> indexes(Connection connection, String databaseName, String schemaName, String tableName) {
String sql = String.format(INDEX_SQL, schemaName, tableName);
String sql = String.format(INDEX_SQL, getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(tableName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
LinkedHashMap<String, TableIndex> map = new LinkedHashMap();
while (resultSet.next()) {
Expand Down Expand Up @@ -294,7 +301,7 @@ private List<TableIndexColumn> getTableIndexColumn(ResultSet resultSet) throws S

@Override
public List<TableColumn> columns(Connection connection, String databaseName, String schemaName, String tableName) {
String sql = String.format(SELECT_TABLE_COLUMNS, databaseName, schemaName, tableName);
String sql = String.format(SELECT_TABLE_COLUMNS, getSQLIdentifierProcessor().escapeString(databaseName), getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(tableName));
List<TableColumn> tableColumns = new ArrayList<>();
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
Expand Down Expand Up @@ -342,9 +349,9 @@ public TableMeta getTableMeta(String databaseName, String schemaName, String tab
@Override
public String getMetaDataName(String... names) {
if (Arrays.stream(names).count() > 1) {
return Arrays.stream(names).skip(1).filter(name -> StringUtils.isNotBlank(name)).map(name -> "\"" + name + "\"").collect(Collectors.joining("."));
return Arrays.stream(names).skip(1).filter(name -> StringUtils.isNotBlank(name)).map(getSQLIdentifierProcessor()::quoteIdentifier).collect(Collectors.joining("."));
}
return Arrays.stream(names).filter(name -> StringUtils.isNotBlank(name)).map(name -> "\"" + name + "\"").collect(Collectors.joining("."));
return Arrays.stream(names).filter(name -> StringUtils.isNotBlank(name)).map(getSQLIdentifierProcessor()::quoteIdentifier).collect(Collectors.joining("."));
}


Expand Down
Loading
Loading