diff --git a/ydb-trino-adapter/README.md b/ydb-trino-adapter/README.md index 38a449d0..5b1a1cbb 100644 --- a/ydb-trino-adapter/README.md +++ b/ydb-trino-adapter/README.md @@ -10,26 +10,19 @@ # Инструкция по сборке -```bash - -mvn -f pom.xml -DskipTests package -mvn -f pom.xml -DskipTests dependency:copy-dependencies -DincludeScope=runtime +Из каталога `ydb-trino-adapter`: -mkdir -p docker/trino/plugin -cp target/ydb-trino-0.1.0.jar docker/trino/plugin -cp target/dependency/*.jar docker/trino/plugin - -cd docker -docker-compose down -docker-compose up -d +```bash +bash start.sh ``` +Скрипт собирает плагин и runtime-зависимости в `examples/trino/plugin`, +затем перезапускает пример через `examples/docker-compose.yml`. + ## Запуск Trino CLI ```bash - - -docker exec -it ydb-trino trino +docker-compose -f examples/docker-compose.yml exec trino trino ``` ## Каталог и схема diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteDivideModulus.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteDivideModulus.java index ee251b62..804d1cf6 100644 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteDivideModulus.java +++ b/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteDivideModulus.java @@ -1,8 +1,10 @@ package tech.ydb.trino; +import com.google.common.collect.ImmutableList; import io.trino.matching.Captures; import io.trino.matching.Pattern; import io.trino.plugin.base.expression.ConnectorExpressionRule; +import io.trino.plugin.jdbc.QueryParameter; import io.trino.plugin.jdbc.expression.ParameterizedExpression; import io.trino.spi.expression.Call; import io.trino.spi.expression.Constant; @@ -49,16 +51,26 @@ public Pattern getPattern() { @Override public Optional rewrite(Call call, Captures captures, RewriteContext context) { + if (!(call.getArguments().get(1) instanceof Constant rightConstant) || + !(rightConstant.getValue() instanceof Number number) || + number.longValue() == 0 || number.longValue() == -1) { + return Optional.empty(); + } + Optional left = context.defaultRewrite(call.getArguments().getFirst()); + if (left.isEmpty()) { + return Optional.empty(); + } + Optional right = context.defaultRewrite(call.getArguments().get(1)); + if (right.isEmpty()) { + return Optional.empty(); + } String operator = call.getFunctionName().equals(DIVIDE_FUNCTION_NAME) ? "/" : "%"; - return RewriteUtils.rewriteBinaryExpression( - call, - context, - () -> call.getArguments().get(1) instanceof Constant rightConstant && - rightConstant.getValue() instanceof Number number && - number.longValue() != 0 && - number.longValue() != -1, - (left, right) -> format("(%s) %s (%s)", left, operator, right) - ); + return Optional.of(new ParameterizedExpression( + format("(%s) %s (%s)", left.get().expression(), operator, right.get().expression()), + ImmutableList.builder() + .addAll(left.get().parameters()) + .addAll(right.get().parameters()) + .build())); } } diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUnaryStringOperations.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUnaryStringOperations.java index a8339420..f9018fbe 100644 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUnaryStringOperations.java +++ b/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUnaryStringOperations.java @@ -14,7 +14,6 @@ import io.trino.spi.expression.FunctionName; import io.trino.spi.type.VarcharType; -import java.util.Objects; import java.util.Optional; import static io.trino.matching.Capture.newCapture; @@ -42,21 +41,9 @@ public Pattern getPattern() { @Override public Optional rewrite(ConnectorTableHandle handle, ConnectorExpression projectionExpression, Captures captures, RewriteContext context) { - JdbcTypeHandle varcharTypeHandle = YdbTypeUtils.toTypeHandle(VarcharType.VARCHAR).orElse(null); - if (Objects.isNull(varcharTypeHandle)) { - return Optional.empty(); - } - + JdbcTypeHandle varcharTypeHandle = YdbTypeUtils.toTypeHandle(VarcharType.VARCHAR).orElseThrow(); Call call = (Call) projectionExpression; - - ConnectorExpression valueExpr; - if (call.getArguments().size() == 1) { - valueExpr = call.getArguments().getFirst(); - } else { - return Optional.empty(); - } - - Optional rewrittenValue = context.rewriteExpression(valueExpr); + Optional rewrittenValue = context.rewriteExpression(captures.get(VALUE)); if (rewrittenValue.isEmpty()) { return Optional.empty(); } diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUtils.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUtils.java deleted file mode 100644 index 9b68a86c..00000000 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUtils.java +++ /dev/null @@ -1,40 +0,0 @@ -package tech.ydb.trino; - -import com.google.common.collect.ImmutableList; -import io.trino.plugin.base.expression.ConnectorExpressionRule.RewriteContext; -import io.trino.plugin.jdbc.QueryParameter; -import io.trino.plugin.jdbc.expression.ParameterizedExpression; -import io.trino.spi.expression.Call; -import io.trino.spi.expression.ConnectorExpression; -import org.jspecify.annotations.NonNull; - -import java.util.ArrayList; -import java.util.List; -import java.util.Optional; -import java.util.function.BiFunction; -import java.util.function.Supplier; - -public class RewriteUtils { - public static Optional rewriteBinaryExpression( - Call call, - RewriteContext context, - Supplier condition, - // leftSql, rightSql -> resultSql - BiFunction queryCombiner - ) { - if (!condition.get()) { - return Optional.empty(); - } - List sqls = new ArrayList<>(); - ImmutableList.Builder<@NonNull QueryParameter> parameters = ImmutableList.builder(); - for (ConnectorExpression connectorExpression : call.getArguments()) { - Optional expression = context.defaultRewrite(connectorExpression); - if (expression.isEmpty()) { - return Optional.empty(); - } - parameters.addAll(expression.get().parameters()); - sqls.add(expression.get().expression()); - } - return Optional.of(new ParameterizedExpression(queryCombiner.apply(sqls.get(0), sqls.get(1)), parameters.build())); - } -} diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClient.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClient.java index bcb56bec..0d653f3f 100644 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClient.java +++ b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClient.java @@ -30,7 +30,6 @@ import io.trino.plugin.jdbc.QueryBuilder; import io.trino.plugin.jdbc.RemoteTableName; import io.trino.plugin.jdbc.WriteMapping; -import io.trino.plugin.jdbc.aggregation.ImplementAvgDecimal; import io.trino.plugin.jdbc.aggregation.ImplementAvgFloatingPoint; import io.trino.plugin.jdbc.aggregation.ImplementCount; import io.trino.plugin.jdbc.aggregation.ImplementCountAll; @@ -96,17 +95,14 @@ import static io.trino.plugin.jdbc.StandardColumnMappings.doubleWriteFunction; import static io.trino.plugin.jdbc.StandardColumnMappings.integerColumnMapping; import static io.trino.plugin.jdbc.StandardColumnMappings.integerWriteFunction; -import static io.trino.plugin.jdbc.StandardColumnMappings.longDecimalWriteFunction; import static io.trino.plugin.jdbc.StandardColumnMappings.realColumnMapping; import static io.trino.plugin.jdbc.StandardColumnMappings.realWriteFunction; -import static io.trino.plugin.jdbc.StandardColumnMappings.shortDecimalWriteFunction; import static io.trino.plugin.jdbc.StandardColumnMappings.smallintColumnMapping; import static io.trino.plugin.jdbc.StandardColumnMappings.smallintWriteFunction; import static io.trino.plugin.jdbc.StandardColumnMappings.tinyintWriteFunction; import static io.trino.plugin.jdbc.StandardColumnMappings.tinyintColumnMapping; import static io.trino.plugin.jdbc.StandardColumnMappings.varbinaryColumnMapping; import static io.trino.plugin.jdbc.StandardColumnMappings.varbinaryWriteFunction; -import static io.trino.plugin.jdbc.StandardColumnMappings.varcharColumnMapping; import static io.trino.plugin.jdbc.StandardColumnMappings.varcharReadFunction; import static io.trino.plugin.jdbc.StandardColumnMappings.varcharWriteFunction; import static io.trino.plugin.jdbc.TypeHandlingJdbcSessionProperties.getUnsupportedTypeHandling; @@ -127,7 +123,6 @@ import static io.trino.spi.type.TinyintType.TINYINT; import static io.trino.spi.type.VarbinaryType.VARBINARY; import static io.trino.spi.type.VarcharType.createUnboundedVarcharType; -import static io.trino.spi.type.VarcharType.createVarcharType; import static java.lang.Math.max; import static java.lang.String.format; import static java.util.stream.Collectors.joining; @@ -174,8 +169,6 @@ public YdbClient( .add(new RewriteIn()) .add(new RewriteDivideModulus()) .add(new RewriteNullIf()) - .withTypeClass("integer_type", ImmutableSet.of("tinyint", "smallint", "integer", "bigint")) - .withTypeClass("numeric_type", ImmutableSet.of("tinyint", "smallint", "integer", "bigint", "decimal", "real", "double")) .withTypeClass("comparable_type", ImmutableSet.of( "tinyint", "smallint", "integer", "bigint", "decimal", "real", "double", "varchar", "char", "date", "timestamp")) .map("$equal(left, right)").to("left = right") @@ -206,7 +199,6 @@ public YdbClient( .add(new ImplementCountDistinct(bigintTypeHandle, true)) .add(new ImplementSum(YdbTypeUtils::toTypeHandle)) .add(new ImplementAvgFloatingPoint()) - .add(new ImplementAvgDecimal()) .build()); } @@ -426,64 +418,30 @@ public Optional toColumnMapping( return getUnsupportedTypeHandling(session) == CONVERT_TO_VARCHAR ? mapToUnboundedVarchar(typeHandle) : Optional.empty(); } - Optional columnMapping = switch (typeHandle.jdbcType()) { - case Types.BIT, Types.BOOLEAN -> Optional.of(booleanColumnMapping()); - case Types.TINYINT, Types.SMALLINT -> Optional.of(smallintColumnMapping()); - case Types.INTEGER -> Optional.of(integerColumnMapping()); - case Types.BIGINT -> Optional.of(bigintColumnMapping()); - case Types.REAL -> Optional.of(realColumnMapping()); - case Types.FLOAT -> Optional.of(jdbcTypeName.equals("float") ? realColumnMapping() : doubleColumnMapping()); - case Types.DOUBLE -> Optional.of(doubleColumnMapping()); - case Types.DECIMAL -> { - String typeName = typeHandle.jdbcTypeName().orElse("Decimal"); - int precision = typeHandle.columnSize().orElse(YDB_DEFAULT_DECIMAL_PRECISION); - int scale = typeHandle.decimalDigits().orElse(YDB_DEFAULT_DECIMAL_SCALE); - int start = typeName.indexOf('('); - int end = typeName.indexOf(')'); - if (start >= 0 && end > start) { - String[] parts = typeName.substring(start + 1, end).split(","); - if (parts.length == 2) { - Integer typeNamePrecision = Ints.tryParse(parts[0].trim()); - Integer typeNameScale = Ints.tryParse(parts[1].trim()); - if (typeNamePrecision != null && typeNameScale != null) { - precision = typeNamePrecision; - scale = typeNameScale; - } - } + String typeName = typeHandle.jdbcTypeName().orElse("Decimal"); + int precision = typeHandle.columnSize().orElse(YDB_DEFAULT_DECIMAL_PRECISION); + int scale = typeHandle.decimalDigits().orElse(YDB_DEFAULT_DECIMAL_SCALE); + int start = typeName.indexOf('('); + int end = typeName.indexOf(')'); + if (start >= 0 && end > start) { + String[] parts = typeName.substring(start + 1, end).split(","); + if (parts.length == 2) { + Integer typeNamePrecision = Ints.tryParse(parts[0].trim()); + Integer typeNameScale = Ints.tryParse(parts[1].trim()); + if (typeNamePrecision != null && typeNameScale != null) { + precision = typeNamePrecision; + scale = typeNameScale; } - - DecimalType decimalType = createDecimalType(precision, max(scale, 0)); - ColumnMapping decimalMapping = decimalColumnMapping(decimalType); - yield Optional.of(ColumnMapping.mapping( - decimalType, - decimalMapping.getReadFunction(), - decimalMapping.getWriteFunction(), - DISABLE_PUSHDOWN)); - } - case Types.CHAR, Types.NCHAR -> { - String typeName = typeHandle.jdbcTypeName().orElseThrow(); - int length = typeName.toLowerCase().startsWith("char(") - ? Integer.parseInt(typeName.substring(5, typeName.length() - 1)) - : typeHandle.columnSize().orElse(VarcharType.MAX_LENGTH); - yield Optional.of(varcharColumnMapping(length)); - } - case Types.VARCHAR, Types.LONGVARCHAR, Types.NVARCHAR -> { - String typeName = typeHandle.jdbcTypeName().orElseThrow(); - int length = typeName.toLowerCase().startsWith("varchar(") - ? Integer.parseInt(typeName.substring(8, typeName.length() - 1)) - : typeHandle.columnSize().orElse(VarcharType.MAX_LENGTH); - yield Optional.of(varcharColumnMapping(length)); } - case Types.DATE -> Optional.of(dateColumnMapping(typeHandle)); - case Types.TIMESTAMP -> Optional.of(timestampColumnMapping(typeHandle)); - default -> Optional.empty(); - }; - - if (columnMapping.isPresent()) { - return columnMapping; } - return mapToUnboundedVarchar(typeHandle); + DecimalType decimalType = createDecimalType(precision, max(scale, 0)); + ColumnMapping decimalMapping = decimalColumnMapping(decimalType); + return Optional.of(ColumnMapping.mapping( + decimalType, + decimalMapping.getReadFunction(), + decimalMapping.getWriteFunction(), + DISABLE_PUSHDOWN)); } private static ColumnMapping unboundedVarcharColumnMapping() { @@ -495,17 +453,6 @@ private static ColumnMapping unboundedVarcharColumnMapping() { FULL_PUSHDOWN); } - private static ColumnMapping varcharColumnMapping(int varcharLength) { - VarcharType varcharType = varcharLength <= VarcharType.MAX_LENGTH - ? createVarcharType(varcharLength) - : createUnboundedVarcharType(); - return ColumnMapping.sliceMapping( - varcharType, - varcharReadFunction(varcharType), - varcharWriteFunction(), - FULL_PUSHDOWN); - } - @Override public WriteMapping toWriteMapping(ConnectorSession session, Type type) { if (type == BOOLEAN) { @@ -849,7 +796,8 @@ public OptionalLong delete(ConnectorSession session, JdbcTableHandle handle) { handle.getRequiredNamedRelation(), handle.getConstraint(), getAdditionalPredicate(handle.getConstraintExpressions(), Optional.empty())); - return OptionalLong.of(executeReturningDml(session, connection, handle, preparedQuery)); + List primaryKeys = getPrimaryKeys(session, handle.getRequiredNamedRelation().getRemoteTableName()); + return OptionalLong.of(executeReturningDml(session, connection, primaryKeys, preparedQuery)); } catch (SQLException e) { throw new TrinoException(JDBC_ERROR, e); @@ -858,8 +806,9 @@ public OptionalLong delete(ConnectorSession session, JdbcTableHandle handle) { @Override public OptionalLong update(ConnectorSession session, JdbcTableHandle handle) { + List primaryKeys = getPrimaryKeys(session, handle.getRequiredNamedRelation().getRemoteTableName()); verifyNoPrimaryKeyUpdate( - getPrimaryKeys(session, handle.getRequiredNamedRelation().getRemoteTableName()), + primaryKeys, handle.getUpdateAssignments().stream().map(assignment -> (ColumnHandle) assignment.column()).toList()); try (Connection connection = connectionFactory.openConnection(session)) { PreparedQuery preparedQuery = queryBuilder.prepareUpdateQuery( @@ -870,7 +819,7 @@ public OptionalLong update(ConnectorSession session, JdbcTableHandle handle) { handle.getConstraint(), getAdditionalPredicate(handle.getConstraintExpressions(), Optional.empty()), handle.getUpdateAssignments()); - return OptionalLong.of(executeReturningDml(session, connection, handle, preparedQuery)); + return OptionalLong.of(executeReturningDml(session, connection, primaryKeys, preparedQuery)); } catch (SQLException e) { throw new TrinoException(JDBC_ERROR, e); @@ -887,11 +836,8 @@ private static void verifyNoPrimaryKeyUpdate(List primaryKeys, private long executeReturningDml( ConnectorSession session, Connection connection, - JdbcTableHandle handle, + List primaryKeys, PreparedQuery preparedQuery) throws SQLException { - List primaryKeys = getPrimaryKeys( - session, - handle.getRequiredNamedRelation().getRemoteTableName()); if (primaryKeys.isEmpty()) { throw new TrinoException(NOT_SUPPORTED, "YDB DML requires a table primary key"); } diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadata.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadata.java deleted file mode 100644 index 4389043a..00000000 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadata.java +++ /dev/null @@ -1,18 +0,0 @@ -package tech.ydb.trino; - -import io.trino.plugin.jdbc.DefaultJdbcMetadata; -import io.trino.plugin.jdbc.JdbcClient; -import io.trino.plugin.jdbc.JdbcQueryEventListener; -import io.trino.plugin.jdbc.TimestampTimeZoneDomain; - -import java.util.Set; - -public class YdbMetadata extends DefaultJdbcMetadata { - public YdbMetadata( - JdbcClient jdbcClient, - TimestampTimeZoneDomain timestampTimeZoneDomain, - Set jdbcQueryEventListeners - ) { - super(jdbcClient, timestampTimeZoneDomain, false, jdbcQueryEventListeners); - } -} diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadataFactory.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadataFactory.java index faa9960d..627fbaa8 100644 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadataFactory.java +++ b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadataFactory.java @@ -2,6 +2,7 @@ import com.google.inject.Inject; import io.trino.plugin.base.cache.identity.IdentityCacheMapping; +import io.trino.plugin.jdbc.DefaultJdbcMetadata; import io.trino.plugin.jdbc.DefaultJdbcMetadataFactory; import io.trino.plugin.jdbc.JdbcClient; import io.trino.plugin.jdbc.JdbcMetadata; @@ -28,6 +29,6 @@ public YdbMetadataFactory( @Override protected JdbcMetadata create(JdbcClient transactionCachingJdbcClient) { - return new YdbMetadata(transactionCachingJdbcClient, timestampTimeZoneDomain, jdbcQueryEventListeners); + return new DefaultJdbcMetadata(transactionCachingJdbcClient, timestampTimeZoneDomain, false, jdbcQueryEventListeners); } } diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbColumnMappings.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbColumnMappings.java index 33c8202e..15236a25 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbColumnMappings.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbColumnMappings.java @@ -99,9 +99,13 @@ public void testUnsignedWriterBounds() throws Exception { public void testDecimalWriterPreservesPrecisionAndScale() throws Exception { List> calls = new ArrayList<>(); PreparedStatement statement = statement(calls); - LongWriteFunction shortWriter = (LongWriteFunction) YdbColumnMappings.decimalColumnMapping(createDecimalType(10, 3)).getWriteFunction(); + JdbcTypeHandle shortType = new JdbcTypeHandle(Types.DECIMAL, Optional.of("Decimal(10,3)"), + Optional.of(22), Optional.of(9), Optional.empty(), Optional.empty()); + LongWriteFunction shortWriter = (LongWriteFunction) client.toColumnMapping(SESSION, null, shortType).orElseThrow().getWriteFunction(); shortWriter.set(statement, 1, 12345); - ObjectWriteFunction longWriter = (ObjectWriteFunction) YdbColumnMappings.decimalColumnMapping(createDecimalType(35, 10)).getWriteFunction(); + JdbcTypeHandle longType = new JdbcTypeHandle(Types.DECIMAL, Optional.of("Decimal"), + Optional.of(35), Optional.of(10), Optional.empty(), Optional.empty()); + ObjectWriteFunction longWriter = (ObjectWriteFunction) client.toColumnMapping(SESSION, null, longType).orElseThrow().getWriteFunction(); longWriter.set(statement, 2, Int128.valueOf(new BigInteger("123456789012345678901234567890"))); longWriter.setNull(statement, 3); assertThat(calls).containsExactly( diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbConnectorTest.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbConnectorTest.java index 7dd4dcc7..e138a3db 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbConnectorTest.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbConnectorTest.java @@ -46,6 +46,24 @@ public void testDefaultSchema() { assertQueryFails("SELECT * FROM local.\"%\".orders", ".*Schema '%' does not exist"); } + @Test + public void testPrimaryKeyNamedColumn() { + try (TestTable table = newTrinoTable("primary_key_named_column_", "(\"primary_key\" bigint)")) { + assertUpdate("INSERT INTO " + table.getName() + " VALUES 7", 1); + assertQuery("SELECT \"primary_key\" FROM " + table.getName(), "VALUES 7"); + // Quoting does not bypass YDB's column naming restrictions: + // https://ydb.tech/docs/en/concepts/datamodel/table#column-naming-rules + String invalidTable = table.getName() + "_invalid"; + try { + assertQueryFails("CREATE TABLE " + invalidTable + " (\"primary key\" bigint)", + "(?s).*Invalid name for user column 'primary key'.*"); + } + finally { + assertUpdate("DROP TABLE IF EXISTS " + invalidTable); + } + } + } + @Test public void testJoinPushdown() { Session session = Session.builder(getSession()) diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbExpressionRewrites.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbExpressionRewrites.java index cd49145e..8314404d 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbExpressionRewrites.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbExpressionRewrites.java @@ -45,36 +45,51 @@ public class TestYdbExpressionRewrites { @Test public void testNullIfParameterOrder() { Call expression = new Call(BIGINT, NULLIF_FUNCTION_NAME, List.of( - new Constant(11L, BIGINT), new Variable("v", BIGINT))); - JdbcColumnHandle column = new JdbcColumnHandle("value", YdbTypeUtils.toTypeHandle(BIGINT).orElseThrow(), BIGINT); - var result = client.convertPredicate(SESSION, expression, Map.of("v", column)).orElseThrow(); + new Constant(11L, BIGINT), new Constant(22L, BIGINT))); + var result = client.convertPredicate(SESSION, expression, Map.of()).orElseThrow(); assertThat(result.parameters()).extracting(parameter -> parameter.getValue().orElseThrow()) - .containsExactly(11L, 11L); + .containsExactly(11L, 22L, 11L); } @Test public void testStringPositionParameterOrder() { Call expression = new Call(BIGINT, new FunctionName("strpos"), List.of( - new Constant(utf8Slice("hello"), VARCHAR), new Variable("v", VARCHAR))); - JdbcColumnHandle column = new JdbcColumnHandle("value", YdbTypeUtils.toTypeHandle(VARCHAR).orElseThrow(), VARCHAR); + new Constant(utf8Slice("hello"), VARCHAR), new Constant(utf8Slice("ll"), VARCHAR))); JdbcTableHandle table = new JdbcTableHandle( new SchemaTableName("default", "test"), new RemoteTableName(Optional.empty(), Optional.empty(), "test"), Optional.empty()); - var result = client.convertProjection(SESSION, table, expression, Map.of("v", column)).orElseThrow(); + var result = client.convertProjection(SESSION, table, expression, Map.of()).orElseThrow(); assertThat(result.getParameters()).extracting(parameter -> parameter.getValue().orElseThrow()) - .containsExactly(utf8Slice("hello"), utf8Slice("hello")); + .containsExactly(utf8Slice("hello"), utf8Slice("ll"), utf8Slice("hello"), utf8Slice("ll")); } @Test - public void testUnicodeTrimIsNotPushedDown() { + public void testUnicodeCaseRewritesAndTrimFallback() { JdbcColumnHandle column = new JdbcColumnHandle("value", YdbTypeUtils.toTypeHandle(VARCHAR).orElseThrow(), VARCHAR); JdbcTableHandle table = new JdbcTableHandle(new SchemaTableName("default", "test"), new RemoteTableName(Optional.empty(), Optional.empty(), "test"), Optional.empty()); + for (String function : List.of("upper", "lower")) { + assertThat(client.convertProjection(SESSION, table, + new Call(VARCHAR, new FunctionName(function), List.of(new Variable("v", VARCHAR))), Map.of("v", column))).isPresent(); + } assertThat(client.convertProjection(SESSION, table, new Call(VARCHAR, new FunctionName("trim"), List.of(new Variable("v", VARCHAR))), Map.of("v", column))).isEmpty(); } + @Test + public void testIntegralDivisionAndModulusPushdown() { + for (var entry : Map.of(DIVIDE_FUNCTION_NAME, "(?) / (?)", MODULO_FUNCTION_NAME, "(?) % (?)").entrySet()) { + var result = client.convertPredicate(SESSION, new Call(BIGINT, entry.getKey(), List.of( + new Constant(11L, BIGINT), new Constant(2L, BIGINT))), Map.of()).orElseThrow(); + assertThat(result.expression()).isEqualTo(entry.getValue()); + assertThat(result.parameters()).extracting(parameter -> parameter.getValue().orElseThrow()) + .containsExactly(11L, 2L); + assertThat(client.convertPredicate(SESSION, new Call(BIGINT, entry.getKey(), List.of( + new Constant(11L, BIGINT), new Constant(0L, BIGINT))), Map.of())).isEmpty(); + } + } + @Test public void testSignedMinimumModuloFallsBack() { assertThat(BigintOperators.modulo(Long.MIN_VALUE, -1)).isZero(); @@ -84,9 +99,10 @@ public void testSignedMinimumModuloFallsBack() { } @Test - public void testTimestamp64ExpressionBounds() { + public void testTimestampLiteralPredicatesAreNotRewritten() { + // Literal rewriting is unsupported even in range; domain and writer bounds have separate tests. JdbcColumnHandle column = new JdbcColumnHandle("value", YdbTypeUtils.toTypeHandle(TIMESTAMP_MICROS).orElseThrow(), TIMESTAMP_MICROS); - for (long bound : List.of(-4611669897600000001L, 4611669811200000000L)) { + for (long bound : List.of(-4611669897600000001L, 0L, 4611669811200000000L)) { assertThat(client.convertPredicate(SESSION, new Call(BOOLEAN, LESS_THAN_OPERATOR_FUNCTION_NAME, List.of( new Variable("v", TIMESTAMP_MICROS), new Constant(bound, TIMESTAMP_MICROS))), Map.of("v", column))).isEmpty(); } diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbPlugin.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbPlugin.java index 8af74c06..e803d760 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbPlugin.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbPlugin.java @@ -2,7 +2,8 @@ import io.trino.spi.connector.Connector; import io.trino.testing.TestingConnectorContext; -import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; import java.util.Map; @@ -11,9 +12,11 @@ import static org.assertj.core.api.Assertions.assertThat; public class TestYdbPlugin { - @Test - public void testProductionConnectorBootstrap() { - Connector connector = getOnlyElement(new YdbPlugin().getConnectorFactories()) + @ParameterizedTest + @ValueSource(booleans = {false, true}) + public void testConnectorBootstrap(boolean testingClient) { + YdbPlugin plugin = new YdbPlugin(testingClient ? new TestingYdbJdbcModule() : new YdbClientModule()); + Connector connector = getOnlyElement(plugin.getConnectorFactories()) .create("ydb_bootstrap", Map.of("connection-url", "jdbc:ydb:grpc://127.0.0.1:2136/local"), new TestingConnectorContext()); try { diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbWriteMetadata.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbWriteMetadata.java index 0332978e..3674419e 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbWriteMetadata.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestYdbWriteMetadata.java @@ -2,25 +2,47 @@ import io.trino.plugin.base.mapping.DefaultIdentifierMapping; import io.trino.plugin.jdbc.BaseJdbcConfig; +import io.trino.plugin.jdbc.RemoteTableName; import io.trino.plugin.jdbc.logging.RemoteQueryModifier; +import io.trino.spi.connector.ColumnMetadata; +import io.trino.spi.connector.ConnectorTableMetadata; +import io.trino.spi.connector.SchemaTableName; import org.junit.jupiter.api.Test; import java.lang.reflect.Proxy; import java.sql.Connection; import java.sql.DatabaseMetaData; import java.sql.ResultSet; +import java.sql.SQLException; import java.sql.Statement; import java.sql.Types; import java.util.ArrayList; import java.util.List; import java.util.Map; +import java.util.Optional; import java.util.concurrent.atomic.AtomicBoolean; import java.util.function.Supplier; +import static io.trino.spi.type.BigintType.BIGINT; import static io.trino.testing.TestingConnectorSession.SESSION; import static org.assertj.core.api.Assertions.assertThat; public class TestYdbWriteMetadata { + @Test + public void testHiddenKeyWithPrimaryKeyNamedColumn() { + TestingYdbJdbcClient client = new TestingYdbJdbcClient(new BaseJdbcConfig(), + _ -> { + throw new SQLException("This test must not open a connection"); + }, + new YdbQueryBuilder(RemoteQueryModifier.NONE), new DefaultIdentifierMapping(), RemoteQueryModifier.NONE); + ConnectorTableMetadata metadata = new ConnectorTableMetadata( + new SchemaTableName("default", "fixture"), List.of(new ColumnMetadata("primary key", BIGINT))); + assertThat(client.createTableSqls( + new RemoteTableName(Optional.empty(), Optional.empty(), "fixture"), + List.of("`primary key` Int64"), metadata)) + .containsExactly("CREATE TABLE `fixture` (`primary key` Int64, `_ydb_trino_test_pk` Serial, PRIMARY KEY (`_ydb_trino_test_pk`))"); + } + @Test public void testStagingUsesDriverTableIdentity() throws Exception { Map row = Map.of( diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcClient.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcClient.java index 82e4f1ff..e30493ae 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcClient.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcClient.java @@ -66,14 +66,11 @@ protected List createTableSqls(RemoteTableName remoteTableName, List col.startsWith(quoted(YDB_HIDDEN_PK_COLUMN)) || col.startsWith(YDB_HIDDEN_PK_COLUMN)); String sql; - if (hasPrimaryKey) { - sql = String.format("CREATE TABLE %s (%s)", tableName, columnsDeclaration); - } else if (hasHiddenPkColumn) { + if (hasHiddenPkColumn) { sql = String.format("CREATE TABLE %s (%s, PRIMARY KEY (%s))", tableName, columnsDeclaration, quoted(YDB_HIDDEN_PK_COLUMN)); } else { String hiddenPkColumn = quoted(YDB_HIDDEN_PK_COLUMN) + " Serial"; diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcModule.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcModule.java index 60f0c3d8..d84c2bac 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcModule.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/TestingYdbJdbcModule.java @@ -1,64 +1,33 @@ package tech.ydb.trino; -import com.google.inject.Binder; -import com.google.inject.Module; +import com.google.inject.AbstractModule; import com.google.inject.Provides; -import com.google.inject.Scopes; import com.google.inject.Singleton; +import com.google.inject.util.Modules; import io.trino.plugin.base.mapping.IdentifierMapping; import io.trino.plugin.jdbc.BaseJdbcConfig; import io.trino.plugin.jdbc.ConnectionFactory; import io.trino.plugin.jdbc.ForBaseJdbc; import io.trino.plugin.jdbc.JdbcClient; -import io.trino.plugin.jdbc.JdbcMetadataFactory; import io.trino.plugin.jdbc.QueryBuilder; -import io.trino.plugin.jdbc.credential.CredentialProvider; import io.trino.plugin.jdbc.logging.RemoteQueryModifier; -import io.trino.spi.connector.ConnectorPageSinkProvider; -import static com.google.inject.multibindings.OptionalBinder.newOptionalBinder; -import static io.trino.plugin.jdbc.JdbcModule.bindTablePropertiesProvider; - -public class TestingYdbJdbcModule implements Module { +public class TestingYdbJdbcModule extends AbstractModule { @Override - public void configure(Binder binder) { - newOptionalBinder(binder, QueryBuilder.class) - .setBinding() - .to(YdbQueryBuilder.class) - .in(Scopes.SINGLETON); - - newOptionalBinder(binder, JdbcMetadataFactory.class) - .setBinding() - .to(YdbMetadataFactory.class) - .in(Scopes.SINGLETON); - - bindTablePropertiesProvider(binder, YdbTableProperties.class); - newOptionalBinder(binder, ConnectorPageSinkProvider.class) - .setBinding() - .to(YdbPageSinkProvider.class) - .in(Scopes.SINGLETON); - binder.bind(YdbConnector.class).in(Scopes.SINGLETON); - } - - @Provides - @Singleton - @ForBaseJdbc - public JdbcClient provideJdbcClient( - BaseJdbcConfig config, - ConnectionFactory connectionFactory, - QueryBuilder queryBuilder, - IdentifierMapping identifierMapping, - RemoteQueryModifier remoteQueryModifier) { - return new TestingYdbJdbcClient(config, connectionFactory, queryBuilder, identifierMapping, remoteQueryModifier); - } - - @Provides - @Singleton - @ForBaseJdbc - public static ConnectionFactory createConnectionFactory( - BaseJdbcConfig config, - CredentialProvider credentialProvider) { - return YdbClientModule.createConnectionFactory(config, credentialProvider); + protected void configure() { + install(Modules.override(new YdbClientModule()).with(new AbstractModule() { + @Provides + @Singleton + @ForBaseJdbc + public JdbcClient provideJdbcClient( + BaseJdbcConfig config, + ConnectionFactory connectionFactory, + QueryBuilder queryBuilder, + IdentifierMapping identifierMapping, + RemoteQueryModifier remoteQueryModifier) { + return new TestingYdbJdbcClient(config, connectionFactory, queryBuilder, identifierMapping, remoteQueryModifier); + } + })); } } diff --git a/ydb-trino-adapter/src/test/java/tech/ydb/trino/YdbQueryRunner.java b/ydb-trino-adapter/src/test/java/tech/ydb/trino/YdbQueryRunner.java index c274004a..5b0e14b5 100644 --- a/ydb-trino-adapter/src/test/java/tech/ydb/trino/YdbQueryRunner.java +++ b/ydb-trino-adapter/src/test/java/tech/ydb/trino/YdbQueryRunner.java @@ -3,13 +3,8 @@ import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; import com.google.inject.Module; -import io.trino.Session; import io.trino.testing.DistributedQueryRunner; -import io.trino.testing.MaterializedResult; -import io.trino.testing.QueryRunner; -import io.trino.tpch.TpchColumn; import io.trino.tpch.TpchTable; -import org.intellij.lang.annotations.Language; import tech.ydb.test.junit5.YdbHelperExtension; import java.util.HashMap; @@ -28,7 +23,7 @@ private YdbQueryRunner() {} public static Builder builder(YdbHelperExtension ydb) { String jdbcUrl = buildJdbcUrl(ydb); return new Builder() - // Avoid temporary-table CTAS during INSERT; YDB does not support CREATE TABLE AS SELECT. + // Transactional INSERT staging has dedicated coverage in TestYdbCreateTable. .addConnectorProperty("insert.non-transactional-insert.enabled", "true") .addConnectorProperty("merge.non-transactional-merge.enabled", "true") .addConnectorProperty("connection-url", jdbcUrl); diff --git a/ydb-trino-adapter/start.sh b/ydb-trino-adapter/start.sh index aa02dc75..e15747bc 100755 --- a/ydb-trino-adapter/start.sh +++ b/ydb-trino-adapter/start.sh @@ -1,10 +1,15 @@ +#!/usr/bin/env bash +set -euo pipefail + +cd -- "$(dirname -- "${BASH_SOURCE[0]}")" + mvn -f pom.xml -DskipTests package mvn -f pom.xml -DskipTests dependency:copy-dependencies -DincludeScope=runtime -mkdir -p docker/trino/plugin -cp target/ydb-trino-0.1.0.jar docker/trino/plugin -cp target/dependency/*.jar docker/trino/plugin +mkdir -p examples/trino/plugin +cp target/ydb-trino-0.1.0.jar examples/trino/plugin +cp target/dependency/*.jar examples/trino/plugin -cd docker +cd examples docker-compose down -docker-compose up -d \ No newline at end of file +docker-compose up -d