diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/pom.xml b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/pom.xml index 14e1ba688b..9da6701192 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/pom.xml +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/pom.xml @@ -35,6 +35,11 @@ ai.chat2db chat2db-community-tdengine + + org.junit.jupiter + junit-jupiter + test + diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericDBManager.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericDBManager.java index fa38eb44e4..2a31ed0015 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericDBManager.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericDBManager.java @@ -1,6 +1,8 @@ package ai.chat2db.plugin.generic; import ai.chat2db.spi.IDbManager; +import ai.chat2db.community.domain.api.config.DBConfig; +import ai.chat2db.community.domain.api.constant.DBConfigConstants; import ai.chat2db.spi.DefaultDBManager; import ai.chat2db.spi.sql.Chat2DBContext; import ai.chat2db.spi.DefaultSQLExecutor; @@ -17,7 +19,13 @@ public void connectDatabase(Connection connection, String database) { return; } String schema = Chat2DBContext.getConnectInfo().getSchemaName(); - String changeDatabase = Chat2DBContext.getDBConfig().getChangeDatabase(database, schema); + DBConfig dbConfig = Chat2DBContext.getDBConfig(); + String template = dbConfig.getSql(DBConfigConstants.SQL_CHANGE_DATABASE); + if (template != null) { + database = GenericSqlGuards.sanitizeTemplateValue(template, "{database}", database); + schema = GenericSqlGuards.sanitizeTemplateValue(template, "{schema}", schema); + } + String changeDatabase = dbConfig.getChangeDatabase(database, schema); if(StringUtils.isEmpty(changeDatabase)){ return; } diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericMetaData.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericMetaData.java index 33c75af6c7..cee6a5d2cc 100644 --- a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericMetaData.java +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericMetaData.java @@ -1,8 +1,11 @@ package ai.chat2db.plugin.generic; +import ai.chat2db.plugin.generic.identifier.GenericIdentifierProcessor; import ai.chat2db.spi.IDbMetaData; import ai.chat2db.community.domain.api.config.DBConfig; +import ai.chat2db.community.domain.api.constant.DBConfigConstants; import ai.chat2db.spi.DefaultMetaService; +import ai.chat2db.spi.ISQLIdentifierProcessor; import ai.chat2db.community.domain.api.model.account.*; import ai.chat2db.community.domain.api.model.async.*; import ai.chat2db.community.domain.api.config.*; @@ -23,10 +26,24 @@ public class GenericMetaData extends DefaultMetaService implements IDbMetaData { + @Override + public ISQLIdentifierProcessor getSQLIdentifierProcessor() { + return GenericIdentifierProcessor.INSTANCE; + } + @Override public String tableDDL(Connection connection, String databaseName, String schemaName, String tableName) { DBConfig dbConfig = Chat2DBContext.getDBConfig(); - String sql = dbConfig.getTableDdl(databaseName, schemaName, tableName); + String template = dbConfig.getSql(DBConfigConstants.SQL_TABLE_DDL); + String templateDatabaseName = databaseName; + String templateSchemaName = schemaName; + String templateTableName = tableName; + if (template != null) { + templateDatabaseName = GenericSqlGuards.sanitizeTemplateValue(template, "{database}", databaseName); + templateSchemaName = GenericSqlGuards.sanitizeTemplateValue(template, "{schema}", schemaName); + templateTableName = GenericSqlGuards.sanitizeTemplateValue(template, "{table}", tableName); + } + String sql = dbConfig.getTableDdl(templateDatabaseName, templateSchemaName, templateTableName); String sqlResult = dbConfig.getTableDdlResult(); String ddl = null; if (sql != null && sqlResult!=null) { diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericSqlGuards.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericSqlGuards.java new file mode 100644 index 0000000000..511de687d7 --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/GenericSqlGuards.java @@ -0,0 +1,50 @@ +package ai.chat2db.plugin.generic; + +import ai.chat2db.plugin.generic.identifier.GenericIdentifierProcessor; +import org.apache.commons.lang3.StringUtils; + +import java.util.regex.Pattern; + +/** + * Validation helpers for values substituted into generic adapter SQL templates + * (generic.json sqlMap) (#1914). + * + * The generic adapter serves mixed dialects via DBConfig templates (e.g. DuckDB wraps + * placeholders in single quotes, TDengine uses bare identifier positions), so treatment + * is chosen per placeholder by inspecting the template; no single dialect quote char is + * hard-coded. Escaping itself lives in {@link GenericIdentifierProcessor}. + */ +public final class GenericSqlGuards { + + private static final Pattern SAFE_IDENTIFIER_PATTERN = Pattern.compile("^[A-Za-z0-9_$]+$"); + + private GenericSqlGuards() { + } + + /** + * Validate a strict identifier token for bare-identifier template positions, where the + * generic adapter cannot know the dialect's identifier quote char. + */ + public static String requireSafeIdentifier(String value, String what) { + if (value == null || !SAFE_IDENTIFIER_PATTERN.matcher(value).matches()) { + throw new IllegalArgumentException("Invalid generic " + what + ": " + value); + } + return value; + } + + /** + * Sanitize a value that DBConfig substitutes for {@code placeholder} in the given + * generic.json SQL template. A placeholder wrapped in single quotes ('{database}') + * lands in string-literal position and gets literal escaping; a bare placeholder + * ({database}) lands in identifier position and must pass the identifier whitelist. + */ + public static String sanitizeTemplateValue(String template, String placeholder, String value) { + if (template == null || StringUtils.isBlank(value)) { + return value; + } + if (template.contains("'" + placeholder + "'")) { + return GenericIdentifierProcessor.INSTANCE.escapeString(value); + } + return requireSafeIdentifier(value, placeholder); + } +} diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/identifier/GenericIdentifierProcessor.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/identifier/GenericIdentifierProcessor.java new file mode 100644 index 0000000000..78e15999ec --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/main/java/ai/chat2db/plugin/generic/identifier/GenericIdentifierProcessor.java @@ -0,0 +1,112 @@ +package ai.chat2db.plugin.generic.identifier; + +import ai.chat2db.spi.DefaultSQLIdentifierProcessor; +import org.apache.commons.lang3.StringUtils; + +/** + * Generic dialect identifier processor. The SPI-facing {@link #quoteIdentifier(String)} + * is conditional: identifiers that are already valid plain identifiers (and not + * reserved keywords) are returned unquoted so completion/matching consumers keep + * working; anything else is wrapped in double quotes (ANSI default) with embedded + * quotes doubled. Call sites that historically always quoted use + * {@link #quoteIdentifierAlways(String)} (or the SPI always-quote variant + * {@link #quoteIdentifierIgnoreCase(String)}). The generic adapter serves mixed + * dialects via DBConfig templates, so a dialect-parameterized + * {@link #quoteIdentifier(String, char)} variant is also provided for call sites that + * know the target quote char. String literals are escaped by doubling single quotes. + * Shared stateless instance available via {@link #INSTANCE}. + */ +public class GenericIdentifierProcessor extends DefaultSQLIdentifierProcessor { + + public static final GenericIdentifierProcessor INSTANCE = new GenericIdentifierProcessor(); + + /** + * Conditional quoting for SPI/completion paths: null/blank pass through; valid + * plain identifiers that are not reserved keywords are returned unquoted; + * everything else is double-quoted like {@link #quoteIdentifierAlways}. + */ + @Override + public String quoteIdentifier(String identifier) { + if (StringUtils.isBlank(identifier)) { + return identifier; + } + if (isValidIdentifier(identifier) && !isReservedKeyword(identifier.toUpperCase(), null, null)) { + return identifier; + } + return quoteIdentifierAlways(identifier); + } + + @Override + public String quoteIdentifier(String identifier, Integer majorVersion, Integer minorVersion) { + return quoteIdentifier(identifier); + } + + /** + * SPI always-quote variant (preserve case, always quote). + */ + @Override + public String quoteIdentifierIgnoreCase(String identifier) { + return quoteIdentifierAlways(identifier); + } + + /** + * Unconditional double-quote wrapping. Every quote in the raw identifier, + * including boundary quotes, is treated as identifier content so the SPI + * always-quote/remove-quote round-trip contract is preserved. + */ + @Override + public String quoteIdentifierAlways(String identifier) { + return super.quoteIdentifierAlways(identifier); + } + + /** + * Escapes a value interpolated into a single-quoted SQL string literal by + * doubling every single quote. + */ + @Override + public String escapeString(String str) { + return str == null ? null : StringUtils.replace(str, "'", "''"); + } + + /** + * Escapes identifier content for positions already inside quoted templates: + * strips one surrounding pair of the given quote char, then doubles every + * embedded quote char. + */ + public static String escapeIdentifier(String identifier) { + return escapeIdentifier(identifier, '"'); + } + + /** + * Dialect-parameterized content escaping for positions already inside quoted + * templates. + */ + public static String escapeIdentifier(String identifier, char quote) { + if (identifier == null) { + return ""; + } + String q = String.valueOf(quote); + String stripped = identifier; + if (stripped.length() >= 2 && stripped.startsWith(q) && stripped.endsWith(q)) { + stripped = stripped.substring(1, stripped.length() - 1); + } + return StringUtils.replace(stripped, q, q + q); + } + + /** + * Quotes an identifier with the given dialect quote char: strips one surrounding + * pair of that quote, then doubles every embedded quote char. Blank input is + * returned unchanged. + */ + public static String quoteIdentifier(String name, char quote) { + if (StringUtils.isBlank(name)) { + return name; + } + String q = String.valueOf(quote); + String identifier = name; + if (identifier.length() >= 2 && identifier.startsWith(q) && identifier.endsWith(q)) { + identifier = identifier.substring(1, identifier.length() - 1); + } + return q + identifier.replace(q, q + q) + q; + } +} diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/test/java/ai/chat2db/plugin/generic/GenericIdentifierProcessorTest.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/test/java/ai/chat2db/plugin/generic/GenericIdentifierProcessorTest.java new file mode 100644 index 0000000000..3d132fb39e --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/test/java/ai/chat2db/plugin/generic/GenericIdentifierProcessorTest.java @@ -0,0 +1,83 @@ +package ai.chat2db.plugin.generic; + +import ai.chat2db.plugin.generic.identifier.GenericIdentifierProcessor; +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; + +class GenericIdentifierProcessorTest { + + @Test + void escapeStringDoublesSingleQuotes() { + assertNull(GenericIdentifierProcessor.INSTANCE.escapeString(null)); + assertEquals("plain", GenericIdentifierProcessor.INSTANCE.escapeString("plain")); + assertEquals("a''b", GenericIdentifierProcessor.INSTANCE.escapeString("a'b")); + assertEquals("''; DROP TABLE x; --", + GenericIdentifierProcessor.INSTANCE.escapeString("'; DROP TABLE x; --")); + } + + @Test + void quoteIdentifierDoublesDialectQuoteChar() { + // DuckDB-style double quotes + assertEquals("\"plain\"", GenericIdentifierProcessor.quoteIdentifier("plain", '"')); + assertEquals("\"a\"\"b\"", GenericIdentifierProcessor.quoteIdentifier("a\"b", '"')); + assertEquals("\"we\"\"\"\"ird\"", GenericIdentifierProcessor.quoteIdentifier("\"we\"\"ird\"", '"')); + // TDengine-style backticks + assertEquals("`a``b`", GenericIdentifierProcessor.quoteIdentifier("a`b", '`')); + assertEquals("`a``; DROP TABLE b; --`", + GenericIdentifierProcessor.quoteIdentifier("a`; DROP TABLE b; --", '`')); + } + + @Test + void quoteIdentifierIsConditionalForSpiConsumers() { + // null/blank pass through + assertNull(GenericIdentifierProcessor.INSTANCE.quoteIdentifier(null)); + assertEquals("", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("")); + assertEquals(" ", GenericIdentifierProcessor.INSTANCE.quoteIdentifier(" ")); + // valid plain identifiers are returned unquoted + assertEquals("plain", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("plain")); + assertEquals("Plain_Case1", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("Plain_Case1")); + // anything else is wrapped with embedded quotes doubled + assertEquals("\"a\"\"b\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("a\"b")); + assertEquals("\"with space\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("with space")); + assertEquals("\"1abc\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("1abc")); + // versioned overload delegates to the conditional variant + assertEquals("plain", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("plain", null, null)); + assertEquals("\"a\"\"b\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifier("a\"b", null, null)); + } + + @Test + void quoteIdentifierAlwaysWrapsUnconditionally() { + assertNull(GenericIdentifierProcessor.INSTANCE.quoteIdentifierAlways(null)); + assertEquals("\"\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifierAlways("")); + assertEquals("\" \"", GenericIdentifierProcessor.INSTANCE.quoteIdentifierAlways(" ")); + assertEquals("\"plain\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifierAlways("plain")); + assertEquals("\"a\"\"b\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifierAlways("a\"b")); + assertEquals("\"\"\"we\"\"\"\"ird\"\"\"", + GenericIdentifierProcessor.INSTANCE.quoteIdentifierAlways("\"we\"\"ird\"")); + } + + @Test + void alwaysQuoteAndRemoveRoundTripRawIdentifiers() { + for (String raw : new String[] {"", " ", "plain", "\"", "\"edge", "edge\"", "\"quoted\"", "a\"\"b"}) { + assertEquals(raw, GenericIdentifierProcessor.INSTANCE.removeIdentifierQuote( + GenericIdentifierProcessor.INSTANCE.quoteIdentifierAlways(raw)), raw); + } + } + + @Test + void quoteIdentifierIgnoreCaseIsTheAlwaysQuoteVariant() { + assertNull(GenericIdentifierProcessor.INSTANCE.quoteIdentifierIgnoreCase(null)); + assertEquals("\"plain\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifierIgnoreCase("plain")); + assertEquals("\"a\"\"b\"", GenericIdentifierProcessor.INSTANCE.quoteIdentifierIgnoreCase("a\"b")); + } + + @Test + void escapeIdentifierStripsSurroundingPairAndDoublesEmbeddedQuotes() { + assertEquals("a\"\"b", GenericIdentifierProcessor.escapeIdentifier("a\"b")); + assertEquals("we\"\"\"\"ird", GenericIdentifierProcessor.escapeIdentifier("\"we\"\"ird\"")); + assertEquals("a``b", GenericIdentifierProcessor.escapeIdentifier("a`b", '`')); + assertEquals("", GenericIdentifierProcessor.escapeIdentifier(null)); + } +} diff --git a/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/test/java/ai/chat2db/plugin/generic/GenericSqlGuardsTest.java b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/test/java/ai/chat2db/plugin/generic/GenericSqlGuardsTest.java new file mode 100644 index 0000000000..4636bf29c7 --- /dev/null +++ b/chat2db-community-server/chat2db-community-plugins/chat2db-community-generic/src/test/java/ai/chat2db/plugin/generic/GenericSqlGuardsTest.java @@ -0,0 +1,115 @@ +package ai.chat2db.plugin.generic; + +import ai.chat2db.community.domain.api.config.DBConfig; +import ai.chat2db.community.domain.api.constant.DBConfigConstants; +import org.junit.jupiter.api.Test; + +import java.util.HashMap; +import java.util.Map; + +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 GenericSqlGuardsTest { + + private static final String DUCKDB_TABLE_DDL_TEMPLATE = + "select sql from duckdb_tables() where database_name = '{database}' and schema_name = '{schema}' and table_name = '{table}'"; + private static final String TDENGINE_TABLE_DDL_TEMPLATE = "SHOW CREATE TABLE {database}.{table}"; + private static final String TDENGINE_CHANGE_DATABASE_TEMPLATE = "USE {database}"; + + @Test + void requireSafeIdentifierAcceptsStrictNames() { + assertEquals("db1", GenericSqlGuards.requireSafeIdentifier("db1", "database")); + assertEquals("_sys", GenericSqlGuards.requireSafeIdentifier("_sys", "schema")); + assertEquals("a$B", GenericSqlGuards.requireSafeIdentifier("a$B", "table")); + assertEquals("T2$x_y", GenericSqlGuards.requireSafeIdentifier("T2$x_y", "table")); + } + + @Test + void requireSafeIdentifierRejectsUnsafeNames() { + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("a b", "table")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("a'b", "table")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("a\"b", "table")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("a`b", "table")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("a;b", "table")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("a.b", "table")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("db; SHUTDOWN", "database")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier("", "database")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.requireSafeIdentifier(null, "database")); + } + + @Test + void sanitizeTemplateValueTreatsQuotedPlaceholderAsLiteral() { + assertEquals("x'' OR ''1''=''1", + GenericSqlGuards.sanitizeTemplateValue(DUCKDB_TABLE_DDL_TEMPLATE, "{table}", "x' OR '1'='1")); + } + + @Test + void sanitizeTemplateValueTreatsBarePlaceholderAsIdentifier() { + assertEquals("test", + GenericSqlGuards.sanitizeTemplateValue(TDENGINE_CHANGE_DATABASE_TEMPLATE, "{database}", "test")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.sanitizeTemplateValue(TDENGINE_CHANGE_DATABASE_TEMPLATE, "{database}", + "test; SHUTDOWN")); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.sanitizeTemplateValue(TDENGINE_TABLE_DDL_TEMPLATE, "{table}", "a b")); + } + + @Test + void sanitizeTemplateValuePassesThroughBlankAndNullTemplate() { + assertEquals("a b", GenericSqlGuards.sanitizeTemplateValue(null, "{table}", "a b")); + assertNull(GenericSqlGuards.sanitizeTemplateValue(DUCKDB_TABLE_DDL_TEMPLATE, "{table}", null)); + assertEquals("", GenericSqlGuards.sanitizeTemplateValue(DUCKDB_TABLE_DDL_TEMPLATE, "{table}", "")); + } + + @Test + void duckdbTableDdlNeutralizesMaliciousLiteral() { + DBConfig config = configWith(DBConfigConstants.SQL_TABLE_DDL, DUCKDB_TABLE_DDL_TEMPLATE); + String template = config.getSql(DBConfigConstants.SQL_TABLE_DDL); + String databaseName = GenericSqlGuards.sanitizeTemplateValue(template, "{database}", "main"); + String schemaName = GenericSqlGuards.sanitizeTemplateValue(template, "{schema}", "main"); + String tableName = GenericSqlGuards.sanitizeTemplateValue(template, "{table}", "x' OR '1'='1"); + assertEquals("select sql from duckdb_tables() where database_name = 'main' and schema_name = 'main'" + + " and table_name = 'x'' OR ''1''=''1'", + config.getTableDdl(databaseName, schemaName, tableName)); + } + + @Test + void tdengineTableDdlRejectsMaliciousIdentifier() { + DBConfig config = configWith(DBConfigConstants.SQL_TABLE_DDL, TDENGINE_TABLE_DDL_TEMPLATE); + String template = config.getSql(DBConfigConstants.SQL_TABLE_DDL); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.sanitizeTemplateValue(template, "{database}", "test`; SHUTDOWN; --")); + String databaseName = GenericSqlGuards.sanitizeTemplateValue(template, "{database}", "test"); + String tableName = GenericSqlGuards.sanitizeTemplateValue(template, "{table}", "t1"); + assertEquals("SHOW CREATE TABLE test.t1", config.getTableDdl(databaseName, null, tableName)); + } + + @Test + void tdengineChangeDatabaseRejectsMaliciousIdentifier() { + DBConfig config = configWith(DBConfigConstants.SQL_CHANGE_DATABASE, TDENGINE_CHANGE_DATABASE_TEMPLATE); + String template = config.getSql(DBConfigConstants.SQL_CHANGE_DATABASE); + assertThrows(IllegalArgumentException.class, + () -> GenericSqlGuards.sanitizeTemplateValue(template, "{database}", "test; DROP DATABASE x")); + String database = GenericSqlGuards.sanitizeTemplateValue(template, "{database}", "test"); + assertEquals("USE test", config.getChangeDatabase(database, null)); + } + + private static DBConfig configWith(String key, String template) { + DBConfig config = new DBConfig(); + Map sqlMap = new HashMap<>(); + sqlMap.put(key, template); + config.setSqlMap(sqlMap); + return config; + } +}