diff --git a/build.gradle b/build.gradle index 0e7f51fb..62a107db 100644 --- a/build.gradle +++ b/build.gradle @@ -132,6 +132,7 @@ dependencies { testImplementation "io.projectreactor:reactor-test:$reactorVersion" testImplementation 'org.objenesis:objenesis:3.1' testImplementation 'org.mock-server:mockserver-netty:5.11.2' + testImplementation "org.java-websocket:Java-WebSocket:1.5.1" testImplementation "nl.jqno.equalsverifier:equalsverifier:3.3" testImplementation "org.codehaus.groovy:groovy:${groovyVersion}" } diff --git a/src/main/kotlin/io/emeraldpay/dshackle/upstream/ethereum/EthereumWsFactory.kt b/src/main/kotlin/io/emeraldpay/dshackle/upstream/ethereum/EthereumWsFactory.kt index 3ff2883e..6d24e20b 100644 --- a/src/main/kotlin/io/emeraldpay/dshackle/upstream/ethereum/EthereumWsFactory.kt +++ b/src/main/kotlin/io/emeraldpay/dshackle/upstream/ethereum/EthereumWsFactory.kt @@ -77,6 +77,14 @@ class EthereumWsFactory( private const val START_REQUEST = "{\"jsonrpc\":\"2.0\", \"method\":\"eth_subscribe\", \"id\":\"blocks\", \"params\":[\"newHeads\"]}" } + var retryInterval = Defaults.retryConnection.seconds + set(value) { + if (retryInterval <= 0) { + throw IllegalArgumentException("Reconnect interval cannot be zero or less: $retryInterval") + } + field = value + } + private val parser = ResponseWSParser() private val blocks = Sinks @@ -103,15 +111,24 @@ class EthereumWsFactory( } private fun tryReconnectLater() { + if (!keepConnection) { + return + } + log.info("Reconnect to $uri in $retryInterval seconds...") Global.control.schedule( { connectInternal() }, - Defaults.retryConnection.seconds, TimeUnit.SECONDS) + retryInterval, TimeUnit.SECONDS) } private fun connectInternal() { log.info("Connecting to WebSocket: $uri") connection?.dispose() connection = HttpClient.create() + .doOnDisconnected { + if (keepConnection) { + tryReconnectLater() + } + } .doOnError( { _, t -> log.warn("Failed to connect to $uri. Error: ${t.message}") @@ -120,6 +137,7 @@ class EthereumWsFactory( }, { _, _ -> } ) + .headers { headers -> headers.add(HttpHeaderNames.ORIGIN, origin) basicAuth?.let { auth -> @@ -141,8 +159,9 @@ class EthereumWsFactory( .handle { inbound, outbound -> handle(inbound, outbound) } - .doOnError { - log.error("Failed to setup WS connection", it) + .onErrorResume { t -> + log.debug("Dropping WS connection to $uri. Error: ${t.message}") + Mono.empty() } .subscribe() } diff --git a/src/test/groovy/io/emeraldpay/dshackle/test/MockWSServer.groovy b/src/test/groovy/io/emeraldpay/dshackle/test/MockWSServer.groovy new file mode 100644 index 00000000..5a6ad76a --- /dev/null +++ b/src/test/groovy/io/emeraldpay/dshackle/test/MockWSServer.groovy @@ -0,0 +1,90 @@ +package io.emeraldpay.dshackle.test + +import com.fasterxml.jackson.databind.util.ByteBufferBackedInputStream +import org.java_websocket.WebSocket +import org.java_websocket.handshake.ClientHandshake +import org.java_websocket.server.WebSocketServer +import org.joda.time.format.DateTimeFormat + +import java.nio.ByteBuffer +import java.time.Instant +import java.time.LocalDate +import java.time.ZoneId +import java.time.format.DateTimeFormatter +import java.time.format.DateTimeFormatterBuilder + +class MockWSServer extends WebSocketServer { + + private def format = DateTimeFormatter.ofPattern("HH:mm:ss.SSS") + + List received = [] + private WebSocket conn + + private String next + + MockWSServer(int port) { + super(new InetSocketAddress("127.0.0.1", port)) + } + + void log(String msg) { + println(format.format(Instant.now().atZone(ZoneId.systemDefault())) + " MOCKWS: " + msg) + } + + void reply(String message) { + log(">> $message") + if (conn == null) { + println("MOCKWS: ERROR, no active connection") + } + conn.send(message) + } + + void onNextReply(String message) { + next = message + } + + @Override + void onOpen(WebSocket conn, ClientHandshake handshake) { + this.conn = conn + log("Opened connection from ${conn.remoteSocketAddress}") + } + + @Override + void onClose(WebSocket conn, int code, String reason, boolean remote) { + this.conn = null + log("Connection closed, code ${code} with msg '${reason}' ${remote ? 'by remote' : 'by server'}") + } + + @Override + void onMessage(WebSocket conn, String message) { + log("<< $message") + received.add(new ReceivedMessage(message)) + if (next != null) { + reply(next) + next = null + } + } + + @Override + void onMessage(WebSocket conn, ByteBuffer message) { + onMessage(conn, new ByteBufferBackedInputStream(message).text) + } + + @Override + void onError(WebSocket conn, Exception ex) { + log("ERROR, $ex.message") + received.add(new ReceivedMessage("Err: ${ex.message}")) + } + + @Override + void onStart() { + log("Server started") + } + + class ReceivedMessage { + final String value + + ReceivedMessage(String value) { + this.value = value + } + } +} diff --git a/src/test/groovy/io/emeraldpay/dshackle/upstream/ethereum/EthereumWsFactoryRealSpec.groovy b/src/test/groovy/io/emeraldpay/dshackle/upstream/ethereum/EthereumWsFactoryRealSpec.groovy new file mode 100644 index 00000000..ca2a3b59 --- /dev/null +++ b/src/test/groovy/io/emeraldpay/dshackle/upstream/ethereum/EthereumWsFactoryRealSpec.groovy @@ -0,0 +1,108 @@ +package io.emeraldpay.dshackle.upstream.ethereum + +import io.emeraldpay.dshackle.test.MockWSServer +import io.emeraldpay.dshackle.upstream.rpcclient.JsonRpcRequest +import reactor.test.StepVerifier +import spock.lang.Shared +import spock.lang.Specification + +import java.time.Duration + +class EthereumWsFactoryRealSpec extends Specification { + + // needs large timeouts and sleep, especially on CI where it's much slower to run + static TIMEOUT = 15 + static SLEEP = 500 + + static int port = 19900 + new Random().nextInt(100) + @Shared + MockWSServer server + @Shared + EthereumWsFactory.EthereumWs conn + + def setup() { + port++ + server = new MockWSServer(port) + server.start() + Thread.sleep(SLEEP) + conn = new EthereumWsFactory("ws://localhost:${port}".toURI(), "http://localhost:${port}".toURI()).create(null) + } + + def cleanup() { + conn.close() + server.stop() + } + + def "Connects to server"() { + when: + conn.connect() + Thread.sleep(SLEEP) + println("verify....") + def act = server.received + then: + act.size() > 0 + act[0].value.contains("\"method\":\"eth_subscribe\"") + act[0].value.contains("\"params\":[\"newHeads\"]") + } + + def "Makes RPC request"() { + when: + conn.connect() + def resp = conn.call(new JsonRpcRequest("foo_bar", [])) + then: + StepVerifier.create(resp) + .then { + server.reply('{"jsonrpc":"2.0", "id":100, "result": "baz"}') + } + .expectNextMatches { + it.hasResult() && it.resultAsProcessedString == "baz" + } + .expectComplete() + .verify(Duration.ofSeconds(3)) + + when: + Thread.sleep(SLEEP) + def act = server.received + then: + act.size() == 2 + act[1].value.contains("\"method\":\"foo_bar\"") + } + + def "Reconnects after server disconnect"() { + when: + conn.connect() + conn.retryInterval = 2 + Thread.sleep(SLEEP) + server.stop() + Thread.sleep(SLEEP) + server = new MockWSServer(port) + server.start() + def resp = conn.call(new JsonRpcRequest("foo_bar", [])) + // reconnects in 2 seconds, give 1 extra + Thread.sleep(3_000) + def act = server.received + + then: + act.size() > 0 + act[0].value.contains("\"method\":\"eth_subscribe\"") + act[0].value.contains("\"params\":[\"newHeads\"]") + } + + def "Try to connects to server until it's available"() { + when: + server.stop() + Thread.sleep(SLEEP) + conn.retryInterval = 1 + conn.connect() + Thread.sleep(3_000) + server = new MockWSServer(port) + server.start() + Thread.sleep(2_000) + def act = server.received + then: + act.size() > 0 + act[0].value.contains("\"method\":\"eth_subscribe\"") + act[0].value.contains("\"params\":[\"newHeads\"]") + } + +}