diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/pom.xml b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/pom.xml index 70e02a35a..209f9b89d 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/pom.xml +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/pom.xml @@ -15,6 +15,11 @@ ai.chat2db chat2db-community-spi + + org.junit.jupiter + junit-jupiter + test + diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbDBManager.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbDBManager.java index 0990cd651..533fdb075 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbDBManager.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbDBManager.java @@ -29,7 +29,7 @@ public void connectDatabase(Connection connection, String database) { return; } try { - DefaultSQLExecutor.getInstance().execute(connection, String.format(SCRIPT_USE_SCHEMA, schemaName)); + DefaultSQLExecutor.getInstance().execute(connection, String.format(SCRIPT_USE_SCHEMA, MongodbSqlGuards.requireMongoName(schemaName, "database name"))); } catch (SQLException e) { throw new RuntimeException(e); } @@ -37,17 +37,18 @@ public void connectDatabase(Connection connection, String database) { @Override public String dropTable(Connection connection, String databaseName, String schemaName, String tableName) { - return String.format(SCRIPT_DROP_COLLECTION, tableName); + return String.format(SCRIPT_DROP_COLLECTION, MongodbSqlGuards.requireMongoName(tableName, "collection name")); } @Override public String truncateTable(Connection connection, String databaseName, String schemaName, String tableName) throws SQLException { - return String.format(SCRIPT_TRUNCATE_COLLECTION, tableName); + return String.format(SCRIPT_TRUNCATE_COLLECTION, MongodbSqlGuards.requireMongoName(tableName, "collection name")); } @Override public void copyTable(Connection connection, String databaseName, String schemaName, String tableName, String newTableName,boolean copyData) throws SQLException { - String sql = String.format(SCRIPT_COPY_COLLECTION, newTableName, tableName); + String sql = String.format(SCRIPT_COPY_COLLECTION, MongodbSqlGuards.requireMongoName(newTableName, "collection name"), + MongodbSqlGuards.requireMongoName(tableName, "collection name")); DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> null); } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbMetaData.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbMetaData.java index f1c37a565..ab7f7190e 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbMetaData.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbMetaData.java @@ -82,7 +82,7 @@ public PageResult tables(Connection connection, String databaseName, Stri @Override public List columns(Connection connection, String databaseName, String schemaName, String tableName) { - String sql = String.format(SELECT_TABLE_COLUMNS, tableName); + String sql = String.format(SELECT_TABLE_COLUMNS, MongodbSqlGuards.requireMongoName(tableName, "collection name")); List tableColumns = new ArrayList(); return (List) DefaultSQLExecutor.getInstance().execute(connection, sql, (resultSet) -> { while (resultSet.next()) { @@ -106,7 +106,7 @@ public List columns(Connection connection, String databaseName, String schemaNam } public List indexes(Connection connection, String databaseName, String schemaName, String tableName) { - String sql = String.format(SELECT_TABLE_INDEX, tableName); + String sql = String.format(SELECT_TABLE_INDEX, MongodbSqlGuards.requireMongoName(tableName, "collection name")); List tableIndexes = new ArrayList<>(); DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> { while (resultSet.next()) { @@ -141,7 +141,7 @@ private void executeUse(String schemaName) { return; } Connection connection = Chat2DBContext.getConnection(); - String sql = String.format(SCRIPT_USE_SCHEMA, schemaName); + String sql = String.format(SCRIPT_USE_SCHEMA, MongodbSqlGuards.requireMongoName(schemaName, "database name")); try { DefaultSQLExecutor.getInstance().execute(connection, sql); } catch (SQLException e) { diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbScriptExecutor.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbScriptExecutor.java index 6582e8502..5cc2a5258 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbScriptExecutor.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbScriptExecutor.java @@ -62,7 +62,7 @@ public List executeSelectTable(SqlExecuteRequest command) { if (StringUtils.isEmpty(command.getTableName())) { return Collections.emptyList(); } - String sql = String.format(EXECUTE_SQL, command.getTableName()); + String sql = String.format(EXECUTE_SQL, MongodbSqlGuards.requireMongoName(command.getTableName(), "collection name")); command.setScript(sql); return execute(command); } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbSqlBuilder.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbSqlBuilder.java index 4b0f5cf3b..dc5c2f075 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbSqlBuilder.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbSqlBuilder.java @@ -79,7 +79,8 @@ private String buildDeleteSql(String tableName, List deleteSqlCommands, if (StringUtils.isEmpty(idValue)) { return StringUtils.EMPTY; } - String sql = String.format(SQL_DB_DOT_FORMAT_DOT_DELETEONE_OPEN_PAREN_OPEN_BRACE, tableName, idValue); + String sql = String.format(SQL_DB_DOT_FORMAT_DOT_DELETEONE_OPEN_PAREN_OPEN_BRACE, + MongodbSqlGuards.requireMongoName(tableName, "collection name"), MongodbSqlGuards.escapeJsonString(idValue)); log.info(LOG_DELETE_SQL, sql); return sql; @@ -95,13 +96,13 @@ private String buildInsertSql(String tableName, List
headerList, ResultO for (int i = 2; i < newDataList.size(); i++) { Header header = headerList.get(i); String newValue = newDataList.get(i); - sql.append(header.getName()).append(SQLConstants.COLON).append(SQLConstants.DOUBLE_QUOTE).append(newValue).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.COMMA); + sql.append(MongodbSqlGuards.requireMongoName(header.getName(), "field name")).append(SQLConstants.COLON).append(SQLConstants.DOUBLE_QUOTE).append(MongodbSqlGuards.escapeJsonString(newValue)).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.COMMA); } if (sql.isEmpty()) { return StringUtils.EMPTY; } StringBuffer insertSql = new StringBuffer(); - insertSql.append(String.format(SQL_DB_DOT_FORMAT_DOT_INSERTONE, tableName)).append(SQLConstants.OPEN_PARENTHESIS) + insertSql.append(String.format(SQL_DB_DOT_FORMAT_DOT_INSERTONE, MongodbSqlGuards.requireMongoName(tableName, "collection name"))).append(SQLConstants.OPEN_PARENTHESIS) .append(SQLConstants.OPEN_CURLY_BRACE) .append(sql.deleteCharAt(sql.length() - 1)) .append(SQLConstants.CLOSE_CURLY_BRACE) @@ -127,20 +128,20 @@ private String buildUpdate(String tableName, List
headerList, ResultOper if (_idValue.isEmpty()) { _idValue.append(oldDataList.get(1)); } - setSql.append(header.getName()).append(SQLConstants.COLON).append(SQLConstants.DOUBLE_QUOTE).append(newValue).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.COMMA); + setSql.append(MongodbSqlGuards.requireMongoName(header.getName(), "field name")).append(SQLConstants.COLON).append(SQLConstants.DOUBLE_QUOTE).append(MongodbSqlGuards.escapeJsonString(newValue)).append(SQLConstants.DOUBLE_QUOTE).append(SQLConstants.COMMA); } if (_idValue.isEmpty() || setSql.isEmpty()) { return StringUtils.EMPTY; } StringBuffer sql = new StringBuffer(); - sql.append(String.format(SQL_DB_DOT_FORMAT_DOT_UPDATEONE, tableName)).append(SQLConstants.OPEN_PARENTHESIS) + sql.append(String.format(SQL_DB_DOT_FORMAT_DOT_UPDATEONE, MongodbSqlGuards.requireMongoName(tableName, "collection name"))).append(SQLConstants.OPEN_PARENTHESIS) .append(SQLConstants.OPEN_CURLY_BRACE) .append(MONGODB_ID_FIELD) .append(SQLConstants.COLON) .append(MONGODB_OBJECT_ID_TYPE) .append(SQLConstants.OPEN_PARENTHESIS) .append(SQLConstants.DOUBLE_QUOTE) - .append(_idValue.toString()) + .append(MongodbSqlGuards.escapeJsonString(_idValue.toString())) .append(SQLConstants.DOUBLE_QUOTE) .append(SQLConstants.CLOSE_PARENTHESIS) .append(SQLConstants.CLOSE_CURLY_BRACE) diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbSqlGuards.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbSqlGuards.java new file mode 100644 index 000000000..398519e33 --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/main/java/ai/chat2db/plugin/mongodb/MongodbSqlGuards.java @@ -0,0 +1,71 @@ +package ai.chat2db.plugin.mongodb; + +import java.util.regex.Pattern; + +/** + * Validation and JSON-escaping guards for values interpolated into Mongo shell command + * text (#1914). Shell commands are not plain SQL: database/collection names are validated + * against a strict allowlist, and string values inside JSON documents are JSON-escaped. + */ +public final class MongodbSqlGuards { + + private static final Pattern MONGO_NAME_PATTERN = Pattern.compile("^[A-Za-z0-9_$-]+$"); + + private MongodbSqlGuards() { + } + + /** + * Validate a database/collection/field name interpolated into shell command text such as + * {@code use } or {@code db..find()}. Rejects anything outside the allowlist. + */ + public static String requireMongoName(String name, String what) { + if (name == null || !MONGO_NAME_PATTERN.matcher(name).matches()) { + throw new IllegalArgumentException("Invalid MongoDB " + what + ": " + name); + } + return name; + } + + /** + * Escape a value interpolated into a double-quoted JSON string inside a shell command + * (surrounding quotes NOT added). + */ + public static String escapeJsonString(String value) { + if (value == null) { + return null; + } + StringBuilder sb = new StringBuilder(value.length()); + for (int i = 0; i < value.length(); i++) { + char c = value.charAt(i); + switch (c) { + case '\\': + sb.append("\\\\"); + break; + case '"': + sb.append("\\\""); + break; + case '\n': + sb.append("\\n"); + break; + case '\r': + sb.append("\\r"); + break; + case '\t': + sb.append("\\t"); + break; + case '\b': + sb.append("\\b"); + break; + case '\f': + sb.append("\\f"); + break; + default: + if (c < 0x20) { + sb.append(String.format("\\u%04x", (int) c)); + } else { + sb.append(c); + } + } + } + return sb.toString(); + } +} diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/test/java/ai/chat2db/plugin/mongodb/MongodbSqlGuardsTest.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/test/java/ai/chat2db/plugin/mongodb/MongodbSqlGuardsTest.java new file mode 100644 index 000000000..c4e9939d9 --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-mongodb/src/test/java/ai/chat2db/plugin/mongodb/MongodbSqlGuardsTest.java @@ -0,0 +1,128 @@ +package ai.chat2db.plugin.mongodb; + +import java.util.List; + +import ai.chat2db.community.domain.api.model.result.Header; +import ai.chat2db.community.domain.api.model.result.QueryResponse; +import ai.chat2db.community.domain.api.model.result.ResultOperation; +import ai.chat2db.spi.constant.SQLConstants; +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; + +class MongodbSqlGuardsTest { + + @Test + void requireMongoNameAcceptsLegitimateNames() { + assertEquals("mydb", MongodbSqlGuards.requireMongoName("mydb", "database name")); + assertEquals("my_db-1$A", MongodbSqlGuards.requireMongoName("my_db-1$A", "collection name")); + } + + @Test + void requireMongoNameRejectsInjection() { + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlGuards.requireMongoName("a.b", "collection name")); + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlGuards.requireMongoName("a b", "collection name")); + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlGuards.requireMongoName("x; db.dropDatabase(); //", "database name")); + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlGuards.requireMongoName("x\")}", "collection name")); + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlGuards.requireMongoName("", "collection name")); + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlGuards.requireMongoName(null, "collection name")); + } + + @Test + void escapeJsonStringEscapesQuotesBackslashAndControls() { + assertEquals("plain", MongodbSqlGuards.escapeJsonString("plain")); + assertEquals("a\\\"b", MongodbSqlGuards.escapeJsonString("a\"b")); + assertEquals("a\\\\b", MongodbSqlGuards.escapeJsonString("a\\b")); + assertEquals("a\\nb", MongodbSqlGuards.escapeJsonString("a\nb")); + assertEquals("\\u0001", MongodbSqlGuards.escapeJsonString("\u0001")); + assertNull(MongodbSqlGuards.escapeJsonString(null)); + } + + @Test + void dropTableRejectsMaliciousName() { + MongodbDBManager manager = new MongodbDBManager(); + assertEquals(" db. users.drop();", manager.dropTable(null, null, null, "users")); + assertThrows(IllegalArgumentException.class, + () -> manager.dropTable(null, null, null, "users; db.dropDatabase(); //")); + } + + @Test + void truncateTableRejectsMaliciousName() throws Exception { + MongodbDBManager manager = new MongodbDBManager(); + assertEquals("db.users.deleteMany({})", manager.truncateTable(null, null, null, "users")); + assertThrows(IllegalArgumentException.class, + () -> manager.truncateTable(null, null, null, "a.b")); + } + + @Test + void deleteCommandEscapesObjectIdAndValidatesTable() { + QueryResponse response = new QueryResponse(); + response.setTableName("users"); + response.setHeaderList(List.of( + Header.builder().name("rn").build(), + Header.builder().name("_id").build())); + ResultOperation operation = new ResultOperation(); + operation.setType(SQLConstants.DELETE_KEYWORD); + operation.setOldDataList(List.of("1", "abc123")); + response.setOperations(List.of(operation)); + assertEquals("db.users.deleteOne({_id: ObjectId(\"abc123\")})", + MongodbSqlBuilder.getInstance().buildByQueryResult(response)); + + operation.setOldDataList(List.of("1", "a\"), $where: 1, x: (\"")); + String sql = MongodbSqlBuilder.getInstance().buildByQueryResult(response); + assertEquals("db.users.deleteOne({_id: ObjectId(\"a\\\"), $where: 1, x: (\\\"\")})", sql); + + response.setTableName("users; //"); + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlBuilder.getInstance().buildByQueryResult(response)); + } + + @Test + void insertCommandEscapesValuesAndValidatesNames() { + QueryResponse response = new QueryResponse(); + response.setTableName("users"); + response.setHeaderList(List.of( + Header.builder().name("rn").build(), + Header.builder().name("_id").build(), + Header.builder().name("name").build())); + ResultOperation operation = new ResultOperation(); + operation.setType(SQLConstants.CREATE_KEYWORD); + operation.setDataList(List.of("1", "id1", "a\"}), db.dropDatabase(), //")); + response.setOperations(List.of(operation)); + String sql = MongodbSqlBuilder.getInstance().buildByQueryResult(response); + assertEquals("db.users.insertOne({name:\"a\\\"}), db.dropDatabase(), //\"})", sql); + + // malicious field name is rejected + response.setHeaderList(List.of( + Header.builder().name("rn").build(), + Header.builder().name("_id").build(), + Header.builder().name("a}), x:(").build())); + assertThrows(IllegalArgumentException.class, + () -> MongodbSqlBuilder.getInstance().buildByQueryResult(response)); + } + + @Test + void updateCommandEscapesValuesAndId() { + QueryResponse response = new QueryResponse(); + response.setTableName("users"); + response.setHeaderList(List.of( + Header.builder().name("rn").build(), + Header.builder().name("_id").build(), + Header.builder().name("name").build())); + ResultOperation operation = new ResultOperation(); + operation.setType(SQLConstants.UPDATE_KEYWORD); + operation.setOldDataList(List.of("1", "id\"1", "old")); + operation.setDataList(List.of("1", "id\"1", "new\"value")); + response.setOperations(List.of(operation)); + String sql = MongodbSqlBuilder.getInstance().buildByQueryResult(response); + assertEquals("db.users.updateOne({_id:ObjectId(\"id\\\"1\")},{$set:{name:\"new\\\"value\"}})", sql); + } +}