This is an automated email from the ASF dual-hosted git repository.

chengzhang pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/shardingsphere.git


The following commit(s) were added to refs/heads/master by this push:
     new dcf1aca1125 Adjust protocolType to sqlParserType in SQLTranslator 
(#28906)
dcf1aca1125 is described below

commit dcf1aca1125a94b789fad710387380a56de19b58
Author: Zhengqiang Duan <[email protected]>
AuthorDate: Tue Oct 31 19:30:16 2023 +0800

    Adjust protocolType to sqlParserType in SQLTranslator (#28906)
---
 .../infra/rewrite/SQLRewriteEntryTest.java         |  1 +
 .../engine/GenericSQLRewriteEngineTest.java        | 17 +++++++--
 .../rewrite/engine/RouteSQLRewriteEngineTest.java  | 34 +++++++++++++----
 .../sqltranslator/rule/SQLTranslatorRule.java      |  4 +-
 .../sqltranslator/rule/SQLTranslatorRuleTest.java  | 44 +++++++++++-----------
 .../sqltranslator/jooq/JooQSQLTranslator.java      |  2 +-
 6 files changed, 65 insertions(+), 37 deletions(-)

diff --git 
a/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/SQLRewriteEntryTest.java
 
b/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/SQLRewriteEntryTest.java
index 2fb3eacfba9..cf42ec961fb 100644
--- 
a/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/SQLRewriteEntryTest.java
+++ 
b/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/SQLRewriteEntryTest.java
@@ -70,6 +70,7 @@ class SQLRewriteEntryTest {
         when(result.getSql()).thenReturn("SELECT ?");
         when(result.getParameters()).thenReturn(Collections.singletonList(1));
         CommonSQLStatementContext sqlStatementContext = 
mock(CommonSQLStatementContext.class);
+        
when(sqlStatementContext.getDatabaseType()).thenReturn(TypedSPILoader.getService(DatabaseType.class,
 "H2"));
         when(result.getSqlStatementContext()).thenReturn(sqlStatementContext);
         when(result.getHintValueContext()).thenReturn(new HintValueContext());
         return result;
diff --git 
a/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/GenericSQLRewriteEngineTest.java
 
b/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/GenericSQLRewriteEngineTest.java
index 8d3512e9bdf..a753669d530 100644
--- 
a/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/GenericSQLRewriteEngineTest.java
+++ 
b/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/GenericSQLRewriteEngineTest.java
@@ -52,9 +52,13 @@ class GenericSQLRewriteEngineTest {
         when(database.getProtocolType()).thenReturn(databaseType);
         Map<String, StorageUnit> storageUnits = mockStorageUnits(databaseType);
         
when(database.getResourceMetaData().getStorageUnits()).thenReturn(storageUnits);
+        CommonSQLStatementContext sqlStatementContext = 
mock(CommonSQLStatementContext.class);
+        when(sqlStatementContext.getDatabaseType()).thenReturn(databaseType);
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(sqlStatementContext);
         GenericSQLRewriteResult actual = new GenericSQLRewriteEngine(rule, 
database, mock(RuleMetaData.class))
-                .rewrite(new SQLRewriteContext(database, 
mock(CommonSQLStatementContext.class), "SELECT 1", Collections.emptyList(), 
mock(ConnectionContext.class),
-                        new HintValueContext()), mock(QueryContext.class));
+                .rewrite(new SQLRewriteContext(database, sqlStatementContext, 
"SELECT 1", Collections.emptyList(), mock(ConnectionContext.class),
+                        new HintValueContext()), queryContext);
         assertThat(actual.getSqlRewriteUnit().getSql(), is("SELECT 1"));
         assertThat(actual.getSqlRewriteUnit().getParameters(), 
is(Collections.emptyList()));
     }
@@ -66,9 +70,14 @@ class GenericSQLRewriteEngineTest {
         when(database.getName()).thenReturn(DefaultDatabase.LOGIC_NAME);
         
when(database.getSchemas()).thenReturn(Collections.singletonMap("test", 
mock(ShardingSphereSchema.class)));
         
when(database.getResourceMetaData().getStorageUnits()).thenReturn(Collections.emptyMap());
+        CommonSQLStatementContext sqlStatementContext = 
mock(CommonSQLStatementContext.class);
+        DatabaseType databaseType = mock(DatabaseType.class);
+        when(sqlStatementContext.getDatabaseType()).thenReturn(databaseType);
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(sqlStatementContext);
         GenericSQLRewriteResult actual = new GenericSQLRewriteEngine(rule, 
database, mock(RuleMetaData.class))
-                .rewrite(new SQLRewriteContext(database, 
mock(CommonSQLStatementContext.class), "SELECT 1", Collections.emptyList(), 
mock(ConnectionContext.class),
-                        new HintValueContext()), mock(QueryContext.class));
+                .rewrite(new SQLRewriteContext(database, sqlStatementContext, 
"SELECT 1", Collections.emptyList(), mock(ConnectionContext.class),
+                        new HintValueContext()), queryContext);
         assertThat(actual.getSqlRewriteUnit().getSql(), is("SELECT 1"));
         assertThat(actual.getSqlRewriteUnit().getParameters(), 
is(Collections.emptyList()));
     }
diff --git 
a/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/RouteSQLRewriteEngineTest.java
 
b/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/RouteSQLRewriteEngineTest.java
index ad83303895c..f050e343883 100644
--- 
a/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/RouteSQLRewriteEngineTest.java
+++ 
b/infra/rewrite/src/test/java/org/apache/shardingsphere/infra/rewrite/engine/RouteSQLRewriteEngineTest.java
@@ -57,13 +57,16 @@ class RouteSQLRewriteEngineTest {
     void assertRewriteWithStandardParameterBuilder() {
         DatabaseType databaseType = mock(DatabaseType.class);
         ShardingSphereDatabase database = mockDatabase(databaseType);
-        SQLRewriteContext sqlRewriteContext =
-                new SQLRewriteContext(database, 
mock(CommonSQLStatementContext.class), "SELECT ?", 
Collections.singletonList(1), mock(ConnectionContext.class), new 
HintValueContext());
+        CommonSQLStatementContext sqlStatementContext = 
mock(CommonSQLStatementContext.class);
+        when(sqlStatementContext.getDatabaseType()).thenReturn(databaseType);
+        SQLRewriteContext sqlRewriteContext = new SQLRewriteContext(database, 
sqlStatementContext, "SELECT ?", Collections.singletonList(1), 
mock(ConnectionContext.class), new HintValueContext());
         RouteUnit routeUnit = new RouteUnit(new RouteMapper("ds", "ds_0"), 
Collections.singletonList(new RouteMapper("tbl", "tbl_0")));
         RouteContext routeContext = new RouteContext();
         routeContext.getRouteUnits().add(routeUnit);
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(sqlStatementContext);
         RouteSQLRewriteResult actual = new RouteSQLRewriteEngine(
-                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
mock(QueryContext.class));
+                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
queryContext);
         assertThat(actual.getSqlRewriteUnits().size(), is(1));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getSql(), 
is("SELECT ?"));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getParameters(), 
is(Collections.singletonList(1)));
@@ -85,6 +88,7 @@ class RouteSQLRewriteEngineTest {
         
when(statementContext.getOrderByContext().getItems()).thenReturn(Collections.emptyList());
         
when(statementContext.getPaginationContext().isHasPagination()).thenReturn(false);
         DatabaseType databaseType = mock(DatabaseType.class);
+        when(statementContext.getDatabaseType()).thenReturn(databaseType);
         ShardingSphereDatabase database = mockDatabase(databaseType);
         SQLRewriteContext sqlRewriteContext = new SQLRewriteContext(database, 
statementContext, "SELECT ?", Collections.singletonList(1), 
mock(ConnectionContext.class), new HintValueContext());
         RouteContext routeContext = new RouteContext();
@@ -92,8 +96,10 @@ class RouteSQLRewriteEngineTest {
         RouteUnit secondRouteUnit = new RouteUnit(new RouteMapper("ds", 
"ds_0"), Collections.singletonList(new RouteMapper("tbl", "tbl_1")));
         routeContext.getRouteUnits().add(firstRouteUnit);
         routeContext.getRouteUnits().add(secondRouteUnit);
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(statementContext);
         RouteSQLRewriteResult actual = new RouteSQLRewriteEngine(
-                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
mock(QueryContext.class));
+                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
queryContext);
         assertThat(actual.getSqlRewriteUnits().size(), is(1));
         assertThat(actual.getSqlRewriteUnits().get(firstRouteUnit).getSql(), 
is("SELECT ? UNION ALL SELECT ?"));
         
assertThat(actual.getSqlRewriteUnits().get(firstRouteUnit).getParameters(), 
is(Arrays.asList(1, 1)));
@@ -107,14 +113,17 @@ class RouteSQLRewriteEngineTest {
         
when(statementContext.getGroupedParameters()).thenReturn(Collections.singletonList(Collections.singletonList(1)));
         
when(statementContext.getOnDuplicateKeyUpdateParameters()).thenReturn(Collections.emptyList());
         DatabaseType databaseType = mock(DatabaseType.class);
+        when(statementContext.getDatabaseType()).thenReturn(databaseType);
         ShardingSphereDatabase database = mockDatabase(databaseType);
         SQLRewriteContext sqlRewriteContext =
                 new SQLRewriteContext(database, statementContext, "INSERT INTO 
tbl VALUES (?)", Collections.singletonList(1), mock(ConnectionContext.class), 
new HintValueContext());
         RouteUnit routeUnit = new RouteUnit(new RouteMapper("ds", "ds_0"), 
Collections.singletonList(new RouteMapper("tbl", "tbl_0")));
         RouteContext routeContext = new RouteContext();
         routeContext.getRouteUnits().add(routeUnit);
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(statementContext);
         RouteSQLRewriteResult actual = new RouteSQLRewriteEngine(
-                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
mock(QueryContext.class));
+                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
queryContext);
         assertThat(actual.getSqlRewriteUnits().size(), is(1));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getSql(), 
is("INSERT INTO tbl VALUES (?)"));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getParameters(), 
is(Collections.singletonList(1)));
@@ -128,6 +137,7 @@ class RouteSQLRewriteEngineTest {
         
when(statementContext.getGroupedParameters()).thenReturn(Collections.singletonList(Collections.singletonList(1)));
         
when(statementContext.getOnDuplicateKeyUpdateParameters()).thenReturn(Collections.emptyList());
         DatabaseType databaseType = mock(DatabaseType.class);
+        when(statementContext.getDatabaseType()).thenReturn(databaseType);
         ShardingSphereDatabase database = mockDatabase(databaseType);
         SQLRewriteContext sqlRewriteContext =
                 new SQLRewriteContext(database, statementContext, "INSERT INTO 
tbl VALUES (?)", Collections.singletonList(1), mock(ConnectionContext.class), 
new HintValueContext());
@@ -136,8 +146,10 @@ class RouteSQLRewriteEngineTest {
         routeContext.getRouteUnits().add(routeUnit);
         // TODO check why data node is "ds.tbl_0", not "ds_0.tbl_0"
         routeContext.getOriginalDataNodes().add(Collections.singletonList(new 
DataNode("ds.tbl_0")));
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(statementContext);
         RouteSQLRewriteResult actual = new RouteSQLRewriteEngine(
-                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
mock(QueryContext.class));
+                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
queryContext);
         assertThat(actual.getSqlRewriteUnits().size(), is(1));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getSql(), 
is("INSERT INTO tbl VALUES (?)"));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getParameters(), 
is(Collections.singletonList(1)));
@@ -151,6 +163,7 @@ class RouteSQLRewriteEngineTest {
         
when(statementContext.getGroupedParameters()).thenReturn(Collections.singletonList(Collections.singletonList(1)));
         
when(statementContext.getOnDuplicateKeyUpdateParameters()).thenReturn(Collections.emptyList());
         DatabaseType databaseType = mock(DatabaseType.class);
+        when(statementContext.getDatabaseType()).thenReturn(databaseType);
         ShardingSphereDatabase database = mockDatabase(databaseType);
         SQLRewriteContext sqlRewriteContext =
                 new SQLRewriteContext(database, statementContext, "INSERT INTO 
tbl VALUES (?)", Collections.singletonList(1), mock(ConnectionContext.class), 
new HintValueContext());
@@ -158,8 +171,10 @@ class RouteSQLRewriteEngineTest {
         RouteContext routeContext = new RouteContext();
         routeContext.getRouteUnits().add(routeUnit);
         routeContext.getOriginalDataNodes().add(Collections.emptyList());
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(statementContext);
         RouteSQLRewriteResult actual = new RouteSQLRewriteEngine(
-                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
mock(QueryContext.class));
+                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
queryContext);
         assertThat(actual.getSqlRewriteUnits().size(), is(1));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getSql(), 
is("INSERT INTO tbl VALUES (?)"));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getParameters(), 
is(Collections.singletonList(1)));
@@ -173,6 +188,7 @@ class RouteSQLRewriteEngineTest {
         
when(statementContext.getGroupedParameters()).thenReturn(Collections.singletonList(Collections.singletonList(1)));
         
when(statementContext.getOnDuplicateKeyUpdateParameters()).thenReturn(Collections.emptyList());
         DatabaseType databaseType = mock(DatabaseType.class);
+        when(statementContext.getDatabaseType()).thenReturn(databaseType);
         ShardingSphereDatabase database = mockDatabase(databaseType);
         SQLRewriteContext sqlRewriteContext =
                 new SQLRewriteContext(database, statementContext, "INSERT INTO 
tbl VALUES (?)", Collections.singletonList(1), mock(ConnectionContext.class), 
new HintValueContext());
@@ -180,8 +196,10 @@ class RouteSQLRewriteEngineTest {
         RouteContext routeContext = new RouteContext();
         routeContext.getRouteUnits().add(routeUnit);
         routeContext.getOriginalDataNodes().add(Collections.singletonList(new 
DataNode("ds_1.tbl_1")));
+        QueryContext queryContext = mock(QueryContext.class);
+        
when(queryContext.getSqlStatementContext()).thenReturn(statementContext);
         RouteSQLRewriteResult actual = new RouteSQLRewriteEngine(
-                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
mock(QueryContext.class));
+                new SQLTranslatorRule(new SQLTranslatorRuleConfiguration()), 
database, mock(RuleMetaData.class)).rewrite(sqlRewriteContext, routeContext, 
queryContext);
         assertThat(actual.getSqlRewriteUnits().size(), is(1));
         assertThat(actual.getSqlRewriteUnits().get(routeUnit).getSql(), 
is("INSERT INTO tbl VALUES (?)"));
         
assertTrue(actual.getSqlRewriteUnits().get(routeUnit).getParameters().isEmpty());
diff --git 
a/kernel/sql-translator/core/src/main/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRule.java
 
b/kernel/sql-translator/core/src/main/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRule.java
index 7f8163b6647..a852d79c10f 100644
--- 
a/kernel/sql-translator/core/src/main/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRule.java
+++ 
b/kernel/sql-translator/core/src/main/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRule.java
@@ -58,8 +58,8 @@ public final class SQLTranslatorRule implements GlobalRule {
      */
     public String translate(final String sql, final QueryContext queryContext, 
final DatabaseType storageType, final ShardingSphereDatabase database,
                             final RuleMetaData globalRuleMetaData) {
-        DatabaseType protocolType = database.getProtocolType();
-        if (protocolType.equals(storageType) || null == storageType) {
+        DatabaseType sqlParserType = 
queryContext.getSqlStatementContext().getDatabaseType();
+        if (sqlParserType.equals(storageType) || null == storageType) {
             return sql;
         }
         try {
diff --git 
a/kernel/sql-translator/core/src/test/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRuleTest.java
 
b/kernel/sql-translator/core/src/test/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRuleTest.java
index 8c482a04b91..6a725319e19 100644
--- 
a/kernel/sql-translator/core/src/test/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRuleTest.java
+++ 
b/kernel/sql-translator/core/src/test/java/org/apache/shardingsphere/sqltranslator/rule/SQLTranslatorRuleTest.java
@@ -31,6 +31,7 @@ import java.util.Locale;
 import static org.hamcrest.CoreMatchers.is;
 import static org.hamcrest.MatcherAssert.assertThat;
 import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.mockito.Mockito.RETURNS_DEEP_STUBS;
 import static org.mockito.Mockito.mock;
 import static org.mockito.Mockito.when;
 
@@ -40,9 +41,9 @@ class SQLTranslatorRuleTest {
     void assertTranslateWhenProtocolSameAsStorage() {
         String expected = "select 1";
         DatabaseType databaseType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
-        ShardingSphereDatabase database = mock(ShardingSphereDatabase.class);
-        when(database.getProtocolType()).thenReturn(databaseType);
-        String actual = new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("CONVERT_TO_UPPER_CASE", 
false)).translate(expected, mock(QueryContext.class), databaseType, database,
+        QueryContext queryContext = mock(QueryContext.class, 
RETURNS_DEEP_STUBS);
+        
when(queryContext.getSqlStatementContext().getDatabaseType()).thenReturn(databaseType);
+        String actual = new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("CONVERT_TO_UPPER_CASE", 
false)).translate(expected, queryContext, databaseType, 
mock(ShardingSphereDatabase.class),
                 mock(RuleMetaData.class));
         assertThat(actual, is(expected));
     }
@@ -50,10 +51,10 @@ class SQLTranslatorRuleTest {
     @Test
     void assertTranslateWhenNoStorage() {
         String expected = "select 1";
-        ShardingSphereDatabase database = mock(ShardingSphereDatabase.class);
-        DatabaseType protocolType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
-        when(database.getProtocolType()).thenReturn(protocolType);
-        String actual = new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("CONVERT_TO_UPPER_CASE", 
false)).translate(expected, mock(QueryContext.class), null, database,
+        DatabaseType sqlParserType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
+        QueryContext queryContext = mock(QueryContext.class, 
RETURNS_DEEP_STUBS);
+        
when(queryContext.getSqlStatementContext().getDatabaseType()).thenReturn(sqlParserType);
+        String actual = new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("CONVERT_TO_UPPER_CASE", 
false)).translate(expected, queryContext, null, 
mock(ShardingSphereDatabase.class),
                 mock(RuleMetaData.class));
         assertThat(actual, is(expected));
     }
@@ -61,11 +62,11 @@ class SQLTranslatorRuleTest {
     @Test
     void assertTranslateWithProtocolDifferentWithStorage() {
         String input = "select 1";
-        DatabaseType protocolType = 
TypedSPILoader.getService(DatabaseType.class, "MySQL");
-        ShardingSphereDatabase database = mock(ShardingSphereDatabase.class);
-        when(database.getProtocolType()).thenReturn(protocolType);
+        DatabaseType sqlParserType = 
TypedSPILoader.getService(DatabaseType.class, "MySQL");
+        QueryContext queryContext = mock(QueryContext.class, 
RETURNS_DEEP_STUBS);
+        
when(queryContext.getSqlStatementContext().getDatabaseType()).thenReturn(sqlParserType);
         DatabaseType storageType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
-        String actual = new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("CONVERT_TO_UPPER_CASE", 
false)).translate(input, mock(QueryContext.class), storageType, database,
+        String actual = new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("CONVERT_TO_UPPER_CASE", 
false)).translate(input, queryContext, storageType, 
mock(ShardingSphereDatabase.class),
                 mock(RuleMetaData.class));
         assertThat(actual, is(input.toUpperCase(Locale.ROOT)));
     }
@@ -73,24 +74,23 @@ class SQLTranslatorRuleTest {
     @Test
     void assertUseOriginalSQLWhenTranslatingFailed() {
         String expected = "select 1";
-        DatabaseType protocolType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
-        ShardingSphereDatabase database = mock(ShardingSphereDatabase.class);
-        when(database.getProtocolType()).thenReturn(protocolType);
+        DatabaseType sqlParserType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
+        QueryContext queryContext = mock(QueryContext.class, 
RETURNS_DEEP_STUBS);
+        
when(queryContext.getSqlStatementContext().getDatabaseType()).thenReturn(sqlParserType);
         DatabaseType storageType = 
TypedSPILoader.getService(DatabaseType.class, "MySQL");
-        String actual =
-                new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("ALWAYS_FAILED", true)).translate(expected, 
mock(QueryContext.class), storageType, database, mock(RuleMetaData.class));
+        String actual = new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("ALWAYS_FAILED", true)).translate(expected, 
queryContext, storageType, mock(ShardingSphereDatabase.class),
+                mock(RuleMetaData.class));
         assertThat(actual, is(expected));
     }
     
     @Test
     void assertNotUseOriginalSQLWhenTranslatingFailed() {
-        DatabaseType protocolType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
-        ShardingSphereDatabase database = mock(ShardingSphereDatabase.class);
-        when(database.getProtocolType()).thenReturn(protocolType);
+        QueryContext queryContext = mock(QueryContext.class, 
RETURNS_DEEP_STUBS);
+        DatabaseType sqlParserType = 
TypedSPILoader.getService(DatabaseType.class, "PostgreSQL");
+        
when(queryContext.getSqlStatementContext().getDatabaseType()).thenReturn(sqlParserType);
         DatabaseType storageType = 
TypedSPILoader.getService(DatabaseType.class, "MySQL");
-        assertThrows(UnsupportedTranslatedDatabaseException.class,
-                () -> new SQLTranslatorRule(new 
SQLTranslatorRuleConfiguration("ALWAYS_FAILED", false)).translate("", 
mock(QueryContext.class), storageType, database,
-                        mock(RuleMetaData.class)));
+        assertThrows(UnsupportedTranslatedDatabaseException.class, () -> new 
SQLTranslatorRule(new SQLTranslatorRuleConfiguration("ALWAYS_FAILED", 
false)).translate("", queryContext, storageType,
+                mock(ShardingSphereDatabase.class), mock(RuleMetaData.class)));
     }
     
     @Test
diff --git 
a/kernel/sql-translator/provider/jooq/src/main/java/org/apache/shardingsphere/sqltranslator/jooq/JooQSQLTranslator.java
 
b/kernel/sql-translator/provider/jooq/src/main/java/org/apache/shardingsphere/sqltranslator/jooq/JooQSQLTranslator.java
index 3797b3a1fe1..fdd136ba4d2 100644
--- 
a/kernel/sql-translator/provider/jooq/src/main/java/org/apache/shardingsphere/sqltranslator/jooq/JooQSQLTranslator.java
+++ 
b/kernel/sql-translator/provider/jooq/src/main/java/org/apache/shardingsphere/sqltranslator/jooq/JooQSQLTranslator.java
@@ -35,7 +35,7 @@ public final class JooQSQLTranslator implements SQLTranslator 
{
     public String translate(final String sql, final QueryContext queryContext, 
final DatabaseType storageType, final ShardingSphereDatabase database,
                             final RuleMetaData globalRuleMetaData) {
         try {
-            Query query = 
DSL.using(JooQDialectRegistry.getSQLDialect(database.getProtocolType())).parser().parseQuery(sql);
+            Query query = 
DSL.using(JooQDialectRegistry.getSQLDialect(queryContext.getSqlStatementContext().getDatabaseType())).parser().parseQuery(sql);
             return 
DSL.using(JooQDialectRegistry.getSQLDialect(storageType)).render(query);
             // CHECKSTYLE:OFF
         } catch (final Exception ignored) {

Reply via email to