diff --git a/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaData.java b/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaData.java index f1794a6f197ec..9b195c6605fe1 100644 --- a/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaData.java +++ b/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaData.java @@ -31,8 +31,13 @@ import org.apache.shardingsphere.sqlfederation.compiler.sql.type.SQLFederationDataTypeFactory; import org.apache.shardingsphere.sqlfederation.resultset.converter.DialectSQLFederationColumnTypeConverter; +import java.math.BigDecimal; import java.math.BigInteger; +import java.sql.Date; import java.sql.ResultSetMetaData; +import java.sql.Time; +import java.sql.Timestamp; +import java.sql.Types; import java.util.List; import java.util.Map; import java.util.Optional; @@ -180,7 +185,59 @@ public boolean isDefinitelyWritable(final int column) { @Override public String getColumnClassName(final int column) { - return resultColumnType.getFieldList().get(column - 1).getType().getSqlTypeName().getClass().getName(); + RelDataType relDataType = resultColumnType.getFieldList().get(column - 1).getType(); + if (relDataType instanceof JavaType && BigInteger.class.isAssignableFrom(((JavaType) relDataType).getJavaClass())) { + return BigInteger.class.getName(); + } + SqlTypeName originalSqlTypeName = relDataType.getSqlTypeName(); + if (null != columnTypeConverter) { + Class convertedClass = columnTypeConverter.convertColumnValueClass(originalSqlTypeName); + if (null != convertedClass) { + return convertedClass.getName(); + } + } + return getColumnClassNameByType(getColumnType(column)); + } + + private String getColumnClassNameByType(final int columnType) { + switch (columnType) { + case Types.BOOLEAN: + case Types.BIT: + return Boolean.class.getName(); + case Types.TINYINT: + case Types.SMALLINT: + case Types.INTEGER: + return Integer.class.getName(); + case Types.BIGINT: + return Long.class.getName(); + case Types.FLOAT: + case Types.REAL: + return Float.class.getName(); + case Types.DOUBLE: + return Double.class.getName(); + case Types.NUMERIC: + case Types.DECIMAL: + return BigDecimal.class.getName(); + case Types.DATE: + return Date.class.getName(); + case Types.TIME: + return Time.class.getName(); + case Types.TIMESTAMP: + return Timestamp.class.getName(); + case Types.BINARY: + case Types.VARBINARY: + case Types.LONGVARBINARY: + return byte[].class.getName(); + case Types.CHAR: + case Types.VARCHAR: + case Types.LONGVARCHAR: + case Types.NCHAR: + case Types.NVARCHAR: + case Types.LONGNVARCHAR: + return String.class.getName(); + default: + return Object.class.getName(); + } } private Optional findTableName(final int column) { diff --git a/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/converter/DialectSQLFederationColumnTypeConverter.java b/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/converter/DialectSQLFederationColumnTypeConverter.java index 1426639ff2daf..e3398d6ae7de9 100644 --- a/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/converter/DialectSQLFederationColumnTypeConverter.java +++ b/kernel/sql-federation/core/src/main/java/org/apache/shardingsphere/sqlfederation/resultset/converter/DialectSQLFederationColumnTypeConverter.java @@ -42,4 +42,14 @@ public interface DialectSQLFederationColumnTypeConverter extends DatabaseTypedSP * @return converted column type */ int convertColumnType(SqlTypeName sqlTypeName); + + /** + * Convert column value class. + * + * @param sqlTypeName original SQL type name + * @return actual Java class of the converted value, or null if no special conversion + */ + default Class convertColumnValueClass(final SqlTypeName sqlTypeName) { + return null; + } } diff --git a/kernel/sql-federation/core/src/test/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaDataTest.java b/kernel/sql-federation/core/src/test/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaDataTest.java index 0e23a03cf9c53..b95a2ac48b272 100644 --- a/kernel/sql-federation/core/src/test/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaDataTest.java +++ b/kernel/sql-federation/core/src/test/java/org/apache/shardingsphere/sqlfederation/resultset/SQLFederationResultSetMetaDataTest.java @@ -33,21 +33,30 @@ import org.apache.shardingsphere.infra.spi.type.typed.TypedSPILoader; import org.apache.shardingsphere.sqlfederation.resultset.converter.DialectSQLFederationColumnTypeConverter; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import java.math.BigDecimal; import java.math.BigInteger; +import java.sql.Date; import java.sql.ResultSetMetaData; +import java.sql.Time; +import java.sql.Timestamp; import java.sql.Types; import java.util.ArrayList; +import java.util.Arrays; import java.util.Collections; import java.util.HashMap; import java.util.List; import java.util.Map; -import static org.hamcrest.Matchers.is; import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.is; import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.doReturn; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -355,9 +364,67 @@ void assertGetColumnClassName() { JavaType varcharType = mock(JavaType.class); when(varcharType.getJavaClass()).thenReturn(SqlTypeName.class); when(varcharType.getSqlTypeName()).thenReturn(SqlTypeName.VARCHAR); + DialectSQLFederationColumnTypeConverter converter = mock(DialectSQLFederationColumnTypeConverter.class); + when(converter.convertColumnType(SqlTypeName.VARCHAR)).thenReturn(Types.VARCHAR); + SQLFederationResultSetMetaData metaData = new SQLFederationResultSetMetaData( + mock(), Collections.emptyList(), databaseType, + createResultType(new String[]{"foo_col"}, varcharType), + Collections.singletonMap(1, "foo_label"), converter); + assertThat(metaData.getColumnClassName(1), is(String.class.getName())); + } + + @Test + void assertGetColumnClassNameForBigInteger() { + JavaType javaBigIntegerType = mock(JavaType.class); + when(javaBigIntegerType.getJavaClass()).thenReturn(BigInteger.class); + when(javaBigIntegerType.getSqlTypeName()).thenReturn(SqlTypeName.DECIMAL); + RelDataType resultType = createResultType(new String[]{"foo_col"}, javaBigIntegerType); + SQLFederationResultSetMetaData metaData = new SQLFederationResultSetMetaData( + mock(), Collections.emptyList(), databaseType, resultType, Collections.singletonMap(1, "foo_label"), mock()); + assertThat(metaData.getColumnClassName(1), is(BigInteger.class.getName())); + } + + @Test + void assertGetColumnClassNameForConvertedValueClass() { + RelDataType booleanType = mock(RelDataType.class); + when(booleanType.getSqlTypeName()).thenReturn(SqlTypeName.BOOLEAN); + DialectSQLFederationColumnTypeConverter converter = mock(DialectSQLFederationColumnTypeConverter.class); + doReturn(Integer.class).when(converter).convertColumnValueClass(SqlTypeName.BOOLEAN); + RelDataType resultType = createResultType(new String[]{"foo_col"}, booleanType); + SQLFederationResultSetMetaData metaData = new SQLFederationResultSetMetaData( + mock(), Collections.emptyList(), databaseType, resultType, Collections.singletonMap(1, "foo_label"), converter); + assertThat(metaData.getColumnClassName(1), is(Integer.class.getName())); + } + + @ParameterizedTest(name = "{0}") + @MethodSource("columnClassNameSource") + void assertGetColumnClassNameByType(final String name, final SqlTypeName sqlTypeName, final int jdbcType, final String expectedClassName) { + RelDataType relDataType = mock(RelDataType.class); + when(relDataType.getSqlTypeName()).thenReturn(sqlTypeName); + DialectSQLFederationColumnTypeConverter converter = mock(DialectSQLFederationColumnTypeConverter.class); + when(converter.convertColumnType(sqlTypeName)).thenReturn(jdbcType); + RelDataType resultType = createResultType(new String[]{"foo_col"}, relDataType); SQLFederationResultSetMetaData metaData = new SQLFederationResultSetMetaData( - mock(), Collections.emptyList(), databaseType, createResultType(new String[]{"foo_col"}, varcharType), Collections.singletonMap(1, "foo_label"), mock()); - assertThat(metaData.getColumnClassName(1), is(SqlTypeName.VARCHAR.getClass().getName())); + mock(), Collections.emptyList(), databaseType, resultType, Collections.singletonMap(1, "foo_label"), converter); + assertThat(metaData.getColumnClassName(1), is(expectedClassName)); + } + + private static Iterable columnClassNameSource() { + return Arrays.asList( + Arguments.of("tinyint", SqlTypeName.TINYINT, Types.TINYINT, Integer.class.getName()), + Arguments.of("smallint", SqlTypeName.SMALLINT, Types.SMALLINT, Integer.class.getName()), + Arguments.of("integer", SqlTypeName.INTEGER, Types.INTEGER, Integer.class.getName()), + Arguments.of("bigint", SqlTypeName.BIGINT, Types.BIGINT, Long.class.getName()), + Arguments.of("float", SqlTypeName.FLOAT, Types.FLOAT, Float.class.getName()), + Arguments.of("double", SqlTypeName.DOUBLE, Types.DOUBLE, Double.class.getName()), + Arguments.of("numeric", SqlTypeName.DECIMAL, Types.NUMERIC, BigDecimal.class.getName()), + Arguments.of("decimal", SqlTypeName.DECIMAL, Types.DECIMAL, BigDecimal.class.getName()), + Arguments.of("date", SqlTypeName.DATE, Types.DATE, Date.class.getName()), + Arguments.of("time", SqlTypeName.TIME, Types.TIME, Time.class.getName()), + Arguments.of("timestamp", SqlTypeName.TIMESTAMP, Types.TIMESTAMP, Timestamp.class.getName()), + Arguments.of("binary", SqlTypeName.BINARY, Types.BINARY, byte[].class.getName()), + Arguments.of("varchar", SqlTypeName.VARCHAR, Types.VARCHAR, String.class.getName()), + Arguments.of("default", SqlTypeName.ANY, Types.OTHER, Object.class.getName())); } private RelDataType createRowType(final boolean nullable, final int precision, final int scale) { diff --git a/kernel/sql-federation/dialect/mysql/src/main/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverter.java b/kernel/sql-federation/dialect/mysql/src/main/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverter.java index d8d5a38c5aeeb..517cc809de862 100644 --- a/kernel/sql-federation/dialect/mysql/src/main/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverter.java +++ b/kernel/sql-federation/dialect/mysql/src/main/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverter.java @@ -42,6 +42,14 @@ public int convertColumnType(final SqlTypeName sqlTypeName) { return result; } + @Override + public Class convertColumnValueClass(final SqlTypeName sqlTypeName) { + if (SqlTypeName.BOOLEAN == sqlTypeName) { + return Integer.class; + } + return null; + } + @Override public String getDatabaseType() { return "MySQL"; diff --git a/kernel/sql-federation/dialect/mysql/src/test/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverterTest.java b/kernel/sql-federation/dialect/mysql/src/test/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverterTest.java index e63aec9cd9827..6730aa677bbd0 100644 --- a/kernel/sql-federation/dialect/mysql/src/test/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverterTest.java +++ b/kernel/sql-federation/dialect/mysql/src/test/java/org/apache/shardingsphere/sqlfederation/mysql/MySQLSQLFederationColumnTypeConverterTest.java @@ -22,6 +22,7 @@ import org.apache.shardingsphere.database.connector.core.type.DatabaseType; import org.apache.shardingsphere.infra.spi.type.typed.TypedSPILoader; import org.apache.shardingsphere.sqlfederation.resultset.converter.DialectSQLFederationColumnTypeConverter; +import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.Arguments; import org.junit.jupiter.params.provider.MethodSource; @@ -37,6 +38,11 @@ class MySQLSQLFederationColumnTypeConverterTest { private final DialectSQLFederationColumnTypeConverter converter = DatabaseTypedSPILoader.getService(DialectSQLFederationColumnTypeConverter.class, databaseType); + @Test + void assertConvertColumnValueClass() { + assertThat(converter.convertColumnValueClass(SqlTypeName.BOOLEAN), is(Integer.class)); + } + @ParameterizedTest(name = "{0}") @MethodSource("convertValueSource") void assertConvertColumnValue(final String name, final Object input, final Object expected) {