diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManager.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManager.java index dbc38953bc..84bebcd9d2 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManager.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManager.java @@ -2,13 +2,13 @@ import ai.chat2db.spi.IDbManager; import ai.chat2db.plugin.postgresql.builder.PostgreSQLSqlBuilder; +import ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor; import ai.chat2db.spi.DefaultDBManager; import ai.chat2db.community.domain.api.model.async.AsyncContext; import ai.chat2db.spi.sql.Chat2DBContext; import ai.chat2db.spi.model.datasource.ConnectInfo; import ai.chat2db.spi.model.request.TableMetadataRequest; 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.lang3.StringUtils; @@ -57,7 +57,7 @@ private void exportSequences(Connection connection, String schemaName, AsyncCont if (StringUtils.isBlank(sequenceName)) { continue; } - String quotedSequenceName = SqlUtils.quoteObjectName(sequenceName); + String quotedSequenceName = qualifiedTableName(schemaName, sequenceName, false); sqlBuilder.append(SQL_DROP_SEQUENCE_EXISTS).append(quotedSequenceName).append(";\n"); sqlBuilder.append(SQL_CREATE_SEQUENCE).append(quotedSequenceName).append("\n") .append(" START WITH ").append(startValue).append("\n") @@ -79,7 +79,9 @@ private void exportTypes(Connection connection, String schemaName, AsyncContext StringBuilder typeBuilder = new StringBuilder(); DefaultSQLExecutor.getInstance().preExecute(connection, ENUM_TYPE_DDL_SQL, new String[]{schemaName}, resultSet -> { while (resultSet.next()) { - typeBuilder.append(SQL_DROP_TYPE_EXISTS).append(SqlUtils.quoteObjectName(resultSet.getString("type_name"))).append(";\n"); + typeBuilder.append(SQL_DROP_TYPE_EXISTS) + .append(qualifiedTableName(schemaName, resultSet.getString("type_name"), false)) + .append(";\n"); typeBuilder.append(resultSet.getString("ddl")).append("\n"); asyncContext.write(typeBuilder.toString()); } @@ -87,7 +89,7 @@ private void exportTypes(Connection connection, String schemaName, AsyncContext typeBuilder.setLength(0); DefaultSQLExecutor.getInstance().preExecute(connection, UDT_SQL, new String[]{schemaName}, resultSet -> { while (resultSet.next()) { - String typeName = SqlUtils.quoteObjectName(resultSet.getString("type_name")); + String typeName = qualifiedTableName(schemaName, resultSet.getString("type_name"), false); typeBuilder.append(SQL_DROP_TYPE_EXISTS).append(typeName).append(";\n"); typeBuilder.append(resultSet.getString("create_type_statement")).append("\n"); asyncContext.write(typeBuilder.toString()); @@ -109,7 +111,8 @@ public void exportTable(Connection connection, String databaseName, String schem String tableDDL = Chat2DBContext.getDbMetaData().tableDDL(connection, new TableMetadataRequest(databaseName, schemaName, tableName)); StringBuilder sqlBuilder = new StringBuilder(); - sqlBuilder.append("\n").append(SQL_DROP_TABLE_EXISTS).append(SqlUtils.quoteObjectName(tableName)).append(";").append("\n") + sqlBuilder.append("\n").append(SQL_DROP_TABLE_EXISTS) + .append(qualifiedTableName(schemaName, tableName, false)).append(";").append("\n") .append(tableDDL).append("\n"); asyncContext.write(sqlBuilder.toString()); if (asyncContext.isContainsData()) { @@ -126,7 +129,7 @@ private void exportViews(Connection connection, String schemaName, AsyncContext StringBuilder sqlBuilder = new StringBuilder(); String viewName = resultSet.getString("table_name"); String viewDefinition = resultSet.getString("view_definition"); - String quotedObjectName = SqlUtils.quoteObjectName(viewName); + String quotedObjectName = qualifiedTableName(schemaName, viewName, false); sqlBuilder.append(SQL_DROP_VIEW_EXISTS).append(quotedObjectName).append(";\n"); sqlBuilder.append(SQL_CREATE_REPLACE_VIEW).append(quotedObjectName).append(" AS ").append(viewDefinition).append("\n"); asyncContext.write(sqlBuilder.toString()); @@ -143,9 +146,9 @@ private void exportRoutines(Connection connection, String schemaName, AsyncConte String routineDefinition = resultSet.getString("function_definition"); String prokind = resultSet.getString("prokind"); if (Objects.equals("f", prokind)) { - sqlBuilder.append(SQL_DROP_FUNCTION_EXISTS).append(schemaName).append(".").append(routineName).append(";\n"); + sqlBuilder.append(SQL_DROP_FUNCTION_EXISTS).append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName)).append(".").append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(routineName)).append(";\n"); } else { - sqlBuilder.append(SQL_DROP_PROCEDURE_EXISTS).append(schemaName).append(".").append(routineName).append(";\n"); + sqlBuilder.append(SQL_DROP_PROCEDURE_EXISTS).append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName)).append(".").append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(routineName)).append(";\n"); } sqlBuilder.append(routineDefinition).append(";\n\n"); asyncContext.write(sqlBuilder.toString()); @@ -180,7 +183,7 @@ public Connection getConnection(ConnectInfo connectInfo) { connectInfo.setSchemaName(null); Connection connection = super.getConnection(connectInfo); if (StringUtils.isNotBlank(schemaName)) { - String sql = String.format(SQL_SET_SEARCH_PATH_USER_PUBLIC, schemaName); + String sql = String.format(SQL_SET_SEARCH_PATH_USER_PUBLIC, PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName)); try { DefaultSQLExecutor.getInstance().execute(connection, sql); } catch (SQLException e) { @@ -207,8 +210,12 @@ public String replaceDatabaseInJdbcUrl(String url, String newDatabase) { @Override public String dropTable(Connection connection, String databaseName, String schemaName, String tableName) { - String sql = "DROP TABLE " + SqlUtils.quoteObjectName(tableName); - return sql; + return "DROP TABLE " + qualifiedTableName(schemaName, tableName, false); + } + + @Override + public String truncateTable(Connection connection, String databaseName, String schemaName, String tableName) { + return "TRUNCATE TABLE " + qualifiedTableName(schemaName, tableName, true); } @Override @@ -227,13 +234,16 @@ void executeDropSql(Connection connection, String sql) { @Override public void copyTable(Connection connection, String databaseName, String schemaName, String tableName, String newTableName, boolean copyData) throws SQLException { - String sql = ""; - if (copyData) { - sql = "CREATE TABLE " + SqlUtils.quoteObjectName(newTableName) + " AS TABLE " + SqlUtils.quoteObjectName(tableName) + " WITH DATA"; - } else { - sql = "CREATE TABLE " + SqlUtils.quoteObjectName(newTableName) + " AS TABLE " + SqlUtils.quoteObjectName(tableName) + " WITH NO DATA"; - } - DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> null); + DefaultSQLExecutor.getInstance().execute(connection, + buildCopyTableSql(schemaName, tableName, newTableName, copyData), resultSet -> null); + } + + static String buildCopyTableSql(String schemaName, String tableName, String newTableName, + boolean copyData) { + String source = qualifiedTableName(schemaName, tableName, true); + String target = qualifiedTableName(schemaName, newTableName, true); + return "CREATE TABLE " + target + " AS TABLE " + source + + (copyData ? " WITH DATA" : " WITH NO DATA"); } @Override @@ -244,7 +254,25 @@ 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 " + qualifiedTableName(schemaName, viewName, false); DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> null); } + + private static String qualifiedTableName(String schemaName, String tableName, + boolean normalizeQuotedTable) { + String normalizedTable = normalizeQuotedTable ? normalizeQuotedIdentifier(tableName) : tableName; + String quotedTable = PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(normalizedTable); + if (StringUtils.isBlank(schemaName)) { + return quotedTable; + } + return PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName) + + "." + quotedTable; + } + + private static String normalizeQuotedIdentifier(String identifier) { + if (PostgreSQLIdentifierProcessor.INSTANCE.isQuoteIdentifier(identifier)) { + return PostgreSQLIdentifierProcessor.INSTANCE.removeIdentifierQuote(identifier); + } + return identifier; + } } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLMetaData.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLMetaData.java index 2dde2e6b98..e1612940b8 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLMetaData.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSQLMetaData.java @@ -27,7 +27,6 @@ import ai.chat2db.spi.sql.Chat2DBContext; import ai.chat2db.spi.DefaultSQLExecutor; import ai.chat2db.spi.util.SortUtils; -import ai.chat2db.spi.util.SqlUtils; import com.google.common.collect.Lists; import jakarta.validation.constraints.NotEmpty; import lombok.extern.slf4j.Slf4j; @@ -54,8 +53,6 @@ public class PostgreSQLMetaData extends DefaultMetaService implements IDbMetaDat - public static final ISQLIdentifierProcessor POSTGRE_SQL_IDENTIFIER_PROCESSOR = new PostgreSQLIdentifierProcessor(); - @Override public List databases(Connection connection) { List list = DefaultSQLExecutor.getInstance().execute(connection, SQL_SELECT_DATNAME_PG_DATABASE, resultSet -> { @@ -93,7 +90,7 @@ public List tables(Connection connection, String databaseName, String sch @Override public List triggers(Connection connection, String databaseName, String schemaName) { List triggers = new ArrayList<>(); - String sql = String.format(TRIGGER_SQL_LIST, schemaName); + String sql = String.format(TRIGGER_SQL_LIST, getSQLIdentifierProcessor().escapeString(schemaName)); return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> { while (resultSet.next()) { Trigger trigger = new Trigger(); @@ -108,11 +105,7 @@ public List triggers(Connection connection, String databaseName, String protected String format(String objectName) { - if (StringUtils.isBlank(objectName)) { - return objectName; - } else { - return SqlUtils.quoteObjectName(objectName); - } + return PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(objectName); } @Override @@ -140,7 +133,7 @@ public String tableDDL(Connection connection, String databaseName, String schema StringBuilder ddlBuilder = new StringBuilder(200); - String formatTableName = format(tableName); + String formatTableName = quoteQualifiedName(schemaName, tableName); ddlBuilder.append(SQL_CREATE_TABLE).append(formatTableName); String options = DefaultSQLExecutor.getInstance().preExecute(connection, TABLE_OPTION_SQL, new String[]{schemaName, tableName}, resultSet -> { if (resultSet.next()) { @@ -156,7 +149,7 @@ public String tableDDL(Connection connection, String databaseName, String schema StringBuilder tableSpaceBuilder = new StringBuilder(); tableSpaceBuilder.append(" tablespace "); if (resultSet.next()) { - tableSpaceBuilder.append(resultSet.getString("tablespace")); + tableSpaceBuilder.append(format(resultSet.getString("tablespace"))); } else { tableSpaceBuilder.append("pg_default"); } @@ -178,9 +171,9 @@ public String tableDDL(Connection connection, String databaseName, String schema constraintsBuilder.append(",\n"); } constraintsBuilder.append("\t").append(" constraint ") - .append(constraintName) + .append(format(constraintName)) .append(" ") - .append(constraintDefinition.toLowerCase()); + .append(constraintDefinition); } } if (!constraintsBuilder.isEmpty()) { @@ -197,7 +190,9 @@ public String tableDDL(Connection connection, String databaseName, String schema String partitionDefinition = resultSet.getString("PARTITION_DEFINITION"); boolean isParentTable = resultSet.getBoolean("is_parent_table"); if (StringUtils.isNotBlank(parentTableName) && StringUtils.isNotBlank(partitionDefinition)) { - ddlBuilder.append("\n").append(" partition of ").append(SqlUtils.quoteObjectName(parentTableName)).append("\n"); + ddlBuilder.append("\n").append(" partition of ") + .append(quoteQualifiedName(resultSet.getString("parent_schema"), parentTableName)) + .append("\n"); if (!constraintsBuilder.isEmpty()) { ddlBuilder.append("(\n") .append(constraintsBuilder) @@ -219,9 +214,9 @@ public String tableDDL(Connection connection, String databaseName, String schema String table_name = resultSet.getString("TABLE_NAME"); if (StringUtils.isNotBlank(owner) && StringUtils.isNotBlank(table_name)) { tableOwnerBuilder.append(SQL_ALTER_TABLE) - .append(format(table_name)) + .append(quoteQualifiedName(schemaName, table_name)) .append(" owner to ") - .append(owner) + .append(format(owner)) .append(";").append("\n"); } } @@ -236,11 +231,11 @@ public String tableDDL(Connection connection, String databaseName, String schema String privilegeType = resultSet.getString("PRIVILEGE_TYPE"); if (StringUtils.isNotBlank(privilegeType)) { tablePrivilegeBuilder.append(SQL_GRANT) - .append(privilegeType.toLowerCase()) + .append(PostgreSqlGuards.requirePrivilege(privilegeType)) .append(SQL_ON) .append(formatTableName) .append(" to ") - .append(grantee) + .append(quoteRoleName(grantee)) .append(";").append("\n"); } } @@ -459,7 +454,7 @@ public String tableDDL(Connection connection, String databaseName, String schema boolean isPartitioned = false; if (resultSet.next()) { ddlBuilder.append(" partition by ") - .append(resultSet.getString("partition_key").toLowerCase()) + .append(resultSet.getString("partition_key")) .append(";"); isPartitioned = true; ddlBuilder.append("\n"); @@ -474,18 +469,22 @@ public String tableDDL(Connection connection, String databaseName, String schema String parentTableName = resultSet.getString("PARENT_TABLE"); String partitionDefinition = resultSet.getString("PARTITION_DEFINITION"); if (StringUtils.isNotBlank(parentTableName) && StringUtils.isNotBlank(partitionDefinition)) { - ddlBuilder.append("\n").append(SQL_CREATE_TABLE).append(format(subName)).append("\n") - .append("partition of ").append(parentTableName).append("\n") - .append(partitionDefinition.toLowerCase()).append(";\n"); + // sub_name and PARENT_TABLE are quote_ident() output: already safely quoted + ddlBuilder.append("\n").append(SQL_CREATE_TABLE) + .append(resultSet.getString("schema_name")).append(".").append(subName).append("\n") + .append("partition of ").append(format(schemaName)).append(".") + .append(parentTableName).append("\n") + .append(partitionDefinition).append(";\n"); } } }); } else if (childTableInfo.size() >= 2) { + String parentSchemaName = childTableInfo.get(0); String parentTableName = childTableInfo.get(1); ddlBuilder.append(" ").append(" inherits ") .append("(") - .append(format(parentTableName)) + .append(quoteQualifiedName(parentSchemaName, parentTableName)) .append(")").append("\n"); if (StringUtils.isNotBlank(options)) { ddlBuilder.append(" ").append(options).append("\n"); @@ -519,10 +518,12 @@ public String tableDDL(Connection connection, String databaseName, String schema } String objectType = resultSet.getString("object_type"); + String quoteSchemaName = resultSet.getString("schema_name"); String quoteTableName = resultSet.getString("table_name"); String columnName = resultSet.getString("column_name"); - ddlBuilder.append(SQL_COMMENT).append(objectType.toLowerCase()).append(" ").append(quoteTableName); + ddlBuilder.append(SQL_COMMENT).append(objectType.toLowerCase(Locale.ROOT)).append(" ") + .append(quoteSchemaName).append(".").append(quoteTableName); if (StringUtils.isNotBlank(columnName)) { ddlBuilder.append(".").append(columnName); @@ -537,10 +538,11 @@ public String tableDDL(Connection connection, String databaseName, String schema DefaultSQLExecutor.getInstance().preExecute(connection, TABLE_INDEX_COMMENT_SQL, new String[]{schemaName, tableName}, resultSet -> { while (resultSet.next()) { + String schema_name = resultSet.getString("schema_name"); String index_name = resultSet.getString("index_name"); String index_comment = resultSet.getString("index_comment"); - ddlBuilder.append(SQL_COMMENT_INDEX).append(index_name) + ddlBuilder.append(SQL_COMMENT_INDEX).append(schema_name).append(".").append(index_name) .append(" is ").append(index_comment).append(";\n"); } @@ -576,7 +578,7 @@ public Function function(Connection connection, @NotEmpty String databaseName, S @Override public Table view(Connection connection, String databaseName, String schemaName, String viewName) { - String sql = String.format(VIEW_SQL, schemaName, viewName); + String sql = String.format(VIEW_SQL, getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(viewName)); return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> { Table table = new Table(); table.setDatabaseName(databaseName); @@ -593,7 +595,7 @@ public Table view(Connection connection, String databaseName, String schemaName, public Trigger trigger(Connection connection, @NotEmpty String databaseName, String schemaName, String triggerName) { - String sql = String.format(TRIGGER_SQL, schemaName, triggerName); + String sql = String.format(TRIGGER_SQL, getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(triggerName)); return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> { Trigger trigger = new Trigger(); trigger.setDatabaseName(databaseName); @@ -625,7 +627,7 @@ public Procedure procedure(Connection connection, @NotEmpty String databaseName, @Override public List indexes(Connection connection, String databaseName, String schemaName, String tableName) { - String constraintSql = String.format(SELECT_KEY_INDEX, schemaName, tableName); + String constraintSql = String.format(SELECT_KEY_INDEX, getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(tableName)); Map constraintMap = new HashMap(); LinkedHashMap foreignMap = new LinkedHashMap(); DefaultSQLExecutor.getInstance().execute(connection, constraintSql, resultSet -> { @@ -655,7 +657,7 @@ public List indexes(Connection connection, String databaseName, Stri return null; }); - String sql = String.format(SELECT_TABLE_INDEX, schemaName, tableName); + String sql = String.format(SELECT_TABLE_INDEX, getSQLIdentifierProcessor().escapeString(schemaName), getSQLIdentifierProcessor().escapeString(tableName)); return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> { LinkedHashMap map = new LinkedHashMap(foreignMap); @@ -709,7 +711,7 @@ public List columns(Connection connection, String databaseName, Str if (StringUtils.equalsIgnoreCase(v.getColumnType(), "bpchar")) { v.setColumnType(PostgreSQLColumnTypeEnum.CHAR.getColumnType().getTypeName().toUpperCase()); } else { - v.setColumnType(v.getColumnType().toUpperCase()); + v.setColumnType(v.getColumnType().toUpperCase(Locale.ROOT)); } }); return columnList; @@ -747,12 +749,23 @@ public IValueProcessor getValueProcessor() { @Override public ISQLIdentifierProcessor getSQLIdentifierProcessor() { - return POSTGRE_SQL_IDENTIFIER_PROCESSOR; + return PostgreSQLIdentifierProcessor.INSTANCE; } @Override public String getMetaDataName(String... names) { - return Arrays.stream(names).filter(name -> StringUtils.isNotBlank(name)).map(name -> "\"" + name + "\"").collect(Collectors.joining(".")); + return new PostgreSQLSqlBuilder().quoteQualifiedIdentifier(names); + } + + private String quoteQualifiedName(String schemaName, String objectName) { + if (StringUtils.isBlank(schemaName)) { + return format(objectName); + } + return format(schemaName) + "." + format(objectName); + } + + private String quoteRoleName(String roleName) { + return "PUBLIC".equalsIgnoreCase(roleName) ? "PUBLIC" : format(roleName); } @Override @@ -786,10 +799,7 @@ public ModifyViewConfiguration viewMeta(String databaseName, String schemaName) String sql = "select * from table_name"; StringBuilder sqlBuilder = new StringBuilder(100); sqlBuilder.append(SQL_CREATE).append("view "); - if (StringUtils.isNotBlank(schemaName)) { - sqlBuilder.append("\"").append(schemaName).append("\"").append("."); - } - sqlBuilder.append("\"").append("undefined").append("\""); + sqlBuilder.append(quoteQualifiedName(schemaName, "undefined")); sqlBuilder.append(" AS \n").append(sql).append(";"); configuration.setPreviewSql(sqlBuilder.toString()); configuration.setSql(sql); diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSqlGuards.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSqlGuards.java new file mode 100644 index 0000000000..a3e9cc8959 --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/PostgreSqlGuards.java @@ -0,0 +1,326 @@ +package ai.chat2db.plugin.postgresql; + +import org.apache.commons.lang3.StringUtils; + +import java.util.ArrayDeque; +import java.util.Deque; +import java.util.Locale; +import java.util.Set; +import java.util.regex.Pattern; + +/** + * Validation helpers for non-escapable SQL positions in PostgreSQL DDL/DML generation + * (strict name tokens, raw DEFAULT expressions, bit/hex literal content, enum options). + * Escaping itself lives in + * {@link ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor}. + */ +public final class PostgreSqlGuards { + + private static final Pattern PG_NAME_PATTERN = Pattern.compile("^[A-Za-z_][A-Za-z0-9_]*$"); + private static final Pattern BIT_LITERAL_PATTERN = Pattern.compile("^[01]*$"); + private static final Pattern HEX_LITERAL_PATTERN = Pattern.compile("^[0-9a-fA-F]*$"); + private static final Set DEFAULT_BREAKOUT_KEYWORDS = Set.of( + "CHECK", "CONSTRAINT", "DEFAULT", "GENERATED", "PRIMARY", "REFERENCES", "UNIQUE", + "DROP", "ALTER", "CREATE", "GRANT", "REVOKE", "TRUNCATE"); + private static final Set TYPE_BREAKOUT_KEYWORDS = Set.of( + "CHECK", "COLLATE", "CONSTRAINT", "DEFAULT", "GENERATED", "NOT", "NULL", + "PRIMARY", "REFERENCES", "UNIQUE"); + private static final Set PRIVILEGES = Set.of( + "SELECT", "INSERT", "UPDATE", "DELETE", "TRUNCATE", "REFERENCES", "TRIGGER", "MAINTAIN"); + private static final Set VIEW_STORAGE_CLAUSES = Set.of( + "TEMP", "LOCAL TEMP", "GLOBAL TEMP", "UNLOGGED"); + + private PostgreSqlGuards() { + } + + /** + * Validate a strict PostgreSQL name token (index method / role / keyword-style positions where + * escaping is impossible by design). + */ + public static String requirePgName(String value, String what) { + if (value == null || !PG_NAME_PATTERN.matcher(value).matches()) { + throw new IllegalArgumentException("Invalid PostgreSQL " + what + ": " + value); + } + return value; + } + + /** + * Validates one quote-aware PostgreSQL DEFAULT expression. Function calls, casts, sequence + * expressions, array subscripts, dollar-quoted strings, and nested parentheses are preserved, + * while statement terminators, comments, unbalanced delimiters, top-level commas, and column + * constraint suffixes are rejected. + */ + public static String requireDefaultExpression(String value) { + if (StringUtils.isBlank(value)) { + throw invalid("default value", value); + } + String expression = value.trim(); + scanDefaultExpression(expression); + return expression; + } + + /** + * Validates a PostgreSQL type expression used when the type is not an exact enum match. + * Supports qualified and quoted user-defined types, type modifiers, multi-word built-ins, + * and array suffixes without allowing the value to append a column constraint. + */ + public static String requireColumnTypeExpression(String value) { + if (StringUtils.isBlank(value)) { + throw invalid("column type", value); + } + String expression = value.trim(); + Deque delimiters = new ArrayDeque<>(); + boolean sawName = false; + for (int i = 0; i < expression.length(); i++) { + char c = expression.charAt(i); + if (c == '"') { + int quoteEnd = scanQuoted(expression, i, '"', false, "column type"); + sawName = true; + i = quoteEnd; + continue; + } + if (startsWith(expression, i, "--") || startsWith(expression, i, "/*") + || startsWith(expression, i, "*/") || c == ';' || Character.isISOControl(c)) { + throw invalid("column type", value); + } + if (Character.isLetter(c) || c == '_') { + int tokenEnd = scanWord(expression, i); + String token = expression.substring(i, tokenEnd).toUpperCase(Locale.ROOT); + if (TYPE_BREAKOUT_KEYWORDS.contains(token)) { + throw invalid("column type", value); + } + sawName = true; + i = tokenEnd - 1; + continue; + } + if (Character.isDigit(c) || Character.isWhitespace(c) || c == '.' || c == '$') { + continue; + } + if (c == '(') { + delimiters.push(c); + continue; + } + if (c == ')') { + if (delimiters.isEmpty() || delimiters.pop() != '(') { + throw invalid("column type", value); + } + continue; + } + if (c == ',') { + if (delimiters.isEmpty()) { + throw invalid("column type", value); + } + continue; + } + if (c == '[' && i + 1 < expression.length() && expression.charAt(i + 1) == ']') { + i++; + continue; + } + throw invalid("column type", value); + } + if (!sawName || !delimiters.isEmpty()) { + throw invalid("column type", value); + } + return expression; + } + + /** + * Returns whether a temporal default should be treated as SQL syntax instead of a text value. + */ + public static boolean isTemporalExpression(String value) { + if (StringUtils.isBlank(value)) { + return false; + } + String expression = value.trim(); + String upper = expression.toUpperCase(Locale.ROOT); + return expression.indexOf('(') >= 0 || expression.contains("::") + || upper.startsWith("CURRENT_") || upper.startsWith("CURRENT ") + || upper.startsWith("LOCALTIME") || upper.startsWith("LOCALTIMESTAMP") + || upper.startsWith("DATE '") || upper.startsWith("TIME '") + || upper.startsWith("TIMESTAMP '") || upper.startsWith("INTERVAL '"); + } + + public static boolean isFunctionOrCastExpression(String value) { + if (StringUtils.isBlank(value)) { + return false; + } + String expression = value.trim(); + int openParenthesis = expression.indexOf('('); + return expression.contains("::") || openParenthesis > 0; + } + + public static String requirePrivilege(String value) { + String privilege = StringUtils.trimToEmpty(value).toUpperCase(Locale.ROOT); + if (!PRIVILEGES.contains(privilege)) { + throw invalid("privilege", value); + } + return privilege; + } + + public static String requireViewStorageClause(String value) { + String storageClause = StringUtils.normalizeSpace(value).toUpperCase(Locale.ROOT); + if (!VIEW_STORAGE_CLAUSES.contains(storageClause)) { + throw invalid("view storage clause", value); + } + return storageClause; + } + + /** + * Validate content of a B'...' bit literal. + */ + public static String requireBitLiteral(String value) { + if (value == null || !BIT_LITERAL_PATTERN.matcher(value).matches()) { + throw new IllegalArgumentException("Invalid PostgreSQL bit literal: " + value); + } + return value; + } + + /** + * Validate content of a \x... bytea hex literal. + */ + public static String requireHexLiteral(String value) { + if (value == null || !HEX_LITERAL_PATTERN.matcher(value).matches()) { + throw new IllegalArgumentException("Invalid PostgreSQL bytea hex literal: " + value); + } + return value; + } + + /** + * Validate an option that must be one of the given enum constants (e.g. view check option). + * Returns the canonical enum name. + */ + public static > String requireEnumConstant(String value, E[] constants, String what) { + for (E constant : constants) { + if (constant.name().equalsIgnoreCase(StringUtils.trimToEmpty(value))) { + return constant.name(); + } + } + throw new IllegalArgumentException("Invalid PostgreSQL " + what + ": " + value); + } + + private static void scanDefaultExpression(String expression) { + Deque delimiters = new ArrayDeque<>(); + boolean sawContent = false; + for (int i = 0; i < expression.length(); i++) { + char c = expression.charAt(i); + if (c == '\'' || c == '"') { + boolean escapeBackslash = c == '\'' && i > 0 && (expression.charAt(i - 1) == 'E' + || expression.charAt(i - 1) == 'e') + && (i == 1 || !isIdentifierCharacter(expression.charAt(i - 2))); + i = scanQuoted(expression, i, c, escapeBackslash, "default value"); + sawContent = true; + continue; + } + if (c == '$') { + int dollarEnd = scanDollarQuoted(expression, i); + if (dollarEnd >= 0) { + i = dollarEnd; + sawContent = true; + continue; + } + } + if (startsWith(expression, i, "--") || startsWith(expression, i, "/*") + || startsWith(expression, i, "*/") || c == ';' || c == '\n' || c == '\r' + || Character.isISOControl(c)) { + throw invalid("default value", expression); + } + if (c == '(' || c == '[') { + delimiters.push(c); + sawContent = true; + continue; + } + if (c == ')' || c == ']') { + if (delimiters.isEmpty() || !matches(delimiters.pop(), c)) { + throw invalid("default value", expression); + } + sawContent = true; + continue; + } + if (c == ',' && delimiters.isEmpty()) { + throw invalid("default value", expression); + } + if (Character.isLetter(c) || c == '_') { + int tokenEnd = scanWord(expression, i); + if (delimiters.isEmpty()) { + String token = expression.substring(i, tokenEnd).toUpperCase(Locale.ROOT); + if (DEFAULT_BREAKOUT_KEYWORDS.contains(token) + || (("NOT".equals(token) || "NULL".equals(token)) && sawContent)) { + throw invalid("default value", expression); + } + } + sawContent = true; + i = tokenEnd - 1; + continue; + } + if (!Character.isWhitespace(c)) { + sawContent = true; + } + } + if (!sawContent || !delimiters.isEmpty()) { + throw invalid("default value", expression); + } + } + + private static int scanQuoted(String expression, int quoteStart, char quote, + boolean escapeBackslash, String description) { + for (int i = quoteStart + 1; i < expression.length(); i++) { + char c = expression.charAt(i); + if (escapeBackslash && c == '\\') { + if (++i >= expression.length()) { + throw invalid(description, expression); + } + continue; + } + if (c == quote) { + if (i + 1 < expression.length() && expression.charAt(i + 1) == quote) { + i++; + } else { + return i; + } + } + } + throw invalid(description, expression); + } + + private static int scanDollarQuoted(String expression, int start) { + int delimiterEnd = expression.indexOf('$', start + 1); + if (delimiterEnd < 0) { + return -1; + } + for (int i = start + 1; i < delimiterEnd; i++) { + if (!isIdentifierCharacter(expression.charAt(i))) { + return -1; + } + } + String delimiter = expression.substring(start, delimiterEnd + 1); + int contentEnd = expression.indexOf(delimiter, delimiterEnd + 1); + if (contentEnd < 0) { + throw invalid("default value", expression); + } + return contentEnd + delimiter.length() - 1; + } + + private static int scanWord(String expression, int start) { + int current = start + 1; + while (current < expression.length() && isIdentifierCharacter(expression.charAt(current))) { + current++; + } + return current; + } + + private static boolean isIdentifierCharacter(char c) { + return Character.isLetterOrDigit(c) || c == '_' || c == '$'; + } + + private static boolean startsWith(String value, int offset, String candidate) { + return offset + candidate.length() <= value.length() && value.startsWith(candidate, offset); + } + + private static boolean matches(char open, char close) { + return open == '(' && close == ')' || open == '[' && close == ']'; + } + + private static IllegalArgumentException invalid(String description, String value) { + return new IllegalArgumentException("Invalid PostgreSQL " + description + ": " + value); + } +} diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilder.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilder.java index def3656e8c..d3367c43b1 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilder.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilder.java @@ -2,12 +2,15 @@ import ai.chat2db.spi.constant.SQLConstants; -import ai.chat2db.plugin.postgresql.PostgreSQLMetaData; +import ai.chat2db.community.domain.api.enums.plugin.DmlTypeEnum; +import ai.chat2db.plugin.postgresql.PostgreSqlGuards; +import ai.chat2db.plugin.postgresql.enums.PostgreSQLViewCheckOptionEnum; import ai.chat2db.plugin.postgresql.enums.type.PostgreSQLColumnTypeEnum; import ai.chat2db.plugin.postgresql.enums.type.PostgreSQLIndexTypeEnum; -import ai.chat2db.spi.ISQLIdentifierProcessor; +import ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor; import ai.chat2db.spi.DefaultSqlBuilder; import ai.chat2db.spi.model.request.PageLimitRequest; +import ai.chat2db.spi.model.request.UpdateSqlRequest; import ai.chat2db.community.domain.api.model.account.*; import ai.chat2db.community.domain.api.model.async.*; import ai.chat2db.community.domain.api.config.*; @@ -20,6 +23,7 @@ import ai.chat2db.community.domain.api.model.view.*; import ai.chat2db.community.domain.api.config.TableBuilderConfig; import org.apache.commons.collections4.CollectionUtils; +import org.apache.commons.collections4.MapUtils; import org.apache.commons.lang3.BooleanUtils; import org.apache.commons.lang3.StringUtils; @@ -39,6 +43,9 @@ public String quoteIdentifier(String identifier) { @Override public String quoteQualifiedIdentifier(String... identifiers) { + if (identifiers.length == 3) { + return quoteQualifiedIdentifier(identifiers[1], identifiers[2]); + } return Arrays.stream(identifiers) .filter(StringUtils::isNotBlank) .map(PostgreSQLSqlBuilder::quotePostgreSqlIdentifier) @@ -50,6 +57,53 @@ public String quoteAlias(String alias) { return quoteIdentifier(alias); } + @Override + public String buildUpdate(UpdateSqlRequest request) { + StringBuilder script = new StringBuilder(SQLConstants.UPDATE_KEYWORD + SQLConstants.SPACE); + buildTableName(request.getDatabaseName(), request.getSchemaName(), request.getTableName(), script); + script.append(" SET "); + script.append(request.getRow().entrySet().stream() + .map(entry -> PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(entry.getKey()) + + SQLConstants.EQUAL_SQL + entry.getValue()) + .collect(Collectors.joining(SQLConstants.COMMA))); + if (MapUtils.isNotEmpty(request.getPrimaryKeyMap())) { + script.append(" WHERE "); + script.append(request.getPrimaryKeyMap().entrySet().stream() + .map(entry -> PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(entry.getKey()) + + SQLConstants.EQUAL_SQL + entry.getValue()) + .collect(Collectors.joining(SQLConstants.SQL_AND))); + } + return script.toString(); + } + + @Override + public String buildTemplate(Table table, String type) { + if (table == null || CollectionUtils.isEmpty(table.getColumnList()) || StringUtils.isBlank(type)) { + return SQLConstants.EMPTY; + } + String tableName = quoteQualifiedIdentifier(table.getSchemaName(), table.getName()); + List columnNames = table.getColumnList().stream() + .map(column -> PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())) + .toList(); + if (DmlTypeEnum.INSERT.name().equalsIgnoreCase(type)) { + return "INSERT INTO " + tableName + " (" + String.join(SQLConstants.COMMA, columnNames) + + ") VALUES (" + columnNames.stream().map(name -> SQLConstants.SPACE) + .collect(Collectors.joining(SQLConstants.COMMA)) + ")"; + } + if (DmlTypeEnum.UPDATE.name().equalsIgnoreCase(type)) { + return "UPDATE " + tableName + " SET " + columnNames.stream() + .map(name -> name + SQLConstants.EQUAL_SQL + SQLConstants.SPACE) + .collect(Collectors.joining(SQLConstants.COMMA)) + " WHERE "; + } + if (DmlTypeEnum.DELETE.name().equalsIgnoreCase(type)) { + return "DELETE FROM " + tableName + " WHERE "; + } + if (DmlTypeEnum.SELECT.name().equalsIgnoreCase(type)) { + return "SELECT " + String.join(SQLConstants.COMMA, columnNames) + " FROM " + tableName; + } + return SQLConstants.EMPTY; + } + @@ -91,18 +145,18 @@ public String buildCreateTable(Table table, TableBuilderConfig tableBuilderConfi StringBuilder script = new StringBuilder(); script.append(SQL_CREATE_TABLE); if (needFullTableName) { - script.append(SQLConstants.DOUBLE_QUOTE).append(table.getSchemaName()).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.DOT); + script.append(quoteQualifiedIdentifier(table.getSchemaName(), table.getName())); + } else { + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(table.getName())); } - script.append(SQLConstants.DOUBLE_QUOTE).append(table.getName()).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.SPACE_OPEN_PARENTHESIS).append(SQLConstants.SPACE).append(SQLConstants.LINE_SEPARATOR); + script.append(SQLConstants.SPACE_OPEN_PARENTHESIS).append(SQLConstants.SPACE).append(SQLConstants.LINE_SEPARATOR); for (TableColumn column : table.getColumnList()) { if (StringUtils.isBlank(column.getName()) || StringUtils.isBlank(column.getColumnType())) { continue; } - PostgreSQLColumnTypeEnum typeEnum = PostgreSQLColumnTypeEnum.getByType(column.getColumnType()); - if (typeEnum == null) { - continue; - } - script.append(SQLConstants.TAB).append(typeEnum.buildCreateColumnSql(column)).append(SQLConstants.COMMA_LINE_SEPARATOR); + script.append(SQLConstants.TAB) + .append(PostgreSQLColumnTypeEnum.buildCreateColumnSqlSafely(column)) + .append(SQLConstants.COMMA_LINE_SEPARATOR); } Map> tableIndexMap = table.getIndexList().stream() .collect(Collectors.partitioningBy(v -> PostgreSQLIndexTypeEnum.NORMAL.getName().equals(v.getType()))); @@ -137,8 +191,12 @@ public String buildCreateTable(Table table, TableBuilderConfig tableBuilderConfi } if (StringUtils.isNotBlank(table.getComment())) { script.append(SQLConstants.LINE_SEPARATOR); - script.append(SQL_COMMENT_TABLE).append(SQLConstants.SPACE).append(SQLConstants.DOUBLE_QUOTE).append(table.getName()).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE) - .append(table.getComment()).append(SQLConstants.SINGLE_QUOTE_SEMICOLON_LINE_SEPARATOR); + script.append(SQL_COMMENT_TABLE).append(SQLConstants.SPACE) + .append(quoteQualifiedIdentifier( + BooleanUtils.isTrue(needFullTableName) ? table.getSchemaName() : null, + table.getName())) + .append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE) + .append(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(table.getComment())).append(SQLConstants.SINGLE_QUOTE_SEMICOLON_LINE_SEPARATOR); } List tableColumnList = table.getColumnList().stream().filter(v -> StringUtils.isNotBlank(v.getComment())).toList(); for (TableColumn tableColumn : tableColumnList) { @@ -166,17 +224,16 @@ public String buildCreateTable(Table table, TableBuilderConfig tableBuilderConfi public String buildAITableSchema(Table table) { StringBuilder script = new StringBuilder(); script.append(SQL_CREATE_TABLE); - script.append(SQLConstants.DOUBLE_QUOTE).append(table.getSchemaName()).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.DOT); - script.append(SQLConstants.DOUBLE_QUOTE).append(table.getName()).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.SPACE_OPEN_PARENTHESIS).append(SQLConstants.SPACE).append(SQLConstants.LINE_SEPARATOR); + script.append(quoteQualifiedIdentifier(table.getSchemaName(), table.getName())) + .append(SQLConstants.SPACE_OPEN_PARENTHESIS).append(SQLConstants.SPACE) + .append(SQLConstants.LINE_SEPARATOR); for (TableColumn column : table.getColumnList()) { if (StringUtils.isBlank(column.getName()) || StringUtils.isBlank(column.getColumnType())) { continue; } - PostgreSQLColumnTypeEnum typeEnum = PostgreSQLColumnTypeEnum.getByType(column.getColumnType()); - if (typeEnum == null) { - continue; - } - script.append(SQLConstants.TAB).append(typeEnum.buildAICreateColumnSql(column)).append(SQLConstants.COMMA_LINE_SEPARATOR); + script.append(SQLConstants.TAB) + .append(PostgreSQLColumnTypeEnum.buildAICreateColumnSqlSafely(column)) + .append(SQLConstants.COMMA_LINE_SEPARATOR); } if (CollectionUtils.isEmpty(table.getIndexList())) { table.setIndexList(List.of()); @@ -214,8 +271,10 @@ public String buildAITableSchema(Table table) { } if (StringUtils.isNotBlank(table.getComment())) { script.append(SQLConstants.LINE_SEPARATOR); - script.append(SQL_COMMENT_TABLE).append(SQLConstants.SPACE).append(SQLConstants.DOUBLE_QUOTE).append(table.getName()).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE) - .append(table.getComment()).append(SQLConstants.SINGLE_QUOTE_SEMICOLON_LINE_SEPARATOR); + script.append(SQL_COMMENT_TABLE).append(SQLConstants.SPACE) + .append(quoteQualifiedIdentifier(table.getSchemaName(), table.getName())) + .append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE) + .append(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(table.getComment())).append(SQLConstants.SINGLE_QUOTE_SEMICOLON_LINE_SEPARATOR); } List indexList = table.getIndexList().stream().filter(v -> StringUtils.isNotBlank(v.getComment())).toList(); for (TableIndex index : indexList) { @@ -233,35 +292,37 @@ public String buildAITableSchema(Table table) { @Override public String buildAlterTable(Table oldTable, Table newTable) { StringBuilder script = new StringBuilder(); - if (!StringUtils.equalsIgnoreCase(oldTable.getName(), newTable.getName())) { - script.append(SQL_ALTER_TABLE).append(SQLConstants.DOUBLE_QUOTE).append(oldTable.getName()).append(SQLConstants.DOUBLE_QUOTE); - script.append(SQLConstants.TAB).append(SQL_RENAME).append(SQLConstants.DOUBLE_QUOTE).append(newTable.getName()).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.SEMICOLON_LINE_SEPARATOR); + String oldQualifiedName = quoteQualifiedIdentifier(oldTable.getSchemaName(), oldTable.getName()); + String newQualifiedName = quoteQualifiedIdentifier(newTable.getSchemaName(), newTable.getName()); + if (!StringUtils.equals(oldTable.getName(), newTable.getName())) { + script.append(SQL_ALTER_TABLE).append(oldQualifiedName); + script.append(SQLConstants.TAB).append(SQL_RENAME).append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(newTable.getName())).append(SQLConstants.SEMICOLON_LINE_SEPARATOR); } newTable.setIndexList(newTable.getIndexList().stream().filter(v -> StringUtils.isNotBlank(v.getEditStatus())).toList()); List columnNameList = newTable.getColumnList().stream().filter(v -> v.getOldName() != null && !StringUtils.equals(v.getOldName(), v.getName())).toList(); for (TableColumn tableColumn : columnNameList) { - script.append(SQL_ALTER_TABLE).append(SQLConstants.DOUBLE_QUOTE).append(newTable.getName()).append(VALUE_DOUBLE_QUOTE).append(SQL_RENAME_COLUMN) - .append(tableColumn.getOldName()).append(VALUE_DOUBLE_QUOTE_TO_DOUBLE_QUOTE).append(tableColumn.getName()).append(SQLConstants.DOUBLE_QUOTE_SEMICOLON_LINE_SEPARATOR); + script.append(SQL_ALTER_TABLE).append(newQualifiedName).append(VALUE_DOUBLE_QUOTE) + .append("RENAME COLUMN ") + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableColumn.getOldName())) + .append(" TO ") + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableColumn.getName())) + .append(SQLConstants.SEMICOLON_LINE_SEPARATOR); } Map> tableIndexMap = newTable.getIndexList().stream() .collect(Collectors.partitioningBy(v -> PostgreSQLIndexTypeEnum.NORMAL.getName().equals(v.getType()))); StringBuilder scriptModify = new StringBuilder(); Boolean modify = false; - scriptModify.append(SQL_ALTER_TABLE).append(SQLConstants.DOUBLE_QUOTE).append(newTable.getName()).append(VALUE_DOUBLE_QUOTE_2); + scriptModify.append(SQL_ALTER_TABLE).append(newQualifiedName).append(VALUE_DOUBLE_QUOTE_2); List columnList = newTable.getColumnList(); for (TableColumn tableColumn : columnList) { String editStatus = tableColumn.getEditStatus(); if (StringUtils.isBlank(editStatus)) { continue; } - PostgreSQLColumnTypeEnum typeEnum = PostgreSQLColumnTypeEnum.getByType(tableColumn.getColumnType()); - if (typeEnum == null) { - continue; - } - String modifyColumn = typeEnum.buildModifyColumn(tableColumn); + String modifyColumn = PostgreSQLColumnTypeEnum.buildModifyColumnSafely(tableColumn); if (StringUtils.isNotBlank(modifyColumn)) { scriptModify.append(SQLConstants.TAB).append(modifyColumn).append(SQLConstants.COMMA_LINE_SEPARATOR); modify = true; @@ -295,8 +356,8 @@ public String buildAlterTable(Table oldTable, Table newTable) { } if (!StringUtils.equals(oldTable.getComment(), newTable.getComment())) { script.append(SQLConstants.LINE_SEPARATOR); - script.append(SQL_COMMENT_TABLE).append(SQLConstants.SPACE).append(SQLConstants.DOUBLE_QUOTE).append(newTable.getName()).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE) - .append(newTable.getComment()).append(SQLConstants.SINGLE_QUOTE_SEMICOLON_LINE_SEPARATOR); + script.append(SQL_COMMENT_TABLE).append(SQLConstants.SPACE).append(newQualifiedName).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE) + .append(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(newTable.getComment())).append(SQLConstants.SINGLE_QUOTE_SEMICOLON_LINE_SEPARATOR); } for (TableColumn tableColumn : newTable.getColumnList()) { PostgreSQLColumnTypeEnum typeEnum = PostgreSQLColumnTypeEnum.getByType(tableColumn.getColumnType()); @@ -340,17 +401,17 @@ public String buildPageLimit(PageLimitRequest request) { @Override public String buildCreateDatabase(Database database) { StringBuilder sqlBuilder = new StringBuilder(); - sqlBuilder.append(SQL_CREATE_DATABASE + database.getName() + SQLConstants.DOUBLE_QUOTE); + sqlBuilder.append(SQL_CREATE_DATABASE).append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(database.getName())); sqlBuilder.append(SQLConstants.LINE_SEPARATOR_SQL_WITH); if (StringUtils.isNotBlank(database.getCharset())) { - sqlBuilder.append(VALUE_LC_CTYPE_EQUAL_SINGLE_QUOTE).append(database.getCharset()).append(VALUE_SINGLE_QUOTE); + sqlBuilder.append(VALUE_LC_CTYPE_EQUAL_SINGLE_QUOTE).append(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(database.getCharset())).append(VALUE_SINGLE_QUOTE); } if (StringUtils.isNotBlank(database.getCollation())) { - sqlBuilder.append(SQL_LC_COLLATE_EQUAL_SINGLE_QUOTE).append(database.getCollation()).append(VALUE_SINGLE_QUOTE); + sqlBuilder.append(SQL_LC_COLLATE_EQUAL_SINGLE_QUOTE).append(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(database.getCollation())).append(VALUE_SINGLE_QUOTE); } if (StringUtils.isNotBlank(database.getComment())) { - sqlBuilder.append(SQL_SEMICOLON_COMMENT_ON_DATABASE_DOUBLE_QUOTE).append(database.getName()).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE).append(database.getComment()).append(SQLConstants.SINGLE_QUOTE_SEMICOLON); + sqlBuilder.append(SQL_SEMICOLON_COMMENT_ON_DATABASE_DOUBLE_QUOTE).append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(database.getName())).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE).append(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(database.getComment())).append(SQLConstants.SINGLE_QUOTE_SEMICOLON); } return sqlBuilder.toString(); } @@ -363,23 +424,16 @@ public String buildDropDatabase(String databaseName) { @Override protected void buildTableName(String databaseName, String schemaName, String tableName, StringBuilder script) { - ISQLIdentifierProcessor postgreSqlIdentifierProcessor = PostgreSQLMetaData.POSTGRE_SQL_IDENTIFIER_PROCESSOR; - if (StringUtils.isNotBlank(databaseName)) { - script.append(postgreSqlIdentifierProcessor.quoteIdentifier(databaseName)).append('.'); - } - if (StringUtils.isNotBlank(schemaName)) { - script.append(postgreSqlIdentifierProcessor.quoteIdentifier(schemaName)).append('.'); - } - - script.append(postgreSqlIdentifierProcessor.quoteIdentifier(tableName)); + script.append(quoteQualifiedIdentifier(databaseName, schemaName, tableName)); } @Override protected void buildColumns(List columnList, StringBuilder script) { - ISQLIdentifierProcessor postgreSqlIdentifierProcessor = PostgreSQLMetaData.POSTGRE_SQL_IDENTIFIER_PROCESSOR; if (CollectionUtils.isNotEmpty(columnList)) { script.append(SQLConstants.SPACE_OPEN_PARENTHESIS) - .append(columnList.stream().map(postgreSqlIdentifierProcessor::quoteIdentifier).collect(Collectors.joining(SQLConstants.COMMA))) + .append(columnList.stream() + .map(PostgreSQLIdentifierProcessor.INSTANCE::quoteIdentifierAlways) + .collect(Collectors.joining(SQLConstants.COMMA))) .append(SQLConstants.CLOSE_PARENTHESIS_SPACE); } } @@ -387,12 +441,13 @@ protected void buildColumns(List columnList, StringBuilder script) { @Override public String buildCreateSchema(Schema schema) { StringBuilder sqlBuilder = new StringBuilder(); - sqlBuilder.append(SQL_CREATE_SCHEMA + schema.getName() + SQLConstants.DOUBLE_QUOTE); + sqlBuilder.append(SQL_CREATE_SCHEMA).append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schema.getName())); if (StringUtils.isNotBlank(schema.getOwner())) { - sqlBuilder.append(SQLConstants.SCHEMA_AUTHORIZATION_SQL).append(schema.getOwner()); + sqlBuilder.append(SQLConstants.SCHEMA_AUTHORIZATION_SQL) + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schema.getOwner())); } if (StringUtils.isNotBlank(schema.getComment())) { - sqlBuilder.append(SQL_SEMICOLON_COMMENT_ON_SCHEMA_DOUBLE_QUOTE).append(schema.getName()).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE).append(schema.getComment()).append(SQLConstants.SINGLE_QUOTE_SEMICOLON); + sqlBuilder.append(SQL_SEMICOLON_COMMENT_ON_SCHEMA_DOUBLE_QUOTE).append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schema.getName())).append(VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE).append(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(schema.getComment())).append(SQLConstants.SINGLE_QUOTE_SEMICOLON); } return sqlBuilder.toString(); } @@ -403,24 +458,12 @@ public String buildDropSchema(String schemaName) { } private static String quotePostgreSqlIdentifier(String name) { - if (StringUtils.isBlank(name)) { - return name; - } - String identifier = name; - if (identifier.length() >= 2 && identifier.startsWith(SQLConstants.DOUBLE_QUOTE) - && identifier.endsWith(SQLConstants.DOUBLE_QUOTE)) { - identifier = identifier.substring(1, identifier.length() - 1); - } - return SQLConstants.DOUBLE_QUOTE - + identifier.replace(SQLConstants.DOUBLE_QUOTE, - SQLConstants.DOUBLE_QUOTE + SQLConstants.DOUBLE_QUOTE) - + SQLConstants.DOUBLE_QUOTE; + return PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(name); } private static String quotePostgreSqlStringLiteral(String value) { return SQLConstants.SINGLE_QUOTE - + value.replace(SQLConstants.SINGLE_QUOTE, - SQLConstants.SINGLE_QUOTE + SQLConstants.SINGLE_QUOTE) + + PostgreSQLIdentifierProcessor.INSTANCE.escapeString(value) + SQLConstants.SINGLE_QUOTE; } @@ -435,7 +478,8 @@ public String buildCreateView(ModifyView modifyView) { String tempClause = modifyView.getStorageClause(); if (StringUtils.isNotBlank(tempClause)) { - createViewSqlBuilder.append(tempClause).append(SQLConstants.SPACE); + createViewSqlBuilder.append(PostgreSqlGuards.requireViewStorageClause(tempClause)) + .append(SQLConstants.SPACE); } if (modifyView.isUseRecursive()) { @@ -456,6 +500,7 @@ public String buildCreateView(ModifyView modifyView) { createViewSqlBuilder.append(SQLConstants.LINE_SEPARATOR).append(viewBody).append(SQLConstants.SPACE); String checkOption = modifyView.getCheckOption(); if (StringUtils.isNotBlank(checkOption)) { + checkOption = PostgreSqlGuards.requireEnumConstant(checkOption, PostgreSQLViewCheckOptionEnum.values(), "view check option"); createViewSqlBuilder.append(SQLConstants.LINE_SEPARATOR_SQL_WITH).append(checkOption).append(SQLConstants.CHECK_OPTION_SQL); } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLColumnTypeEnumConstants.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLColumnTypeEnumConstants.java index 5b0ed840e7..a78286be77 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLColumnTypeEnumConstants.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLColumnTypeEnumConstants.java @@ -13,9 +13,9 @@ public final class PostgreSQLColumnTypeEnumConstants { - public static final String SQL_ALTER_COLUMN = "ALTER COLUMN \""; + public static final String SQL_ALTER_COLUMN = "ALTER COLUMN "; public static final String SQL_COMMENT_COLUMN = "COMMENT ON COLUMN"; - public static final String SQL_DROP_COLUMN = "DROP COLUMN \""; + public static final String SQL_DROP_COLUMN = "DROP COLUMN "; private PostgreSQLColumnTypeEnumConstants() { } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLDBManagerConstants.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLDBManagerConstants.java index c61b3ebeac..eeaccaf4be 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLDBManagerConstants.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLDBManagerConstants.java @@ -30,7 +30,7 @@ public final class PostgreSQLDBManagerConstants { public static final String SQL_DROP_TABLE_EXISTS = "DROP TABLE IF EXISTS "; public static final String SQL_DROP_TYPE_EXISTS = "DROP TYPE IF EXISTS "; public static final String SQL_DROP_VIEW_EXISTS = "DROP VIEW IF EXISTS "; - public static final String SQL_SET_SEARCH_PATH_USER_PUBLIC = "SET search_path TO \"%s\",\"$user\",\"public\""; + public static final String SQL_SET_SEARCH_PATH_USER_PUBLIC = "SET search_path TO %s,\"$user\",\"public\""; private PostgreSQLDBManagerConstants() { } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLIndexTypeEnumConstants.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLIndexTypeEnumConstants.java index f7d4ca3aba..4a520212bc 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLIndexTypeEnumConstants.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLIndexTypeEnumConstants.java @@ -17,8 +17,8 @@ public final class PostgreSQLIndexTypeEnumConstants { public static final String SQL_COMMENT_CONSTRAINT = "COMMENT ON CONSTRAINT"; public static final String SQL_COMMENT_INDEX = "COMMENT ON INDEX"; public static final String SQL_CREATE = "CREATE"; - public static final String SQL_DROP_CONSTRAINT = "DROP CONSTRAINT \""; - public static final String SQL_DROP_INDEX = "DROP INDEX \""; + public static final String SQL_DROP_CONSTRAINT = "DROP CONSTRAINT "; + public static final String SQL_DROP_INDEX = "DROP INDEX "; public static final String SQL_ON = "ON "; private PostgreSQLIndexTypeEnumConstants() { diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLSqlBuilderConstants.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLSqlBuilderConstants.java index 13046a5f8e..cdb3fd0261 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLSqlBuilderConstants.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/constant/PostgreSQLSqlBuilderConstants.java @@ -33,23 +33,23 @@ public final class PostgreSQLSqlBuilderConstants { public static final String SQL_WHERE_CTID_IN_OPEN_PAREN_SELECT_CTID_FROM = " where ctid in (select ctid from "; public static final String VALUE_LIMIT_1_CLOSE_PAREN = " limit 1)"; - public static final String VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE = "\" IS '"; - public static final String VALUE_DOUBLE_QUOTE = "\" "; + public static final String VALUE_DOUBLE_QUOTE_IS_SINGLE_QUOTE = " IS '"; + public static final String VALUE_DOUBLE_QUOTE = " "; public static final String VALUE_DOUBLE_QUOTE_TO_DOUBLE_QUOTE = "\" TO \""; - public static final String VALUE_DOUBLE_QUOTE_2 = "\" \n"; + public static final String VALUE_DOUBLE_QUOTE_2 = " \n"; public static final String VALUE_LC_CTYPE_EQUAL_SINGLE_QUOTE = "\n LC_CTYPE = '"; public static final String VALUE_SINGLE_QUOTE = "' "; public static final String SQL_LC_COLLATE_EQUAL_SINGLE_QUOTE = "\n LC_COLLATE = '"; - public static final String SQL_SEMICOLON_COMMENT_ON_DATABASE_DOUBLE_QUOTE = "; COMMENT ON DATABASE \""; - public static final String SQL_SEMICOLON_COMMENT_ON_SCHEMA_DOUBLE_QUOTE = "; COMMENT ON SCHEMA \""; + public static final String SQL_SEMICOLON_COMMENT_ON_DATABASE_DOUBLE_QUOTE = "; COMMENT ON DATABASE "; + public static final String SQL_SEMICOLON_COMMENT_ON_SCHEMA_DOUBLE_QUOTE = "; COMMENT ON SCHEMA "; public static final String SQL_RECURSIVE = "RECURSIVE "; public static final String UNDEFINED_KEYWORD = "undefined"; public static final String SQL_ALTER_TABLE = "ALTER TABLE "; public static final String SQL_COMMENT_TABLE = "COMMENT ON TABLE"; public static final String SQL_COMMENT_VIEW = "comment on view"; public static final String SQL_CREATE = "CREATE "; - public static final String SQL_CREATE_DATABASE = "CREATE DATABASE \""; - public static final String SQL_CREATE_SCHEMA = "CREATE SCHEMA \""; + public static final String SQL_CREATE_DATABASE = "CREATE DATABASE "; + public static final String SQL_CREATE_SCHEMA = "CREATE SCHEMA "; public static final String SQL_CREATE_TABLE = "CREATE TABLE "; public static final String SQL_LIMIT = " LIMIT "; public static final String SQL_OFFSET = " OFFSET "; diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLColumnTypeEnum.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLColumnTypeEnum.java index df98d12ebc..713c45cb82 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLColumnTypeEnum.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLColumnTypeEnum.java @@ -1,10 +1,11 @@ package ai.chat2db.plugin.postgresql.enums.type; +import ai.chat2db.plugin.postgresql.PostgreSqlGuards; +import ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor; import ai.chat2db.spi.IColumnBuilder; import ai.chat2db.community.domain.api.enums.plugin.EditStatusEnum; import ai.chat2db.community.domain.api.model.metadata.ColumnType; import ai.chat2db.community.domain.api.model.metadata.TableColumn; -import ai.chat2db.spi.util.SqlUtils; import com.google.common.collect.Maps; import org.apache.commons.lang3.StringUtils; @@ -73,11 +74,11 @@ public enum PostgreSQLColumnTypeEnum implements IColumnBuilder { - private static Map COLUMN_TYPE_MAP = Maps.newHashMap(); + private static final Map COLUMN_TYPE_MAP = Maps.newHashMap(); static { for (PostgreSQLColumnTypeEnum value : PostgreSQLColumnTypeEnum.values()) { - COLUMN_TYPE_MAP.put(value.getColumnType().getTypeName(), value); + COLUMN_TYPE_MAP.put(value.getColumnType().getTypeName().toUpperCase(Locale.ROOT), value); } } @@ -89,7 +90,20 @@ public enum PostgreSQLColumnTypeEnum implements IColumnBuilder { } public static PostgreSQLColumnTypeEnum getByType(String dataType) { - return COLUMN_TYPE_MAP.get(SqlUtils.removeDigits(dataType.toUpperCase())); + if (StringUtils.isBlank(dataType)) { + return null; + } + String typeExpression = PostgreSqlGuards.requireColumnTypeExpression(dataType); + String baseType = typeExpression; + int argumentsStart = baseType.indexOf('('); + if (argumentsStart >= 0) { + baseType = baseType.substring(0, argumentsStart); + } + while (baseType.stripTrailing().endsWith("[]")) { + baseType = baseType.stripTrailing(); + baseType = baseType.substring(0, baseType.length() - 2); + } + return COLUMN_TYPE_MAP.get(baseType.trim().toUpperCase(Locale.ROOT)); } public static List getTypes() { @@ -102,15 +116,45 @@ public ColumnType getColumnType() { return columnType; } + public static String buildCreateColumnSqlSafely(TableColumn column) { + PostgreSQLColumnTypeEnum type = getByType(column.getColumnType()); + return type == null ? buildSafeFallbackColumn(column, false) : type.buildCreateColumnSql(column); + } + + public static String buildAICreateColumnSqlSafely(TableColumn column) { + PostgreSQLColumnTypeEnum type = getByType(column.getColumnType()); + return type == null ? buildSafeFallbackColumn(column, true) : type.buildAICreateColumnSql(column); + } + + public static String buildModifyColumnSafely(TableColumn column) { + PostgreSQLColumnTypeEnum type = getByType(column.getColumnType()); + if (type != null) { + return type.buildModifyColumn(column); + } + if (EditStatusEnum.DELETE.name().equals(column.getEditStatus())) { + return SQL_DROP_COLUMN + PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName()); + } + if (EditStatusEnum.ADD.name().equals(column.getEditStatus())) { + return "ADD COLUMN " + buildSafeFallbackColumn(column, false); + } + if (EditStatusEnum.MODIFY.name().equals(column.getEditStatus())) { + String columnName = PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName()); + String dataType = PostgreSqlGuards.requireColumnTypeExpression(column.getColumnType()); + return SQL_ALTER_COLUMN + columnName + " TYPE " + dataType + + " USING " + columnName + "::" + dataType; + } + return ""; + } + @Override public String buildCreateColumnSql(TableColumn column) { - PostgreSQLColumnTypeEnum type = COLUMN_TYPE_MAP.get(column.getColumnType().toUpperCase()); + PostgreSQLColumnTypeEnum type = COLUMN_TYPE_MAP.get(column.getColumnType().toUpperCase(Locale.ROOT)); if (type == null) { - return buildDefaultColumn(column, false); + return buildSafeFallbackColumn(column, false); } StringBuilder script = new StringBuilder(); - script.append("\"").append(column.getName()).append("\"").append(" "); + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())).append(" "); script.append(buildDataType(column, type)).append(" "); @@ -126,13 +170,13 @@ public String buildCreateColumnSql(TableColumn column) { @Override public String buildAICreateColumnSql(TableColumn column) { - PostgreSQLColumnTypeEnum type = COLUMN_TYPE_MAP.get(column.getColumnType().toUpperCase()); + PostgreSQLColumnTypeEnum type = COLUMN_TYPE_MAP.get(column.getColumnType().toUpperCase(Locale.ROOT)); if (type == null) { - return buildDefaultColumn(column, false); + return buildSafeFallbackColumn(column, true); } StringBuilder script = new StringBuilder(); - script.append("\"").append(column.getName()).append("\"").append(" "); + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())).append(" "); script.append(buildDataType(column, type)).append(" "); @@ -153,13 +197,13 @@ private String buildCollation(TableColumn column, PostgreSQLColumnTypeEnum type) if (!type.getColumnType().isSupportCollation() || StringUtils.isEmpty(column.getCollationName())) { return ""; } - return StringUtils.join("\"", column.getCollationName(), "\""); + return PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getCollationName()); } @Override public String buildModifyColumn(TableColumn column) { if (EditStatusEnum.DELETE.name().equals(column.getEditStatus())) { - return StringUtils.join(SQL_DROP_COLUMN, column.getName() + "\""); + return StringUtils.join(SQL_DROP_COLUMN, PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())); } else if (EditStatusEnum.ADD.name().equals(column.getEditStatus())) { return StringUtils.join("ADD COLUMN ", buildCreateColumnSql(column)); } else if (EditStatusEnum.MODIFY.name().equals(column.getEditStatus())) { @@ -175,26 +219,26 @@ public String buildModifyColumn(TableColumn column) { boolean sizeChanged = oldColumnSize != null && newColumnSize != null && !oldColumnSize.equals(newColumnSize); boolean scaleChanged = oldDecimalDigits != null && newDecimalDigits != null && !oldDecimalDigits.equals(newDecimalDigits); if (!sameType || sizeChanged || scaleChanged) { - String newDataTypeClause = buildDataType(column, this); + String newDataTypeClause = buildModifiedDataType(column); script.append(SQL_ALTER_COLUMN) - .append(column.getName()) - .append("\" TYPE ") + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())) + .append(" TYPE ") .append(newDataTypeClause); - script.append(" USING \"") - .append(column.getName()) - .append("\"::") + script.append(" USING ") + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())) + .append("::") .append(newDataTypeClause) .append(",\n"); } } else { - String newDataTypeClause = buildDataType(column, this); + String newDataTypeClause = buildModifiedDataType(column); script.append(SQL_ALTER_COLUMN) - .append(column.getName()) - .append("\" TYPE ") + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())) + .append(" TYPE ") .append(newDataTypeClause) - .append(" USING \"") - .append(column.getName()) - .append("\"::") + .append(" USING ") + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())) + .append("::") .append(newDataTypeClause) .append(",\n"); } @@ -207,12 +251,12 @@ public String buildModifyColumn(TableColumn column) { if (oldColumn != null) { Integer oldNullable = oldColumn.getNullable(); if (oldNullable != null && newNullable != null && !oldNullable.equals(newNullable)) { - script.append("\tALTER COLUMN \"").append(columnName).append("\" ") + script.append("\tALTER COLUMN ").append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(columnName)).append(" ") .append(shouldDropNotNull ? "DROP" : "SET").append(" NOT NULL ,\n"); } } else { if (newNullable != null) { - script.append("\tALTER COLUMN \"").append(columnName).append("\" ") + script.append("\tALTER COLUMN ").append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(columnName)).append(" ") .append(shouldDropNotNull ? "DROP" : "SET").append(" NOT NULL ,\n"); } } @@ -230,8 +274,8 @@ public String buildModifyColumn(TableColumn column) { if (shouldAppendDefault) { script.append(SQL_ALTER_COLUMN) - .append(column.getName()) - .append("\" SET ") + .append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())) + .append(" SET ") .append(defaultValue) .append(",\n"); } @@ -253,12 +297,22 @@ public String buildComment(TableColumn column, PostgreSQLColumnTypeEnum type) { return ""; } if (column.getOldColumn() == null || !StringUtils.equals(column.getOldColumn().getComment(), column.getComment())) { - return StringUtils.join(SQL_COMMENT_COLUMN, " \"", column.getTableName(), - "\".\"", column.getName(), "\" IS '", column.getComment(), "';"); + String tableName = qualifiedName(column.getSchemaName(), column.getTableName()); + return StringUtils.join(SQL_COMMENT_COLUMN, " ", tableName, ".", + PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName()), " IS '", + PostgreSQLIdentifierProcessor.INSTANCE.escapeString(column.getComment()), "';"); } return ""; } + private static String qualifiedName(String schemaName, String objectName) { + String quotedName = PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(objectName); + if (StringUtils.isBlank(schemaName)) { + return quotedName; + } + return PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName) + "." + quotedName; + } + private String buildDefaultValue(TableColumn column, PostgreSQLColumnTypeEnum type) { if (!type.getColumnType().isSupportDefaultValue() || StringUtils.isEmpty(column.getDefaultValue())) { return ""; @@ -273,17 +327,44 @@ private String buildDefaultValue(TableColumn column, PostgreSQLColumnTypeEnum ty } if (Arrays.asList(CHAR, VARCHAR).contains(type)) { - return StringUtils.join("DEFAULT '", column.getDefaultValue(), "'"); + if (PostgreSqlGuards.isFunctionOrCastExpression(column.getDefaultValue())) { + return StringUtils.join("DEFAULT ", + PostgreSqlGuards.requireDefaultExpression(column.getDefaultValue())); + } + return StringUtils.join("DEFAULT '", PostgreSQLIdentifierProcessor.INSTANCE.escapeString(column.getDefaultValue()), "'"); } if (Arrays.asList(TIMESTAMP, TIME, TIMETZ, TIMESTAMPTZ, DATE).contains(type)) { - if ("CURRENT_TIMESTAMP".equalsIgnoreCase(column.getDefaultValue().trim())) { - return StringUtils.join("DEFAULT ", column.getDefaultValue()); + if (PostgreSqlGuards.isTemporalExpression(column.getDefaultValue())) { + return StringUtils.join("DEFAULT ", + PostgreSqlGuards.requireDefaultExpression(column.getDefaultValue())); } - return StringUtils.join("DEFAULT '", column.getDefaultValue(), "'"); + return StringUtils.join("DEFAULT '", PostgreSQLIdentifierProcessor.INSTANCE.escapeString(column.getDefaultValue()), "'"); } - return StringUtils.join("DEFAULT ", column.getDefaultValue()); + return StringUtils.join("DEFAULT ", PostgreSqlGuards.requireDefaultExpression(column.getDefaultValue())); + } + + private static String buildSafeFallbackColumn(TableColumn column, boolean includeAiComment) { + StringBuilder script = new StringBuilder(); + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getName())) + .append(" ") + .append(PostgreSqlGuards.requireColumnTypeExpression(column.getColumnType())); + if (column.getNullable() != null) { + script.append(column.getNullable() == 1 ? " NULL" : " NOT NULL"); + } + if (StringUtils.isNotEmpty(column.getDefaultValue())) { + if ("EMPTY_STRING".equalsIgnoreCase(column.getDefaultValue().trim())) { + script.append(" DEFAULT ''"); + } else { + script.append(" DEFAULT ") + .append(PostgreSqlGuards.requireDefaultExpression(column.getDefaultValue())); + } + } + if (includeAiComment) { + script.append(" ").append(PostgreSQLColumnTypeEnum.TEXT.buildAICreateColumnCommentSql(column)); + } + return script.toString(); } private String buildNullable(TableColumn column, PostgreSQLColumnTypeEnum type) { @@ -334,4 +415,11 @@ private String buildDataType(TableColumn column, PostgreSQLColumnTypeEnum type) return columnType; } + private String buildModifiedDataType(TableColumn column) { + if (COLUMN_TYPE_MAP.containsKey(column.getColumnType().toUpperCase(Locale.ROOT))) { + return buildDataType(column, this); + } + return PostgreSqlGuards.requireColumnTypeExpression(column.getColumnType()); + } + } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLIndexTypeEnum.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLIndexTypeEnum.java index e4ecc099fa..6de0e4d1c2 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLIndexTypeEnum.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/enums/type/PostgreSQLIndexTypeEnum.java @@ -1,5 +1,6 @@ package ai.chat2db.plugin.postgresql.enums.type; +import ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor; import ai.chat2db.community.domain.api.enums.plugin.EditStatusEnum; import ai.chat2db.community.domain.api.model.metadata.IndexType; import ai.chat2db.community.domain.api.model.metadata.TableIndex; @@ -74,7 +75,7 @@ public String buildIndexScript(TableIndex tableIndex) { script.append(buildIndexUnique(tableIndex)).append(" "); script.append(buildIndexConcurrently(tableIndex)).append(" "); script.append(buildIndexName(tableIndex)).append(" "); - script.append(SQL_ON).append("\"").append(tableIndex.getTableName()).append("\"").append(" "); + script.append(SQL_ON).append(qualifiedName(tableIndex.getSchemaName(), tableIndex.getTableName())).append(" "); script.append(buildIndexMethod(tableIndex)).append(" "); script.append(buildIndexColumn(tableIndex)); } else { @@ -92,16 +93,16 @@ private String buildForeignColum(TableIndex tableIndex) { StringBuilder script = new StringBuilder(); script.append(" REFERENCES "); if (StringUtils.isNotBlank(tableIndex.getForeignSchemaName())) { - script.append(tableIndex.getForeignSchemaName()).append("."); + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableIndex.getForeignSchemaName())).append("."); } if (StringUtils.isNotBlank(tableIndex.getForeignTableName())) { - script.append(tableIndex.getForeignTableName()).append(" "); + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableIndex.getForeignTableName())).append(" "); } if (CollectionUtils.isNotEmpty(tableIndex.getForeignColumnNamelist())) { script.append("("); for (String column : tableIndex.getForeignColumnNamelist()) { if (StringUtils.isNotBlank(column)) { - script.append("\"").append(column).append("\"").append(","); + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column)).append(","); } } script.deleteCharAt(script.length() - 1); @@ -114,7 +115,7 @@ private String buildForeignColum(TableIndex tableIndex) { private String buildIndexMethod(TableIndex tableIndex) { if (StringUtils.isNotBlank(tableIndex.getMethod())) { - return "USING " + tableIndex.getMethod(); + return "USING " + PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableIndex.getMethod()); } else { return ""; } @@ -141,10 +142,12 @@ public String buildIndexComment(TableIndex tableIndex) { return ""; } else if (NORMAL.equals(this)) { return StringUtils.join(SQL_COMMENT_INDEX, " ", - "\"", tableIndex.getName(), "\" IS '", tableIndex.getComment(), "';"); + qualifiedName(tableIndex.getSchemaName(), tableIndex.getName()), " IS '", PostgreSQLIdentifierProcessor.INSTANCE.escapeString(tableIndex.getComment()), "';"); } else { - return StringUtils.join(SQL_COMMENT_CONSTRAINT, " \"", tableIndex.getName(), "\" ON \"", tableIndex.getSchemaName(), - "\".\"", tableIndex.getTableName(), "\" IS '", tableIndex.getComment(), "';"); + return StringUtils.join(SQL_COMMENT_CONSTRAINT, " ", + PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableIndex.getName()), " ON ", + qualifiedName(tableIndex.getSchemaName(), tableIndex.getTableName()), " IS '", + PostgreSQLIdentifierProcessor.INSTANCE.escapeString(tableIndex.getComment()), "';"); } } @@ -153,7 +156,7 @@ private String buildIndexColumn(TableIndex tableIndex) { script.append("("); for (TableIndexColumn column : tableIndex.getColumnList()) { if (StringUtils.isNotBlank(column.getColumnName())) { - script.append("\"").append(column.getColumnName()).append("\"").append(","); + script.append(PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(column.getColumnName())).append(","); } } script.deleteCharAt(script.length() - 1); @@ -162,7 +165,7 @@ private String buildIndexColumn(TableIndex tableIndex) { } private String buildIndexName(TableIndex tableIndex) { - return "\"" + tableIndex.getName() + "\""; + return PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableIndex.getName()); } public String buildModifyIndex(TableIndex tableIndex) { @@ -181,8 +184,17 @@ public String buildModifyIndex(TableIndex tableIndex) { private String buildDropIndex(TableIndex tableIndex) { if (NORMAL.equals(this)) { - return StringUtils.join(SQL_DROP_INDEX, tableIndex.getOldName(), "\""); + return StringUtils.join(SQL_DROP_INDEX, + qualifiedName(tableIndex.getSchemaName(), tableIndex.getOldName())); } - return StringUtils.join(SQL_DROP_CONSTRAINT, tableIndex.getOldName(), "\""); + return StringUtils.join(SQL_DROP_CONSTRAINT, PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableIndex.getOldName())); + } + + private static String qualifiedName(String schemaName, String objectName) { + String quotedName = PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(objectName); + if (StringUtils.isBlank(schemaName)) { + return quotedName; + } + return PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName) + "." + quotedName; } } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/identifier/PostgreSQLIdentifierProcessor.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/identifier/PostgreSQLIdentifierProcessor.java index e4518dc2cc..52f1854b96 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/identifier/PostgreSQLIdentifierProcessor.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/identifier/PostgreSQLIdentifierProcessor.java @@ -4,9 +4,19 @@ import org.apache.commons.lang3.StringUtils; import java.util.HashSet; +import java.util.Locale; import java.util.Set; +/** + * PostgreSQL dialect identifier processor: double-quoted identifiers with embedded-quote + * doubling, and single-quote doubling for string literals + * (standard_conforming_strings=on, so backslash is not an escape character). + * Shared stateless instance available via {@link #INSTANCE} for call sites without MetaData access. + */ public class PostgreSQLIdentifierProcessor extends DefaultSQLIdentifierProcessor { + + public static final PostgreSQLIdentifierProcessor INSTANCE = new PostgreSQLIdentifierProcessor(); + private static final Set PGSQL_RESERVED_KEYWORDS = new HashSet<>(); static { @@ -92,45 +102,78 @@ public class PostgreSQLIdentifierProcessor extends DefaultSQLIdentifierProcessor @Override public boolean isReservedKeyword(String identifier, Integer majorVersion, Integer minorVersion) { - return PGSQL_RESERVED_KEYWORDS.contains(identifier); + return identifier != null && PGSQL_RESERVED_KEYWORDS.contains(identifier.toUpperCase(Locale.ROOT)); } + /** + * SPI-facing conditional quoting: {@code null} stays {@code null}, blank is returned + * unchanged, a valid plain identifier that is not a reserved keyword is returned + * unquoted, and anything else is wrapped in double quotes with one surrounding + * double quotes doubled as raw identifier content. + */ @Override public String quoteIdentifier(String identifier, Integer majorVersion, Integer minorVersion) { - if (isValidIdentifier(identifier)) { - if (containsUpperCase(identifier) || isReservedKeyword(identifier.toUpperCase(), majorVersion, minorVersion)) { - return StringUtils.wrap(identifier, '"'); - } - return identifier; - } - return StringUtils.wrap(identifier, '"'); - + return quoteIdentifier(identifier); } @Override public String quoteIdentifier(String identifier) { - if (isQuoteIdentifier(identifier)) { + if (StringUtils.isBlank(identifier)) { return identifier; } - if (isValidIdentifier(identifier)) { - if (containsUpperCase(identifier) || isReservedKeyword(identifier.toUpperCase(), null, null)) { - return StringUtils.wrap(identifier, '"'); - } + // PostgreSQL folds unquoted identifiers to lowercase, so mixed-case names must stay quoted. + if (isValidIdentifier(identifier) && !containsUpperCase(identifier) + && !isReservedKeyword(identifier, null, null)) { return identifier; } - return StringUtils.wrap(identifier, '"'); + return quoteIdentifierAlways(identifier); } + /** + * Conditional quote variant that preserves the original identifier case. + */ @Override public String quoteIdentifierIgnoreCase(String identifier) { - if (isValidIdentifier(identifier)) { - if (isReservedKeyword(identifier.toUpperCase(), null, null)) { - return StringUtils.wrap(identifier, '"'); - } - return identifier; + return quoteIdentifier(identifier); + } + + /** + * Unconditionally wraps with double quotes and doubles every embedded double quote. + * For DDL-generation call sites that must always emit quoted identifiers. Returns + * {@code null} for {@code null}. + */ + @Override + public String quoteIdentifierAlways(String identifier) { + if (identifier == null) { + return null; } - return StringUtils.wrap(identifier, '"'); + return "\"" + escapeIdentifierContent(identifier) + "\""; + } + + /** + * Escapes a value interpolated into a single-quoted SQL string literal by + * doubling every single quote. Returns {@code null} for {@code null}. + */ + @Override + public String escapeString(String str) { + return str == null ? null : StringUtils.replace(str, "'", "''"); + } + + private static String escapeIdentifierContent(String identifier) { + if (identifier == null) { + return null; + } + return StringUtils.replace(identifier, "\"", "\"\""); + } + + /** + * Escapes identifier content for a position already surrounded by double + * quotes by doubling every embedded double quote. Returns {@code null} for + * {@code null}. + */ + public static String escapeIdentifier(String identifier) { + return escapeIdentifierContent(identifier); } @Override @@ -138,7 +181,7 @@ public String convertIdentifierCase(String identifier) { if (StringUtils.isBlank(identifier)) { return identifier; } else { - return identifier.toLowerCase(); + return identifier.toLowerCase(Locale.ROOT); } } } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/value/template/PostgreSQLDmlValueTemplate.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/value/template/PostgreSQLDmlValueTemplate.java index 87d976ba86..f9feb795f7 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/value/template/PostgreSQLDmlValueTemplate.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/main/java/ai/chat2db/plugin/postgresql/value/template/PostgreSQLDmlValueTemplate.java @@ -1,5 +1,8 @@ package ai.chat2db.plugin.postgresql.value.template; +import ai.chat2db.plugin.postgresql.PostgreSqlGuards; +import ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor; + import static ai.chat2db.plugin.postgresql.constant.PostgreSQLDmlValueTemplateConstants.*; @@ -8,21 +11,18 @@ public class PostgreSQLDmlValueTemplate { - - - public static String wrapBit(String value) { - return String.format(BIT_TEMPLATE, value); + return String.format(BIT_TEMPLATE, PostgreSqlGuards.requireBitLiteral(value)); } public static String wrapBytea(String value) { - return String.format(BYTEA_VALUE, value); + return String.format(BYTEA_VALUE, PostgreSqlGuards.requireHexLiteral(value)); } public static String wrapJsonb(String value) { - return String.format(JSONB_TEMPLATE, value); + return String.format(JSONB_TEMPLATE, PostgreSQLIdentifierProcessor.INSTANCE.escapeString(value)); } public static String wrapJson(String value) { - return String.format(JSON_TEMPLATE, value); + return String.format(JSON_TEMPLATE, PostgreSQLIdentifierProcessor.INSTANCE.escapeString(value)); } } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManagerTest.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManagerTest.java index 2048cf7590..a3c4922ace 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManagerTest.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/PostgreSQLDBManagerTest.java @@ -1,5 +1,6 @@ package ai.chat2db.plugin.postgresql; +import ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor; import org.junit.jupiter.api.Test; import java.sql.Connection; @@ -28,6 +29,23 @@ void buildsDropSchemaSqlWithoutCascade() { assertFalse(manage.sql.contains("CASCADE")); } + @Test + void buildsSchemaQualifiedTableStatementsWithoutDoubleQuotingServiceNames() throws Exception { + PostgreSQLDBManager manager = new PostgreSQLDBManager(); + String source = PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways("ord\"ers"); + String target = PostgreSQLIdentifierProcessor.INSTANCE.quoteIdentifierAlways("ord\"ers_copy"); + + assertEquals("DROP TABLE \"analytics\".\"ord\"\"ers\"", + manager.dropTable(null, "ignored_database", "analytics", "ord\"ers")); + assertEquals("TRUNCATE TABLE \"analytics\".\"ord\"\"ers\"", + manager.truncateTable(null, "ignored_database", "analytics", source)); + assertEquals("CREATE TABLE \"analytics\".\"ord\"\"ers_copy\" AS TABLE " + + "\"analytics\".\"ord\"\"ers\" WITH DATA", + PostgreSQLDBManager.buildCopyTableSql("analytics", source, target, true)); + assertEquals("CREATE TABLE \"ord\"\"ers_copy\" AS TABLE \"ord\"\"ers\" WITH NO DATA", + PostgreSQLDBManager.buildCopyTableSql(null, source, target, false)); + } + private static class TestPostgreSQLDBManager extends PostgreSQLDBManager { private String sql; diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/PostgreSQLIdentifierProcessorTest.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/PostgreSQLIdentifierProcessorTest.java new file mode 100644 index 0000000000..60fe3ea0b1 --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/PostgreSQLIdentifierProcessorTest.java @@ -0,0 +1,411 @@ +package ai.chat2db.plugin.postgresql; + +import ai.chat2db.community.domain.api.config.TableBuilderConfig; +import ai.chat2db.community.domain.api.model.metadata.Database; +import ai.chat2db.community.domain.api.model.metadata.Schema; +import ai.chat2db.community.domain.api.model.metadata.Table; +import ai.chat2db.community.domain.api.model.metadata.TableColumn; +import ai.chat2db.community.domain.api.model.metadata.TableIndex; +import ai.chat2db.community.domain.api.model.metadata.TableIndexColumn; +import ai.chat2db.community.domain.api.model.view.ModifyView; +import ai.chat2db.plugin.postgresql.builder.PostgreSQLSqlBuilder; +import ai.chat2db.plugin.postgresql.enums.PostgreSQLViewCheckOptionEnum; +import ai.chat2db.plugin.postgresql.enums.type.PostgreSQLColumnTypeEnum; +import ai.chat2db.plugin.postgresql.enums.type.PostgreSQLIndexTypeEnum; +import ai.chat2db.plugin.postgresql.identifier.PostgreSQLIdentifierProcessor; +import ai.chat2db.plugin.postgresql.value.template.PostgreSQLDmlValueTemplate; +import org.junit.jupiter.api.Test; + +import java.util.List; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class PostgreSQLIdentifierProcessorTest { + + @Test + void escapeSqlLiteralDoublesSingleQuotes() { + assertEquals("a''b", PostgreSQLIdentifierProcessor.INSTANCE.escapeString("a'b")); + assertEquals("''", PostgreSQLIdentifierProcessor.INSTANCE.escapeString("'")); + assertEquals("plain", PostgreSQLIdentifierProcessor.INSTANCE.escapeString("plain")); + // backslash is NOT an escape character under standard_conforming_strings=on + assertEquals("a\\b", PostgreSQLIdentifierProcessor.INSTANCE.escapeString("a\\b")); + assertNull(PostgreSQLIdentifierProcessor.INSTANCE.escapeString(null)); + } + + @Test + void quoteIdentifierIsConditionalForSpiConsumers() { + PostgreSQLIdentifierProcessor processor = PostgreSQLIdentifierProcessor.INSTANCE; + // null/blank pass through + assertNull(processor.quoteIdentifier(null)); + assertEquals("", processor.quoteIdentifier("")); + assertEquals(" ", processor.quoteIdentifier(" ")); + // valid plain identifiers that are not reserved keywords stay unquoted + assertEquals("plain", processor.quoteIdentifier("plain")); + assertEquals("my_table", processor.quoteIdentifier("my_table")); + // reserved keywords are quoted + assertEquals("\"select\"", processor.quoteIdentifier("select")); + assertEquals("\"USER\"", processor.quoteIdentifier("USER")); + // anything else is wrapped with embedded-quote doubling + assertEquals("\"weird\"\"name\"", processor.quoteIdentifier("weird\"name")); + assertEquals("\"a\"\"; DROP TABLE b; --\"", processor.quoteIdentifier("a\"; DROP TABLE b; --")); + // boundary quotes are raw identifier content and are doubled like embedded quotes + assertEquals("\"\"\"a\"\"b\"\"\"", processor.quoteIdentifier("\"a\"b\"")); + assertEquals("\"\"\"quoted\"\"\"", processor.quoteIdentifier("\"quoted\"")); + // the versioned overload delegates to the same conditional behavior + assertEquals("plain", processor.quoteIdentifier("plain", 15, 0)); + assertEquals("\"select\"", processor.quoteIdentifier("select", 15, 0)); + } + + @Test + void quoteIdentifierIgnoreCaseRemainsConditionalAndPreservesCase() { + PostgreSQLIdentifierProcessor processor = PostgreSQLIdentifierProcessor.INSTANCE; + assertNull(processor.quoteIdentifierIgnoreCase(null)); + assertEquals("plain", processor.quoteIdentifierIgnoreCase("plain")); + assertEquals("\"MyTable\"", processor.quoteIdentifierIgnoreCase("MyTable")); + assertEquals("\"weird\"\"name\"", processor.quoteIdentifierIgnoreCase("weird\"name")); + } + + @Test + void quoteIdentifierAlwaysWrapsUnconditionally() { + PostgreSQLIdentifierProcessor processor = PostgreSQLIdentifierProcessor.INSTANCE; + assertNull(processor.quoteIdentifierAlways(null)); + assertEquals("\"\"", processor.quoteIdentifierAlways("")); + assertEquals("\"plain\"", processor.quoteIdentifierAlways("plain")); + assertEquals("\"my_table\"", processor.quoteIdentifierAlways("my_table")); + assertEquals("\"weird\"\"name\"", processor.quoteIdentifierAlways("weird\"name")); + assertEquals("\"a\"\"; DROP TABLE b; --\"", processor.quoteIdentifierAlways("a\"; DROP TABLE b; --")); + assertEquals("\"\"\"a\"\"b\"\"\"", processor.quoteIdentifierAlways("\"a\"b\"")); + assertEquals("\"\"\"quoted\"\"\"", processor.quoteIdentifierAlways("\"quoted\"")); + } + + @Test + void metadataNameUsesPostgresqlSchemaQualification() { + PostgreSQLMetaData metaData = new PostgreSQLMetaData(); + assertEquals("\"orders\"", metaData.getMetaDataName("orders")); + assertEquals("\"sales\".\"orders\"", + metaData.getMetaDataName("ignored_database", "sales", "orders")); + } + + @Test + void alwaysQuoteAndRemoveQuoteRoundTripExactRawIdentifiers() { + PostgreSQLIdentifierProcessor processor = PostgreSQLIdentifierProcessor.INSTANCE; + for (String raw : List.of("plain", "a\"b", "\"leading", "trailing\"", "\"both\"", "")) { + assertEquals(raw, processor.removeIdentifierQuote(processor.quoteIdentifierAlways(raw)), raw); + } + } + + @Test + void requirePgNameRejectsInjection() { + assertEquals("btree", PostgreSqlGuards.requirePgName("btree", "index method")); + assertEquals("en_US", PostgreSqlGuards.requirePgName("en_US", "role")); + assertThrows(IllegalArgumentException.class, + () -> PostgreSqlGuards.requirePgName("btree; DROP TABLE t", "index method")); + assertThrows(IllegalArgumentException.class, + () -> PostgreSqlGuards.requirePgName("alice\" ", "schema owner")); + } + + @Test + void requireDefaultExpressionAcceptsLegitDefaults() { + String[] valid = {"0", "-1", "1.5", "+2", "true", "FALSE", "NULL", "CURRENT_TIMESTAMP", "now", + "now()", "gen_random_uuid()", "nextval('audit.event_id_seq'::regclass)", + "timezone('UTC'::text, now())", "'{}'::jsonb", "ARRAY[]::integer[]", + "CURRENT_DATE + 1", "NOT FALSE", "'a'||'b'", "'Y'", "'0'", "'O''Brien'", + "E'line\\nfeed'", "$tag$comma, -- and ; stay literal$tag$", "''"}; + for (String value : valid) { + assertEquals(value, PostgreSqlGuards.requireDefaultExpression(value), "should accept: " + value); + } + } + + @Test + void requireDefaultExpressionRejectsDdlReshapePayloads() { + String[] payloads = { + "0) --", "0 --", "1, x INT", "0 NULL", "0 NOT NULL", "0 CHECK (false)", + "0 UNIQUE", "0 DEFAULT 1", + "'abc", "'a'--", "now()); DROP TABLE x", "'a'; DROP TABLE x--", + "0); DROP TABLE t--" + }; + for (String payload : payloads) { + assertThrows(IllegalArgumentException.class, + () -> PostgreSqlGuards.requireDefaultExpression(payload), "should reject: " + payload); + } + } + + @Test + void requireColumnTypeExpressionAcceptsPostgresqlTypesAndRejectsBreakout() { + for (String type : List.of("numeric(10,2)", "timestamp(3) with time zone", + "public.invoice_state", "\"Tenant\".\"InvoiceType\"", "integer[]")) { + assertEquals(type, PostgreSqlGuards.requireColumnTypeExpression(type)); + } + for (String type : List.of("text, injected integer", "text DEFAULT 0", "text); DROP TABLE t;--")) { + assertThrows(IllegalArgumentException.class, + () -> PostgreSqlGuards.requireColumnTypeExpression(type), type); + } + } + + @Test + void requireBitAndHexLiteralsValidateContent() { + assertEquals("0101", PostgreSqlGuards.requireBitLiteral("0101")); + assertThrows(IllegalArgumentException.class, () -> PostgreSqlGuards.requireBitLiteral("2")); + assertThrows(IllegalArgumentException.class, () -> PostgreSqlGuards.requireBitLiteral("1' OR '1'='1")); + assertEquals("deadBEEF", PostgreSqlGuards.requireHexLiteral("deadBEEF")); + assertThrows(IllegalArgumentException.class, () -> PostgreSqlGuards.requireHexLiteral("zz'; DROP TABLE t;--")); + } + + @Test + void requireEnumConstantRejectsUnknownOption() { + assertEquals("CASCADED", PostgreSqlGuards.requireEnumConstant( + "cascaded", PostgreSQLViewCheckOptionEnum.values(), "view check option")); + assertThrows(IllegalArgumentException.class, () -> PostgreSqlGuards.requireEnumConstant( + "CASCADED; DROP TABLE t", PostgreSQLViewCheckOptionEnum.values(), "view check option")); + } + + @Test + void createTableQuotesNamesAndEscapesComment() { + PostgreSQLSqlBuilder builder = new PostgreSQLSqlBuilder(); + Table table = Table.builder() + .schemaName("s\"x") + .name("a\"; DROP TABLE b; --") + .columnList(List.of()) + .indexList(List.of()) + .comment("x'; DROP TABLE u;--") + .build(); + TableBuilderConfig config = TableBuilderConfig.defaultConfig(); + config.setNeedFullTableName(true); + + String sql = builder.buildCreateTable(table, config); + + assertTrue(sql.contains("\"s\"\"x\".\"a\"\"; DROP TABLE b; --\""), sql); + assertTrue(sql.contains("COMMENT ON TABLE \"s\"\"x\".\"a\"\"; DROP TABLE b; --\" " + + "IS 'x''; DROP TABLE u;--';"), sql); + } + + @Test + void createDatabaseQuotesNameAndEscapesComment() { + PostgreSQLSqlBuilder builder = new PostgreSQLSqlBuilder(); + Database database = new Database(); + database.setName("db\"x"); + database.setComment("c'd"); + + String sql = builder.buildCreateDatabase(database); + + assertTrue(sql.contains("CREATE DATABASE \"db\"\"x\""), sql); + assertTrue(sql.contains("COMMENT ON DATABASE \"db\"\"x\" IS 'c''d';"), sql); + } + + @Test + void createSchemaQuotesNameAndOwner() { + PostgreSQLSqlBuilder builder = new PostgreSQLSqlBuilder(); + Schema benign = new Schema(); + benign.setName("s\"x"); + benign.setOwner("postgres"); + assertTrue(builder.buildCreateSchema(benign).contains("CREATE SCHEMA \"s\"\"x\" AUTHORIZATION \"postgres\""), + builder.buildCreateSchema(benign)); + + Schema malicious = new Schema(); + malicious.setName("s"); + malicious.setOwner("alice; DROP TABLE t"); + assertTrue(builder.buildCreateSchema(malicious) + .contains("AUTHORIZATION \"alice; DROP TABLE t\""), builder.buildCreateSchema(malicious)); + } + + @Test + void createViewRejectsCheckOptionInjectionAndEscapesComment() { + PostgreSQLSqlBuilder builder = new PostgreSQLSqlBuilder(); + ModifyView malicious = new ModifyView(); + malicious.setViewName("v"); + malicious.setViewBody("select 1"); + malicious.setCheckOption("CASCADED; DROP TABLE t"); + assertThrows(IllegalArgumentException.class, () -> builder.buildCreateView(malicious)); + + ModifyView maliciousStorage = new ModifyView(); + maliciousStorage.setViewName("v"); + maliciousStorage.setViewBody("select 1"); + maliciousStorage.setStorageClause("TEMP; DROP TABLE t"); + assertThrows(IllegalArgumentException.class, () -> builder.buildCreateView(maliciousStorage)); + + ModifyView benign = new ModifyView(); + benign.setViewName("v\"x"); + benign.setViewBody("select 1"); + benign.setCheckOption("local"); + benign.setComment("c'd"); + String sql = builder.buildCreateView(benign); + assertTrue(sql.contains("VIEW \"v\"\"x\""), sql); + assertTrue(sql.contains("WITH LOCAL CHECK OPTION"), sql); + assertTrue(sql.contains("is 'c''d';"), sql); + } + + @Test + void createColumnSqlQuotesNameAndEscapesStringDefault() { + TableColumn column = TableColumn.builder() + .name("a\"b") + .columnType("VARCHAR") + .columnSize(255) + .defaultValue("O'Brien") + .build(); + + String sql = PostgreSQLColumnTypeEnum.VARCHAR.buildCreateColumnSql(column); + + assertTrue(sql.contains("\"a\"\"b\""), sql); + assertTrue(sql.contains("DEFAULT 'O''Brien'"), sql); + } + + @Test + void createColumnSqlRejectsRawDefaultInjection() { + TableColumn column = TableColumn.builder() + .name("n") + .columnType("INT4") + .defaultValue("0);DROP TABLE t") + .build(); + + assertThrows(IllegalArgumentException.class, () -> PostgreSQLColumnTypeEnum.INT4.buildCreateColumnSql(column)); + } + + @Test + void createColumnSqlPreservesFunctionDefaultAndSafeFallbackType() { + TableColumn timestamp = TableColumn.builder() + .name("createdAt") + .columnType("TIMESTAMP") + .defaultValue("now()") + .build(); + assertTrue(PostgreSQLColumnTypeEnum.TIMESTAMP.buildCreateColumnSql(timestamp) + .contains("DEFAULT now()")); + + TableColumn castText = TableColumn.builder() + .name("payload") + .columnType("VARCHAR") + .defaultValue("'{}'::text") + .build(); + assertTrue(PostgreSQLColumnTypeEnum.VARCHAR.buildCreateColumnSql(castText) + .contains("DEFAULT '{}'::text")); + + TableColumn custom = TableColumn.builder() + .name("amount\"raw") + .columnType("numeric(12,2)") + .nullable(1) + .defaultValue("0::numeric") + .build(); + assertEquals("\"amount\"\"raw\" numeric(12,2) NULL DEFAULT 0::numeric", + PostgreSQLColumnTypeEnum.buildCreateColumnSqlSafely(custom)); + + TableColumn modified = TableColumn.builder() + .name("amount\"raw") + .columnType("numeric(12,2)") + .editStatus("MODIFY") + .oldColumn(TableColumn.builder().columnType("numeric(10,2)").build()) + .build(); + assertEquals("ALTER COLUMN \"amount\"\"raw\" TYPE numeric(12,2) USING " + + "\"amount\"\"raw\"::numeric(12,2)", + PostgreSQLColumnTypeEnum.buildModifyColumnSafely(modified)); + } + + @Test + void columnCommentQuotesNamesAndEscapesComment() { + TableColumn column = TableColumn.builder() + .schemaName("s\"x") + .tableName("t\"x") + .name("c") + .columnType("TEXT") + .comment("it's") + .build(); + + String sql = PostgreSQLColumnTypeEnum.TEXT.buildComment(column, PostgreSQLColumnTypeEnum.TEXT); + + assertEquals("COMMENT ON COLUMN \"s\"\"x\".\"t\"\"x\".\"c\" IS 'it''s';", sql); + } + + @Test + void modifyColumnDeleteQuotesName() { + TableColumn column = TableColumn.builder() + .name("a\"b") + .columnType("TEXT") + .editStatus("DELETE") + .build(); + + assertEquals("DROP COLUMN \"a\"\"b\"", PostgreSQLColumnTypeEnum.TEXT.buildModifyColumn(column)); + } + + @Test + void indexScriptQuotesNamesAndMethod() { + TableIndex tableIndex = TableIndex.builder() + .schemaName("s\"x") + .name("i\"x") + .type("Normal") + .tableName("t\"b") + .method("btree") + .columnList(List.of(TableIndexColumn.builder().columnName("c\"d").build())) + .build(); + + String sql = PostgreSQLIndexTypeEnum.NORMAL.buildIndexScript(tableIndex); + + assertTrue(sql.contains("\"i\"\"x\""), sql); + assertTrue(sql.contains("ON \"s\"\"x\".\"t\"\"b\""), sql); + assertTrue(sql.contains("USING \"btree\""), sql); + assertTrue(sql.contains("(\"c\"\"d\")"), sql); + + TableIndex evilMethod = TableIndex.builder() + .name("i") + .type("Normal") + .tableName("t") + .method("btree; DROP TABLE t") + .columnList(List.of(TableIndexColumn.builder().columnName("c").build())) + .build(); + assertTrue(PostgreSQLIndexTypeEnum.NORMAL.buildIndexScript(evilMethod) + .contains("USING \"btree; DROP TABLE t\"")); + } + + @Test + void indexCommentAndDropQuoteNamesAndEscapeComment() { + TableIndex tableIndex = TableIndex.builder() + .schemaName("s\"x") + .name("i\"x") + .type("Normal") + .comment("c'd") + .build(); + assertEquals("COMMENT ON INDEX \"s\"\"x\".\"i\"\"x\" IS 'c''d';", + PostgreSQLIndexTypeEnum.NORMAL.buildIndexComment(tableIndex)); + + TableIndex dropped = TableIndex.builder() + .schemaName("s\"x") + .name("i") + .oldName("i\"x") + .type("Normal") + .editStatus("DELETE") + .build(); + assertEquals("DROP INDEX \"s\"\"x\".\"i\"\"x\"", + PostgreSQLIndexTypeEnum.NORMAL.buildModifyIndex(dropped)); + + TableIndex constraint = TableIndex.builder() + .schemaName("s\"x") + .tableName("t\"x") + .name("pk\"x") + .type("Primary") + .comment("c'd") + .build(); + assertEquals("COMMENT ON CONSTRAINT \"pk\"\"x\" ON \"s\"\"x\".\"t\"\"x\" IS 'c''d';", + PostgreSQLIndexTypeEnum.PRIMARY.buildIndexComment(constraint)); + } + + @Test + void dmlValueTemplatesEscapeOrValidate() { + assertEquals("B'0101'", PostgreSQLDmlValueTemplate.wrapBit("0101")); + assertThrows(IllegalArgumentException.class, () -> PostgreSQLDmlValueTemplate.wrapBit("1' OR '1'='1")); + assertEquals("E'\\\\xdeadbeef'::bytea", PostgreSQLDmlValueTemplate.wrapBytea("deadbeef")); + assertThrows(IllegalArgumentException.class, () -> PostgreSQLDmlValueTemplate.wrapBytea("zz'; DROP TABLE t;--")); + assertEquals("'{\"a\":\"b\"}'::json", PostgreSQLDmlValueTemplate.wrapJson("{\"a\":\"b\"}")); + assertEquals("'x''y'::jsonb", PostgreSQLDmlValueTemplate.wrapJsonb("x'y")); + } + + @Test + void conditionalQuoteKeepsMixedCaseQuoted() { + PostgreSQLIdentifierProcessor processor = new PostgreSQLIdentifierProcessor(); + // PostgreSQL folds unquoted identifiers to lowercase: mixed-case names must stay quoted. + assertEquals("\"MyTable\"", processor.quoteIdentifier("MyTable")); + assertEquals("mytable", processor.quoteIdentifier("mytable")); + assertEquals("plain_name", processor.quoteIdentifier("plain_name")); + org.junit.jupiter.api.Assertions.assertNull(processor.quoteIdentifier(null)); + assertEquals("\"SELECT\"", processor.quoteIdentifier("SELECT")); + } +} diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilderTest.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilderTest.java index 0ed330d626..4639987e39 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilderTest.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-postgresql/src/test/java/ai/chat2db/plugin/postgresql/builder/PostgreSQLSqlBuilderTest.java @@ -1,8 +1,16 @@ package ai.chat2db.plugin.postgresql.builder; +import ai.chat2db.community.domain.api.model.metadata.Table; +import ai.chat2db.community.domain.api.model.metadata.TableColumn; import ai.chat2db.community.domain.api.model.view.ModifyView; +import ai.chat2db.spi.model.request.DropTableRequest; +import ai.chat2db.spi.model.request.TruncateTableRequest; +import ai.chat2db.spi.model.request.UpdateSqlRequest; import org.junit.jupiter.api.Test; +import java.util.List; +import java.util.Map; + import static org.junit.jupiter.api.Assertions.assertEquals; class PostgreSQLSqlBuilderTest { @@ -42,6 +50,60 @@ void shouldOmitBlankSchemaFromCreateAndCommentViewNames() { builder.buildCreateView(view)); } + @Test + void shouldEscapeInheritedBuilderPathsAndIgnoreDatabaseQualifier() { + PostgreSQLSqlBuilder builder = new PostgreSQLSqlBuilder(); + String schema = "analytics\"x"; + String table = "orders\"x"; + + assertEquals("SELECT COUNT(1) FROM \"analytics\"\"x\".\"orders\"\"x\"", + builder.buildSelectCount(null, schema, table)); + assertEquals("SELECT COUNT(1) FROM \"analytics\"\"x\".\"orders\"\"x\"", + builder.buildSelectCount("ignored_database", schema, table)); + assertEquals("SELECT * FROM \"analytics\"\"x\".\"orders\"\"x\"", + builder.buildSelectTable("ignored_database", schema, table)); + assertEquals("DROP TABLE \"analytics\"\"x\".\"orders\"\"x\"", + builder.buildDropTable(new DropTableRequest("ignored_database", schema, table))); + assertEquals("TRUNCATE TABLE \"analytics\"\"x\".\"orders\"\"x\"", + builder.buildTruncateTable(new TruncateTableRequest("ignored_database", schema, table))); + } + + @Test + void shouldEscapeInheritedUpdateAndTemplateIdentifiers() { + PostgreSQLSqlBuilder builder = new PostgreSQLSqlBuilder(); + UpdateSqlRequest update = UpdateSqlRequest.builder() + .databaseName("ignored_database") + .schemaName("sales\"schema") + .tableName("orders\"table") + .row(Map.of("total\"value", "42")) + .primaryKeyMap(Map.of("order\"id", "7")) + .build(); + + assertEquals("UPDATE \"sales\"\"schema\".\"orders\"\"table\" SET \"total\"\"value\" = 42" + + " WHERE \"order\"\"id\" = 7", + builder.buildUpdate(update)); + + Table table = Table.builder() + .schemaName("sales\"schema") + .name("orders\"table") + .columnList(List.of(TableColumn.builder().name("total\"value").build())) + .build(); + assertEquals("SELECT \"total\"\"value\" FROM \"sales\"\"schema\".\"orders\"\"table\"", + builder.buildTemplate(table, "SELECT")); + } + + @Test + void shouldPreserveCaseOnlyTableRename() { + PostgreSQLSqlBuilder builder = new PostgreSQLSqlBuilder(); + Table oldTable = Table.builder().schemaName("sales").name("orders") + .columnList(List.of()).indexList(List.of()).build(); + Table newTable = Table.builder().schemaName("sales").name("Orders") + .columnList(List.of()).indexList(List.of()).build(); + + assertEquals("ALTER TABLE \"sales\".\"orders\"\tRENAME TO \"Orders\";\n", + builder.buildAlterTable(oldTable, newTable)); + } + private static ModifyView createView(String schemaName, String viewName, String comment) { ModifyView view = new ModifyView(); view.setSchemaName(schemaName);