diff --git a/src/main/java/org/springframework/data/redis/connection/lettuce/LettuceConnection.java b/src/main/java/org/springframework/data/redis/connection/lettuce/LettuceConnection.java index cdf4c7157..21d7dbc51 100644 --- a/src/main/java/org/springframework/data/redis/connection/lettuce/LettuceConnection.java +++ b/src/main/java/org/springframework/data/redis/connection/lettuce/LettuceConnection.java @@ -37,11 +37,13 @@ import io.lettuce.core.output.*; import io.lettuce.core.protocol.Command; import io.lettuce.core.protocol.CommandArgs; import io.lettuce.core.protocol.CommandType; +import io.lettuce.core.protocol.ProtocolKeyword; import io.lettuce.core.pubsub.StatefulRedisPubSubConnection; import io.lettuce.core.sentinel.api.StatefulRedisSentinelConnection; import lombok.RequiredArgsConstructor; import java.lang.reflect.Constructor; +import java.nio.charset.StandardCharsets; import java.util.ArrayList; import java.util.Collections; import java.util.HashMap; @@ -405,7 +407,7 @@ public class LettuceConnection extends AbstractRedisConnection { try { String name = command.trim().toUpperCase(); - CommandType commandType = CommandType.valueOf(name); + ProtocolKeyword commandType = getCommandType(name); validateCommandIfRunningInTransactionMode(commandType, args); @@ -1046,14 +1048,14 @@ public class LettuceConnection extends AbstractRedisConnection { return io.lettuce.core.ScanCursor.of(Long.toString(cursorId)); } - private void validateCommandIfRunningInTransactionMode(CommandType cmd, byte[]... args) { + private void validateCommandIfRunningInTransactionMode(ProtocolKeyword cmd, byte[]... args) { if (this.isQueueing()) { validateCommand(cmd, args); } } - private void validateCommand(CommandType cmd, @Nullable byte[]... args) { + private void validateCommand(ProtocolKeyword cmd, @Nullable byte[]... args) { RedisCommand redisCommand = RedisCommand.failsafeCommandLookup(cmd.name()); if (!RedisCommand.UNKNOWN.equals(redisCommand) && redisCommand.requiresArguments()) { @@ -1106,6 +1108,15 @@ public class LettuceConnection extends AbstractRedisConnection { return connectionProvider; } + private static ProtocolKeyword getCommandType(String name) { + + try { + return CommandType.valueOf(name); + } catch (IllegalArgumentException e) { + return new CustomCommandType(name); + } + } + /** * {@link TypeHints} provide {@link CommandOutput} information for a given {@link CommandType}. * @@ -1114,7 +1125,7 @@ public class LettuceConnection extends AbstractRedisConnection { static class TypeHints { @SuppressWarnings("rawtypes") // - private static final Map> COMMAND_OUTPUT_TYPE_MAPPING = new HashMap<>(); + private static final Map> COMMAND_OUTPUT_TYPE_MAPPING = new HashMap<>(); @SuppressWarnings("rawtypes") // private static final Map, Constructor> CONSTRUCTORS = new ConcurrentHashMap<>(); @@ -1275,7 +1286,7 @@ public class LettuceConnection extends AbstractRedisConnection { * @return {@link ByteArrayOutput} as default when no matching {@link CommandOutput} available. */ @SuppressWarnings("rawtypes") - public CommandOutput getTypeHint(CommandType type) { + public CommandOutput getTypeHint(ProtocolKeyword type) { return getTypeHint(type, new ByteArrayOutput<>(CODEC)); } @@ -1286,7 +1297,7 @@ public class LettuceConnection extends AbstractRedisConnection { * @return */ @SuppressWarnings("rawtypes") - public CommandOutput getTypeHint(CommandType type, CommandOutput defaultType) { + public CommandOutput getTypeHint(ProtocolKeyword type, CommandOutput defaultType) { if (type == null || !COMMAND_OUTPUT_TYPE_MAPPING.containsKey(type)) { return defaultType; @@ -1523,4 +1534,49 @@ public class LettuceConnection extends AbstractRedisConnection { connection.setAutoFlushCommands(true); } } + + /** + * @since 2.3.8 + */ + static class CustomCommandType implements ProtocolKeyword { + + private final String name; + + CustomCommandType(String name) { + this.name = name; + } + + @Override + public byte[] getBytes() { + return name.getBytes(StandardCharsets.US_ASCII); + } + + @Override + public String name() { + return name; + } + + @Override + public boolean equals(Object o) { + + if (this == o) { + return true; + } + if (!(o instanceof CustomCommandType)) { + return false; + } + CustomCommandType that = (CustomCommandType) o; + return ObjectUtils.nullSafeEquals(name, that.name); + } + + @Override + public int hashCode() { + return ObjectUtils.nullSafeHashCode(name); + } + + @Override + public String toString() { + return name; + } + } } diff --git a/src/test/java/org/springframework/data/redis/connection/lettuce/LettuceConnectionUnitTestSuite.java b/src/test/java/org/springframework/data/redis/connection/lettuce/LettuceConnectionUnitTestSuite.java index cb2c5d72f..2919aed8c 100644 --- a/src/test/java/org/springframework/data/redis/connection/lettuce/LettuceConnectionUnitTestSuite.java +++ b/src/test/java/org/springframework/data/redis/connection/lettuce/LettuceConnectionUnitTestSuite.java @@ -19,11 +19,17 @@ import static org.assertj.core.api.Assertions.*; import static org.mockito.Mockito.*; import io.lettuce.core.RedisClient; +import io.lettuce.core.RedisFuture; import io.lettuce.core.XAddArgs; import io.lettuce.core.api.StatefulRedisConnection; import io.lettuce.core.api.async.RedisAsyncCommands; import io.lettuce.core.api.sync.RedisCommands; +import io.lettuce.core.codec.ByteArrayCodec; import io.lettuce.core.codec.RedisCodec; +import io.lettuce.core.output.StatusOutput; +import io.lettuce.core.protocol.AsyncCommand; +import io.lettuce.core.protocol.Command; +import io.lettuce.core.protocol.CommandArgs; import java.lang.reflect.InvocationTargetException; import java.util.Collections; @@ -32,8 +38,8 @@ import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; import org.junit.runners.Suite; - import org.mockito.ArgumentCaptor; + import org.springframework.dao.InvalidDataAccessResourceUsageException; import org.springframework.data.redis.connection.AbstractConnectionUnitTestBase; import org.springframework.data.redis.connection.RedisServerCommands.ShutdownOption; @@ -182,6 +188,21 @@ public class LettuceConnectionUnitTestSuite { assertThat(args.getValue()).extracting("maxlen").isEqualTo(100L); } + + @Test // GH-1979 + public void executeShouldPassThruCustomCommands() { + + Command command = new Command<>(new LettuceConnection.CustomCommandType("FOO.BAR"), + new StatusOutput<>(ByteArrayCodec.INSTANCE)); + AsyncCommand future = new AsyncCommand<>(command); + future.complete(); + + when(asyncCommandsMock.dispatch(any(), any(), any())).thenReturn((RedisFuture) future); + + connection.execute("foo.bar", command.getOutput()); + + verify(asyncCommandsMock).dispatch(eq(command.getType()), eq(command.getOutput()), any(CommandArgs.class)); + } } public static class LettucePipelineConnectionUnitTests extends LettuceConnectionUnitTests {