diff --git a/Sources/Redis/Client/RedisClient.swift b/Sources/Redis/Client/RedisClient.swift index 177277a..168219e 100644 --- a/Sources/Redis/Client/RedisClient.swift +++ b/Sources/Redis/Client/RedisClient.swift @@ -20,7 +20,7 @@ public final class RedisClient: DatabaseConnection, BasicWorker { /// The channel private let channel: Channel - + /// Creates a new Redis client on the provided data source and sink. init(queue: RedisCommandHandler, channel: Channel) { self.queue = queue diff --git a/Sources/Redis/RequestResponseHandler.swift b/Sources/Redis/RequestResponseHandler.swift index 854d76f..762459d 100644 --- a/Sources/Redis/RequestResponseHandler.swift +++ b/Sources/Redis/RequestResponseHandler.swift @@ -116,4 +116,14 @@ final class RequestResponseHandler: ChannelDuplexHandler { ctx.write(self.wrapOutboundOut(request), promise: promise) } } + + public func handlerAdded(ctx: ChannelHandlerContext) { + // handles returning "closed channel" for promises in flight when connection closes + ctx.channel.closeFuture.always { + for promise in self.promiseBuffer { + promise.fail(error: ChannelError.ioOnClosedChannel) + } + self.promiseBuffer.removeAll() + } + } } diff --git a/Tests/RedisTests/RedisDatabaseTests.swift b/Tests/RedisTests/RedisDatabaseTests.swift index 9fdc6cc..7b1562b 100644 --- a/Tests/RedisTests/RedisDatabaseTests.swift +++ b/Tests/RedisTests/RedisDatabaseTests.swift @@ -26,6 +26,34 @@ class RedisDatabaseTests: XCTestCase { try redis.delete("hello").wait() XCTAssertNil(try redis.get("hello", as: String.self).wait()) } + + func testDroppedConnection() throws { + let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) + let config = RedisClientConfig.makeTest() + let database = try RedisDatabase(config: config) + let redis = try database.newConnection(on: group).wait() + defer { redis.close() } + + let timeout: UInt32 = 1 + var dataReceived = false + var errorReceived = false + + let command = RedisData.array(["brpop", "hello", "\(timeout)"].map { RedisData(bulk: $0) }) + _ = redis.send(command).do { data in + dataReceived = true + }.catch { error in + errorReceived = true + } + + // Close the connection + redis.close() + + // Sleep for an extra second seconds to give the transaction time to complete + sleep(timeout+1) + + XCTAssertEqual(dataReceived, false) + XCTAssertEqual(errorReceived, true) + } func testSelect() throws { let group = MultiThreadedEventLoopGroup(numberOfThreads: 1)