Skip to content
Open
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
Expand Up @@ -15,6 +15,11 @@
<groupId>ai.chat2db</groupId>
<artifactId>chat2db-community-spi</artifactId>
</dependency>
<dependency>
<groupId>org.junit.jupiter</groupId>
<artifactId>junit-jupiter</artifactId>
<scope>test</scope>
</dependency>

</dependencies>

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,25 +29,26 @@ 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);
}
}

@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);
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ public PageResult<Table> 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<TableColumn> tableColumns = new ArrayList();
return (List) DefaultSQLExecutor.getInstance().execute(connection, sql, (resultSet) -> {
while (resultSet.next()) {
Expand All @@ -106,7 +106,7 @@ public List columns(Connection connection, String databaseName, String schemaNam
}

public List<TableIndex> 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<TableIndex> tableIndexes = new ArrayList<>();
DefaultSQLExecutor.getInstance().execute(connection, sql, resultSet -> {
while (resultSet.next()) {
Expand Down Expand Up @@ -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) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ public List<ExecuteResponse> 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);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -79,7 +79,8 @@ private String buildDeleteSql(String tableName, List<String> 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;

Expand All @@ -95,13 +96,13 @@ private String buildInsertSql(String tableName, List<Header> 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)
Expand All @@ -127,20 +128,20 @@ private String buildUpdate(String tableName, List<Header> 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)
Expand Down
Original file line number Diff line number Diff line change
@@ -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 <name>} or {@code db.<name>.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();
}
}
Original file line number Diff line number Diff line change
@@ -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);
}
}
Loading