Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 7 additions & 14 deletions ydb-trino-adapter/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
```

## Каталог и схема
Expand Down
Original file line number Diff line number Diff line change
@@ -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;
Expand Down Expand Up @@ -49,16 +51,26 @@ public Pattern<Call> getPattern() {

@Override
public Optional<ParameterizedExpression> rewrite(Call call, Captures captures, RewriteContext<ParameterizedExpression> context) {
if (!(call.getArguments().get(1) instanceof Constant rightConstant) ||
!(rightConstant.getValue() instanceof Number number) ||
number.longValue() == 0 || number.longValue() == -1) {
return Optional.empty();
}
Optional<ParameterizedExpression> left = context.defaultRewrite(call.getArguments().getFirst());
if (left.isEmpty()) {
return Optional.empty();
}
Optional<ParameterizedExpression> 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.<QueryParameter>builder()
.addAll(left.get().parameters())
.addAll(right.get().parameters())
.build()));
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -42,21 +41,9 @@ public Pattern<? extends ConnectorExpression> getPattern() {

@Override
public Optional<JdbcExpression> rewrite(ConnectorTableHandle handle, ConnectorExpression projectionExpression, Captures captures, RewriteContext<ParameterizedExpression> 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<ParameterizedExpression> rewrittenValue = context.rewriteExpression(valueExpr);
Optional<ParameterizedExpression> rewrittenValue = context.rewriteExpression(captures.get(VALUE));
if (rewrittenValue.isEmpty()) {
return Optional.empty();
}
Expand Down
40 changes: 0 additions & 40 deletions ydb-trino-adapter/src/main/java/tech/ydb/trino/RewriteUtils.java

This file was deleted.

106 changes: 26 additions & 80 deletions ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbClient.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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;
Expand All @@ -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;
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -206,7 +199,6 @@ public YdbClient(
.add(new ImplementCountDistinct(bigintTypeHandle, true))
.add(new ImplementSum(YdbTypeUtils::toTypeHandle))
.add(new ImplementAvgFloatingPoint())
.add(new ImplementAvgDecimal())
.build());
}

Expand Down Expand Up @@ -426,64 +418,30 @@ public Optional<ColumnMapping> toColumnMapping(
return getUnsupportedTypeHandling(session) == CONVERT_TO_VARCHAR ? mapToUnboundedVarchar(typeHandle) : Optional.empty();
}

Optional<ColumnMapping> 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() {
Expand All @@ -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) {
Expand Down Expand Up @@ -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<JdbcColumnHandle> primaryKeys = getPrimaryKeys(session, handle.getRequiredNamedRelation().getRemoteTableName());
return OptionalLong.of(executeReturningDml(session, connection, primaryKeys, preparedQuery));
}
catch (SQLException e) {
throw new TrinoException(JDBC_ERROR, e);
Expand All @@ -858,8 +806,9 @@ public OptionalLong delete(ConnectorSession session, JdbcTableHandle handle) {

@Override
public OptionalLong update(ConnectorSession session, JdbcTableHandle handle) {
List<JdbcColumnHandle> 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(
Expand All @@ -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);
Expand All @@ -887,11 +836,8 @@ private static void verifyNoPrimaryKeyUpdate(List<JdbcColumnHandle> primaryKeys,
private long executeReturningDml(
ConnectorSession session,
Connection connection,
JdbcTableHandle handle,
List<JdbcColumnHandle> primaryKeys,
PreparedQuery preparedQuery) throws SQLException {
List<JdbcColumnHandle> primaryKeys = getPrimaryKeys(
session,
handle.getRequiredNamedRelation().getRemoteTableName());
if (primaryKeys.isEmpty()) {
throw new TrinoException(NOT_SUPPORTED, "YDB DML requires a table primary key");
}
Expand Down
18 changes: 0 additions & 18 deletions ydb-trino-adapter/src/main/java/tech/ydb/trino/YdbMetadata.java

This file was deleted.

Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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);
}
}
Loading
Loading