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

Filter by extension

Filter by extension

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

import ai.chat2db.spi.IDbManager;
import ai.chat2db.plugin.h2.identifier.H2IdentifierProcessor;
import ai.chat2db.spi.DefaultDBManager;
import ai.chat2db.community.domain.api.model.async.AsyncContext;
import ai.chat2db.spi.sql.Chat2DBContext;
Expand All @@ -25,10 +26,11 @@ public void exportDatabase(Connection connection, String databaseName, String sc
}

private void exportSchema(Connection connection, String schemaName, AsyncContext asyncContext) throws SQLException {
String sql = String.format("SCRIPT NODATA NOPASSWORDS NOSETTINGS DROP SCHEMA %s;", schemaName);
String template = "SCRIPT NODATA NOPASSWORDS NOSETTINGS DROP SCHEMA %s;";
if (asyncContext.isContainsData()) {
sql = sql.replace("NODATA", "");
template = template.replace("NODATA", "");
}
String sql = String.format(template, H2IdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName));
try (PreparedStatement preparedStatement = connection.prepareStatement(sql); ResultSet resultSet = preparedStatement.executeQuery()) {
while (resultSet.next()) {
String script = resultSet.getString("SCRIPT");
Expand All @@ -51,7 +53,8 @@ public void connectDatabase(Connection connection, String database) {
}
String schemaName = connectInfo.getSchemaName();
try {
DefaultSQLExecutor.getInstance().execute(connection, String.format(SQL_SET_SCHEMA, schemaName));
DefaultSQLExecutor.getInstance().execute(connection,
String.format(SQL_SET_SCHEMA, H2IdentifierProcessor.INSTANCE.quoteIdentifierAlways(schemaName)));
} catch (SQLException e) {

}
Expand All @@ -60,6 +63,6 @@ public void connectDatabase(Connection connection, String database) {

@Override
public String dropTable(Connection connection, String databaseName, String schemaName, String tableName) {
return String.format(SQL_DROP_TABLE, tableName);
return String.format(SQL_DROP_TABLE, H2IdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableName));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,9 @@
import ai.chat2db.community.domain.api.model.sql.*;
import ai.chat2db.spi.model.value.*;
import ai.chat2db.community.domain.api.model.view.*;
import ai.chat2db.plugin.h2.identifier.H2IdentifierProcessor;
import ai.chat2db.spi.DefaultSQLExecutor;
import ai.chat2db.spi.ISQLIdentifierProcessor;
import ai.chat2db.spi.util.SortUtils;
import jakarta.validation.constraints.NotEmpty;
import lombok.extern.slf4j.Slf4j;
Expand All @@ -30,6 +32,11 @@
@Slf4j
public class H2Meta extends DefaultMetaService implements IDbMetaData {

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



@Override
Expand All @@ -50,21 +57,21 @@ private String getDDL(Connection connection, String databaseName, String schemaN
while (columns.next()) {
String columnName = columns.getString("COLUMN_NAME");
String columnType = columns.getString("TYPE_NAME");
int dataType = columns.getInt("DATA_TYPE");
int columnSize = columns.getInt("COLUMN_SIZE");
int decimalDigits = columns.getInt("DECIMAL_DIGITS");
String remarks = columns.getString("REMARKS");
String defaultValue = columns.getString("COLUMN_DEF");
String nullable = columns.getInt("NULLABLE") == ResultSetMetaData.columnNullable ? "NULL" : "NOT NULL";
StringBuilder columnDefinition = new StringBuilder();
columnDefinition.append(columnName).append(" ").append(columnType);
if (columnSize != 0) {
columnDefinition.append("(").append(columnSize).append(")");
}
columnDefinition.append(H2IdentifierProcessor.INSTANCE.quoteIdentifierAlways(columnName)).append(" ")
.append(H2SqlGuards.renderMetadataType(columnType, dataType, columnSize, decimalDigits));
columnDefinition.append(" ").append(nullable);
if (defaultValue != null) {
columnDefinition.append(" DEFAULT ").append(defaultValue);
columnDefinition.append(" DEFAULT ").append(H2SqlGuards.escapeColumnDefault(defaultValue));
}
if (remarks != null) {
columnDefinition.append(SQL_COMMENT).append(remarks).append("'");
columnDefinition.append(SQL_COMMENT).append(getSQLIdentifierProcessor().escapeString(remarks)).append("'");
}
columnDefinitions.add(columnDefinition.toString());
}
Expand All @@ -86,15 +93,16 @@ private String getDDL(Connection connection, String databaseName, String schemaN
}

StringBuilder createTableDDL = new StringBuilder(SQL_CREATE_TABLE);
createTableDDL.append(tableName).append(" (\n");
createTableDDL.append(H2IdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableName)).append(" (\n");
createTableDDL.append(String.join(",\n", columnDefinitions));
createTableDDL.append("\n);\n");
for (Map.Entry<String, List<String>> entry : indexMap.entrySet()) {
String indexName = entry.getKey();
List<String> columnList = entry.getValue();
String indexColumns = String.join(", ", columnList);
String createIndexDDL = String.format(SQL_CREATE_INDEX, indexName, tableName,
indexColumns);
String indexColumns = columnList.stream().map(H2IdentifierProcessor.INSTANCE::quoteIdentifierAlways)
.collect(Collectors.joining(", "));
String createIndexDDL = String.format(SQL_CREATE_INDEX, H2IdentifierProcessor.INSTANCE.quoteIdentifierAlways(indexName),
H2IdentifierProcessor.INSTANCE.quoteIdentifierAlways(tableName), indexColumns);
createTableDDL.append(createIndexDDL);
}
return createTableDDL.toString();
Expand All @@ -110,7 +118,8 @@ private String getDDL(Connection connection, String databaseName, String schemaN
public Function function(Connection connection, @NotEmpty String databaseName, String schemaName,
String functionName) {

String sql = String.format(ROUTINES_SQL, "FUNCTION", databaseName, functionName);
String sql = String.format(ROUTINES_SQL, "FUNCTION", getSQLIdentifierProcessor().escapeString(databaseName),
getSQLIdentifierProcessor().escapeString(functionName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Function function = new Function();
function.setDatabaseName(databaseName);
Expand All @@ -133,7 +142,8 @@ public Function function(Connection connection, @NotEmpty String databaseName, S
@Override
public List<Trigger> triggers(Connection connection, String databaseName, String schemaName) {
List<Trigger> triggers = new ArrayList<>();
String sql = String.format(TRIGGER_SQL_LIST, databaseName,schemaName);
String sql = String.format(TRIGGER_SQL_LIST, getSQLIdentifierProcessor().escapeString(databaseName),
getSQLIdentifierProcessor().escapeString(schemaName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
Trigger trigger = new Trigger();
Expand All @@ -150,7 +160,8 @@ public List<Trigger> triggers(Connection connection, String databaseName, String
public Trigger trigger(Connection connection, @NotEmpty String databaseName, String schemaName,
String triggerName) {

String sql = String.format(TRIGGER_SQL, databaseName, triggerName);
String sql = String.format(TRIGGER_SQL, getSQLIdentifierProcessor().escapeString(databaseName),
getSQLIdentifierProcessor().escapeString(triggerName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Trigger trigger = new Trigger();
trigger.setDatabaseName(databaseName);
Expand All @@ -166,7 +177,8 @@ public Trigger trigger(Connection connection, @NotEmpty String databaseName, Str
@Override
public Procedure procedure(Connection connection, @NotEmpty String databaseName, String schemaName,
String procedureName) {
String sql = String.format(ROUTINES_SQL, "PROCEDURE", databaseName, procedureName);
String sql = String.format(ROUTINES_SQL, "PROCEDURE", getSQLIdentifierProcessor().escapeString(databaseName),
getSQLIdentifierProcessor().escapeString(procedureName));
return DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
Procedure procedure = new Procedure();
procedure.setDatabaseName(databaseName);
Expand All @@ -184,7 +196,8 @@ public Procedure procedure(Connection connection, @NotEmpty String databaseName,

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

@Override
public String getMetaDataName(String... names) {
return Arrays.stream(names).filter(name -> StringUtils.isNotBlank(name)).map(name -> "\"" + name + "\"").collect(Collectors.joining("."));
return Arrays.stream(names).filter(name -> StringUtils.isNotBlank(name)).map(H2IdentifierProcessor.INSTANCE::quoteIdentifierAlways).collect(Collectors.joining("."));
}

@Override
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,127 @@
package ai.chat2db.plugin.h2;

import java.sql.Types;
import java.util.regex.Pattern;

import ai.chat2db.plugin.h2.identifier.H2IdentifierProcessor;

/**
* Validation helpers for non-escapable SQL positions in H2 DDL generation
* (column type names and column default expressions reported by JDBC metadata).
* Escaping itself lives in {@link H2IdentifierProcessor}.
*/
public final class H2SqlGuards {

/**
* Conservative allow-list for column type names reported by JDBC metadata
* (e.g. {@code INTEGER}, {@code CHARACTER VARYING}). Anything else is rejected
* so hostile or corrupt metadata cannot smuggle SQL into generated DDL.
*/
private static final Pattern SAFE_TYPE_NAME = Pattern.compile(
"^[A-Za-z][A-Za-z0-9_]*(?:\\s+[A-Za-z][A-Za-z0-9_]*)*$");

private static final String STRING_LITERAL_SOURCE = "'(?:''|[^'])*'";
private static final String IDENTIFIER_SOURCE = "(?:[A-Za-z_][A-Za-z0-9_$]*|\"(?:\"\"|[^\"])+\")";
private static final Pattern STRING_LITERAL = Pattern.compile("^" + STRING_LITERAL_SOURCE + "$");
private static final Pattern NUMERIC_LITERAL = Pattern.compile(
"^[+-]?(?:\\d+(?:\\.\\d*)?|\\.\\d+)(?:[eE][+-]?\\d+)?$");
private static final Pattern SIMPLE_CONSTANT = Pattern.compile("(?i)^(?:NULL|TRUE|FALSE)$");
private static final Pattern CURRENT_TEMPORAL = Pattern.compile(
"(?i)^(?:CURRENT_DATE|CURRENT_TIME|CURRENT_TIMESTAMP|LOCALTIME|LOCALTIMESTAMP)(?:\\(\\d+\\))?$");
private static final Pattern SAFE_NO_ARG_FUNCTION = Pattern.compile(
"(?i)^(?:NOW|RANDOM_UUID|UUID)\\(\\s*(?:\\d+)?\\s*\\)$");
private static final Pattern TYPED_LITERAL = Pattern.compile(
"(?i)^(?:DATE|TIME(?:\\s+WITH\\s+TIME\\s+ZONE)?|TIMESTAMP(?:\\s+WITH\\s+TIME\\s+ZONE)?|UUID|JSON|GEOMETRY)\\s+"
+ STRING_LITERAL_SOURCE + "$");
private static final Pattern BINARY_LITERAL = Pattern.compile("(?i)^(?:X|BINARY)\\s*'[0-9A-F]*'$");
private static final Pattern SEQUENCE_EXPRESSION = Pattern.compile(
"(?i)^NEXT\\s+VALUE\\s+FOR\\s+" + IDENTIFIER_SOURCE + "(?:\\." + IDENTIFIER_SOURCE + ")?$");

private H2SqlGuards() {
}

/**
* Validates a column type name obtained from JDBC metadata before it is embedded
* into generated DDL. Returns the type name unchanged when it matches the
* allow-list; throws otherwise (fail closed).
*/
public static String requireSafeTypeName(String typeName) {
if (typeName != null && !SAFE_TYPE_NAME.matcher(typeName).matches()) {
throw new IllegalArgumentException("Unsafe column type name from metadata: " + typeName);
}
return typeName;
}

/**
* Reconstructs a type declaration from JDBC metadata without treating display width as a
* type parameter. H2 reports values such as 64 for BIGINT and 26 for TIMESTAMP in
* COLUMN_SIZE, but those values are not legal declarations for these types.
*/
public static String renderMetadataType(String typeName, int dataType, int columnSize, int decimalDigits) {
String safeTypeName = requireSafeTypeName(typeName);
StringBuilder declaration = new StringBuilder(safeTypeName);
switch (dataType) {
case Types.CHAR:
case Types.VARCHAR:
case Types.NCHAR:
case Types.NVARCHAR:
case Types.BINARY:
case Types.VARBINARY:
appendTypeArguments(declaration, columnSize, null);
break;
case Types.DECIMAL:
case Types.NUMERIC:
appendTypeArguments(declaration, columnSize, Math.max(decimalDigits, 0));
break;
case Types.FLOAT:
appendTypeArguments(declaration, columnSize, null);
break;
case Types.TIME:
case Types.TIME_WITH_TIMEZONE:
case Types.TIMESTAMP:
case Types.TIMESTAMP_WITH_TIMEZONE:
if (decimalDigits > 0) {
appendTypeArguments(declaration, decimalDigits, null);
}
break;
default:
break;
}
return declaration.toString();
}

private static void appendTypeArguments(StringBuilder declaration, int precision, Integer scale) {
if (precision <= 0) {
return;
}
declaration.append('(').append(precision);
if (scale != null) {
declaration.append(',').append(scale);
}
declaration.append(')');
}

/**
* Validates a column default obtained from JDBC metadata. Only complete literal forms and
* common H2-generated expressions are accepted; invalid input is rejected rather than
* silently converted into a string literal with different semantics.
* Returns an empty string for {@code null}.
*/
public static String escapeColumnDefault(String columnDefault) {
if (columnDefault == null) {
return "";
}
String trimmed = columnDefault.trim();
if (STRING_LITERAL.matcher(trimmed).matches()
|| NUMERIC_LITERAL.matcher(trimmed).matches()
|| SIMPLE_CONSTANT.matcher(trimmed).matches()
|| CURRENT_TEMPORAL.matcher(trimmed).matches()
|| SAFE_NO_ARG_FUNCTION.matcher(trimmed).matches()
|| TYPED_LITERAL.matcher(trimmed).matches()
|| BINARY_LITERAL.matcher(trimmed).matches()
|| SEQUENCE_EXPRESSION.matcher(trimmed).matches()) {
return trimmed;
}
throw new IllegalArgumentException("Unsafe H2 column default from metadata: " + columnDefault);
}
}
Loading
Loading