diff --git a/ydb-trino-adapter/README.md b/ydb-trino-adapter/README.md index e4aad3c5..c3a5ff36 100644 --- a/ydb-trino-adapter/README.md +++ b/ydb-trino-adapter/README.md @@ -64,6 +64,14 @@ AS SELECT tenant, event_id, payload FROM events; integration-проверка CTAS использует `insert.non-transactional-insert.enabled=true`; transactional staging этой проверкой не подтверждается. +## JOIN pushdown + +По умолчанию JOIN выполняет Trino. Для пробного pushdown задайте +`join_pushdown_enabled=true` в сессии каталога. Адаптер передаёт YDB +`INNER`, `LEFT`, `RIGHT` и `FULL JOIN` только по равенству исходных столбцов +`Int64` (Trino `bigint`), включая составной ключ. Остальные условия JOIN, +вычисляемые ключи и другие типы остаются в Trino. См. [правила YQL JOIN](https://ydb.tech/docs/ru/yql/reference/syntax/select/join). + ## Текст и байты YDB `Text` отображается в Trino как `varchar`, а `Bytes` — как `varbinary` без 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 87d7f8ae..29407615 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 @@ -18,6 +18,7 @@ import io.trino.plugin.jdbc.ConnectionFactory; import io.trino.plugin.jdbc.JdbcColumnHandle; import io.trino.plugin.jdbc.JdbcExpression; +import io.trino.plugin.jdbc.JdbcJoinCondition; import io.trino.plugin.jdbc.JdbcMergeTableHandle; import io.trino.plugin.jdbc.JdbcOutputTableHandle; import io.trino.plugin.jdbc.JdbcSortItem; @@ -44,6 +45,8 @@ import io.trino.spi.connector.ColumnMetadata; import io.trino.spi.connector.ConnectorSession; import io.trino.spi.connector.ConnectorTableMetadata; +import io.trino.spi.connector.JoinStatistics; +import io.trino.spi.connector.JoinType; import io.trino.spi.connector.RetryMode; import io.trino.spi.connector.SchemaTableName; import io.trino.spi.connector.SortOrder; @@ -102,6 +105,7 @@ import static io.trino.plugin.jdbc.StandardColumnMappings.varcharWriteFunction; import static io.trino.spi.StandardErrorCode.INVALID_TABLE_PROPERTY; import static io.trino.spi.StandardErrorCode.NOT_SUPPORTED; +import static io.trino.spi.connector.JoinCondition.Operator.EQUAL; import static io.trino.spi.type.BigintType.BIGINT; import static io.trino.spi.type.BooleanType.BOOLEAN; import static io.trino.spi.type.DateType.DATE; @@ -199,6 +203,42 @@ public YdbClient( .build()); } + @Override + protected boolean isSupportedJoinCondition(ConnectorSession session, JdbcJoinCondition condition) { + return condition.getOperator() == EQUAL + && hasSupportedValueMapping(condition.getLeftColumn()) + && hasSupportedValueMapping(condition.getRightColumn()); + } + + private boolean hasSupportedValueMapping(JdbcColumnHandle column) { + JdbcTypeHandle type = column.getJdbcTypeHandle(); + if (getForcedMappingToVarchar(type).isPresent()) { + return false; + } + String name = type.jdbcTypeName().orElse("").toLowerCase(Locale.ROOT); + return switch (name) { + case "bool", "int8", "int16", "int32", "int64", "uint8", "uint16", "uint32", "uint64", + "float", "double", "utf8", "text", "date", "date32", "datetime", "datetime64", + "timestamp", "timestamp64", "decimal" -> true; + case "string", "bytes" -> column.getColumnType().equals(VARBINARY); + default -> name.startsWith("decimal("); + }; + } + + @Override + public Optional implementJoin( + ConnectorSession session, + JoinType joinType, + PreparedQuery leftSource, + Map leftProjections, + PreparedQuery rightSource, + Map rightProjections, + List joinConditions, + JoinStatistics statistics) { + // The expression API loses which source owns each YQL ON operand. + return Optional.empty(); + } + @Override public Optional implementAggregation( ConnectorSession session, diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClientModule.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClientModule.java index 46d8f723..bb6e18fe 100644 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClientModule.java +++ b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClientModule.java @@ -24,6 +24,11 @@ public class YdbClientModule implements Module { @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) diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbPlugin.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbPlugin.java index 5ef33be0..b44488e9 100644 --- a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbPlugin.java +++ b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbPlugin.java @@ -2,11 +2,13 @@ import com.google.common.collect.ImmutableList; import com.google.inject.Module; +import io.trino.plugin.jdbc.JdbcMetadataConfig; import io.trino.plugin.jdbc.credential.CredentialProviderModule; import io.trino.spi.Plugin; import io.trino.spi.connector.ConnectorFactory; import static io.airlift.configuration.ConfigurationAwareModule.combine; +import static io.airlift.configuration.ConfigBinder.configBinder; public record YdbPlugin(Module module) implements Plugin { private static final String NAME = "ydb"; @@ -21,6 +23,8 @@ public Iterable getConnectorFactories() { NAME, () -> combine( new CredentialProviderModule(), + binder -> configBinder(binder).bindConfigDefaults( + JdbcMetadataConfig.class, config -> config.setComplexJoinPushdownEnabled(false)), module ) )); diff --git a/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbQueryBuilder.java b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbQueryBuilder.java new file mode 100644 index 00000000..6ab88d69 --- /dev/null +++ b/ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbQueryBuilder.java @@ -0,0 +1,41 @@ +package tech.ydb.trino; + +import com.google.inject.Inject; +import io.trino.plugin.jdbc.DefaultQueryBuilder; +import io.trino.plugin.jdbc.JdbcClient; +import io.trino.plugin.jdbc.JdbcColumnHandle; +import io.trino.plugin.jdbc.JdbcJoinCondition; +import io.trino.plugin.jdbc.logging.RemoteQueryModifier; + +import static io.trino.spi.type.DoubleType.DOUBLE; +import static io.trino.spi.type.RealType.REAL; +import static java.lang.String.format; + +public class YdbQueryBuilder extends DefaultQueryBuilder { + @Inject + public YdbQueryBuilder(RemoteQueryModifier queryModifier) { + super(queryModifier); + } + + @Override + protected String formatJoinCondition(JdbcClient client, String leftAlias, String rightAlias, JdbcJoinCondition condition) { + return format("%s %s %s", + formatJoinKey(client, leftAlias, condition.getLeftColumn()), + condition.getOperator().getValue(), + formatJoinKey(client, rightAlias, condition.getRightColumn())); + } + + private String formatJoinKey(JdbcClient client, String alias, JdbcColumnHandle column) { + String reference = alias + "." + client.quoted(column.getColumnName()); + if (column.getJdbcTypeHandle().jdbcTypeName().filter("Uint64"::equalsIgnoreCase).isPresent()) { + return "BITCAST(" + reference + " AS Int64)"; + } + if (column.getColumnType().equals(REAL) || column.getColumnType().equals(DOUBLE)) { + // Match Trino's JOIN behavior for NaN and signed zero. + String type = column.getColumnType().equals(REAL) ? "Float" : "Double"; + return format("NANVL(IF(COALESCE(%1$s = 0, false), CAST(0 AS %2$s), %1$s), CAST(NULL AS %2$s))", + reference, type); + } + return reference; + } +} 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 fdde9656..8e9f2927 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 @@ -1,5 +1,6 @@ package tech.ydb.trino; +import io.trino.Session; import io.trino.spi.type.Type; import io.trino.spi.type.VarcharType; import io.trino.testing.BaseConnectorTest; @@ -13,6 +14,7 @@ import org.junit.jupiter.api.extension.RegisterExtension; import tech.ydb.test.junit5.YdbHelperExtension; +import java.util.List; import java.util.Optional; import java.util.OptionalInt; @@ -42,6 +44,54 @@ public void testDefaultSchema() { assertQueryFails("SELECT * FROM local.\"%\".orders", ".*Schema '%' does not exist"); } + @Test + public void testOptInInt64JoinPushdown() { + Session session = Session.builder(getSession()) + .setCatalogSessionProperty(getSession().getCatalog().orElseThrow(), "join_pushdown_enabled", "true") + .build(); + try (TestTable left = newTrinoTable("join_int64_left_", "(id bigint, k bigint, second_key bigint, d double, s varchar)", + List.of("1, 7, 1, 1.0, 'a'", "2, 7, 2, 2.0, 'b'", "3, NULL, 1, 3.0, 'c'", "4, 8, 1, 4.0, 'd'")); + TestTable right = newTrinoTable("join_int64_right_", "(id bigint, k bigint, second_key bigint, d double, s varchar)", + List.of("10, 7, 1, 1.0, 'a'", "11, 7, 2, 2.0, 'b'", "12, NULL, 1, 3.0, 'c'", "13, 9, 1, 4.0, 'd'"))) { + String join = "SELECT l.id, r.id FROM " + left.getName() + " l %s " + right.getName() + " r ON %s"; + String matches = "VALUES (BIGINT '1', BIGINT '10'), (BIGINT '1', BIGINT '11'), " + + "(BIGINT '2', BIGINT '10'), (BIGINT '2', BIGINT '11')"; + assertThat(query(session, join.formatted("JOIN", "l.k = r.k"))).isFullyPushedDown().matches(matches); + assertThat(query(session, join.formatted("LEFT JOIN", "l.k = r.k"))).isFullyPushedDown() + .matches(matches + ", (BIGINT '3', CAST(NULL AS BIGINT)), (BIGINT '4', CAST(NULL AS BIGINT))"); + assertThat(query(session, join.formatted("RIGHT JOIN", "l.k = r.k"))).isFullyPushedDown() + .matches(matches + ", (CAST(NULL AS BIGINT), BIGINT '12'), (CAST(NULL AS BIGINT), BIGINT '13')"); + assertThat(query(session, join.formatted("FULL JOIN", "l.k = r.k"))).isFullyPushedDown() + .matches(matches + ", (BIGINT '3', CAST(NULL AS BIGINT)), (BIGINT '4', CAST(NULL AS BIGINT)), " + + "(CAST(NULL AS BIGINT), BIGINT '12'), (CAST(NULL AS BIGINT), BIGINT '13')"); + + assertThat(query(session, join.formatted("JOIN", "l.k = r.k AND l.second_key = r.second_key"))) + .isFullyPushedDown() + .matches("VALUES (BIGINT '1', BIGINT '10'), (BIGINT '2', BIGINT '11')"); + assertThat(query(session, "SELECT l.id, r.id FROM (SELECT id, k FROM " + left.getName() + + " WHERE id > 1 AND id < 3) l JOIN (SELECT id, k FROM " + right.getName() + + " WHERE id > 10 AND id < 12) r ON l.k = r.k")) + .isFullyPushedDown() + .matches("VALUES (BIGINT '2', BIGINT '11')"); + + assertThat(query(getSession(), join.formatted("JOIN", "l.k = r.k"))).joinIsNotFullyPushedDown(); + assertThat(query(session, join.formatted("JOIN", "l.k < r.k"))) + .joinIsNotFullyPushedDown(); + assertThat(query(session, join.formatted("JOIN", "l.s = r.s"))) + .matches("VALUES (BIGINT '1', BIGINT '10'), (BIGINT '2', BIGINT '11'), " + + "(BIGINT '3', BIGINT '12'), (BIGINT '4', BIGINT '13')") + .isFullyPushedDown(); + assertThat(query(session, join.formatted("JOIN", "l.d = r.d"))) + .matches("VALUES (BIGINT '1', BIGINT '10'), (BIGINT '2', BIGINT '11'), " + + "(BIGINT '3', BIGINT '12'), (BIGINT '4', BIGINT '13')") + .isFullyPushedDown(); + Session complex = Session.builder(session) + .setCatalogSessionProperty(getSession().getCatalog().orElseThrow(), "complex_join_pushdown_enabled", "true") + .build(); + assertThat(query(complex, join.formatted("JOIN", "l.k = r.k"))).joinIsNotFullyPushedDown(); + } + } + @Override protected void verifyConcurrentAddColumnFailurePermissible(Exception exception) { // YDB serializes overlapping scheme operations on one table and reports the conflict as retryable OVERLOADED: 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 9b380331..3d947848 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 @@ -22,6 +22,11 @@ public class TestingYdbJdbcModule implements Module { @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)