From e001913acb03af42ac820eb308a3df085c26233c Mon Sep 17 00:00:00 2001 From: Igor Artamonov Date: Thu, 28 Oct 2021 23:46:32 -0400 Subject: [PATCH] solution: AccessLog for Websocket requests --- .../monitoring/accesslog/AccessHandlerGrpc.kt | 2 +- .../monitoring/accesslog/AccessHandlerHttp.kt | 161 ++++++++++++++---- .../dshackle/monitoring/accesslog/Events.kt | 2 +- .../monitoring/accesslog/EventsBuilder.kt | 107 ++++++++++-- .../emeraldpay/dshackle/proxy/ProxyServer.kt | 2 +- .../dshackle/proxy/WebsocketHandler.kt | 25 ++- .../proxy/WebsocketHandlerSpec.groovy | 11 +- 7 files changed, 258 insertions(+), 52 deletions(-) diff --git a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerGrpc.kt b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerGrpc.kt index 7e3c849f..aac2d8c7 100644 --- a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerGrpc.kt +++ b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerGrpc.kt @@ -130,7 +130,7 @@ class AccessHandlerGrpc( ): ServerCall.Listener { return process( call, headers, next, - EventsBuilder.NativeSubscribe() as EventsBuilder.RequestReply<*, ReqT, RespT> + EventsBuilder.NativeSubscribe(Events.Channel.GRPC) as EventsBuilder.RequestReply<*, ReqT, RespT> ) } diff --git a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerHttp.kt b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerHttp.kt index eb468ddf..503361dc 100644 --- a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerHttp.kt +++ b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/AccessHandlerHttp.kt @@ -8,6 +8,7 @@ import org.slf4j.LoggerFactory import org.springframework.beans.factory.annotation.Autowired import org.springframework.stereotype.Service import reactor.netty.http.server.HttpServerRequest +import reactor.netty.http.websocket.WebsocketInbound import java.time.Instant import java.util.concurrent.locks.ReentrantLock import kotlin.concurrent.withLock @@ -25,6 +26,9 @@ class AccessHandlerHttp( companion object { private val log = LoggerFactory.getLogger(AccessHandlerHttp::class.java) + + private val NO_SUBSCRIBE = NoOnSubscriptionHandler() + private val NO_REQUEST = NoOpHandler() } /** @@ -38,11 +42,21 @@ class AccessHandlerHttp( interface HandlerFactory { fun create(req: HttpServerRequest, blockchain: Chain): RequestHandler + fun start(req: WebsocketInbound, blockchain: Chain): WsHandlerFactory + } + + interface WsHandlerFactory { + fun call(): RequestHandler + fun subscribe(): SubscriptionHandler } class NoOpFactory : HandlerFactory { override fun create(req: HttpServerRequest, blockchain: Chain): RequestHandler { - return NoOpHandler() + return NO_REQUEST + } + + override fun start(req: WebsocketInbound, blockchain: Chain): WsHandlerFactory { + return NO_REQUEST } } @@ -50,6 +64,10 @@ class AccessHandlerHttp( override fun create(req: HttpServerRequest, blockchain: Chain): RequestHandler { return StandardHandler(accessLogWriter, req, blockchain) } + + override fun start(req: WebsocketInbound, blockchain: Chain): WsHandlerFactory { + return StandardWsHandlerFactory(accessLogWriter, req, blockchain) + } } interface RequestHandler { @@ -58,7 +76,12 @@ class AccessHandlerHttp( fun onResponse(callResult: NativeCall.CallResult) } - class NoOpHandler : RequestHandler { + interface SubscriptionHandler { + fun onRequest(request: Pair) + fun onResponse(msgSize: Long) + } + + class NoOpHandler : RequestHandler, WsHandlerFactory { override fun close() { } @@ -67,37 +90,47 @@ class AccessHandlerHttp( override fun onResponse(callResult: NativeCall.CallResult) { } + + override fun call(): RequestHandler { + return this + } + + override fun subscribe(): SubscriptionHandler { + return NO_SUBSCRIBE + } } - class StandardHandler( - private val accessLogWriter: AccessLogWriter, - private val httpRequest: HttpServerRequest, - private val blockchain: Chain - ) : RequestHandler { - - private var request: BlockchainOuterClass.NativeCallRequest? = null - private val responses = ArrayList() - private val updateLock = ReentrantLock() - - override fun close() { - if (request == null) { - return - } - val responseTime = Instant.now() - val builder = EventsBuilder.NativeCall() - builder.withChain(blockchain.id) - builder.start(httpRequest) - builder.onRequest(request!!) - responses - .map { - builder.onReply(it, Events.Channel.JSONRPC).also { item -> - // since for JSON RPC you get a single response then the timestamp of all items included in it must have the same timestamp - item.ts = responseTime - } - } - .let(accessLogWriter::submit) + class NoOnSubscriptionHandler : SubscriptionHandler { + override fun onRequest(request: Pair) { } + override fun onResponse(msgSize: Long) { + } + } + + class StandardWsHandlerFactory( + private val accessLogWriter: AccessLogWriter, + private val wsRequest: WebsocketInbound, + private val blockchain: Chain + ) : WsHandlerFactory { + + override fun call(): RequestHandler { + return WsRequestHandler(accessLogWriter, wsRequest, blockchain) + } + + override fun subscribe(): SubscriptionHandler { + return WsSubscriptionHandler(accessLogWriter, wsRequest, blockchain) + } + } + + abstract class AbstractRequestHandler( + private val accessLogWriter: AccessLogWriter, + private val channel: Events.Channel + ) : RequestHandler { + protected var request: BlockchainOuterClass.NativeCallRequest? = null + protected val responses = ArrayList() + protected val updateLock = ReentrantLock() + override fun onRequest(request: BlockchainOuterClass.NativeCallRequest) { this.request = request } @@ -107,5 +140,75 @@ class AccessHandlerHttp( responses.add(callResult) } } + + fun onClose(builder: EventsBuilder.NativeCall) { + val responseTime = Instant.now() + responses + .map { + builder.onReply(it, channel).also { item -> + // since for JSON RPC you get a single response then the timestamp of all items included in it must have the same timestamp + item.ts = responseTime + } + } + .let(accessLogWriter::submit) + } + } + + class StandardHandler( + accessLogWriter: AccessLogWriter, + private val httpRequest: HttpServerRequest, + private val blockchain: Chain + ) : RequestHandler, AbstractRequestHandler(accessLogWriter, Events.Channel.JSONRPC) { + + override fun close() { + if (request == null) { + return + } + val builder = EventsBuilder.NativeCall() + builder.withChain(blockchain.id) + builder.start(httpRequest) + builder.onRequest(request!!) + onClose(builder) + } + } + + class WsRequestHandler( + accessLogWriter: AccessLogWriter, + private val wsRequest: WebsocketInbound, + private val blockchain: Chain + ) : RequestHandler, AbstractRequestHandler(accessLogWriter, Events.Channel.WSJSONRPC) { + + override fun close() { + if (request == null) { + return + } + val builder = EventsBuilder.NativeCall() + builder.withChain(blockchain.id) + builder.start(wsRequest) + builder.onRequest(request!!) + onClose(builder) + } + } + + class WsSubscriptionHandler( + private val accessLogWriter: AccessLogWriter, + private val wsRequest: WebsocketInbound, + private val blockchain: Chain + ) : SubscriptionHandler { + + private var builder: EventsBuilder.NativeSubscribeHttp? = null + + override fun onRequest(request: Pair) { + val builder = EventsBuilder.NativeSubscribeHttp(Events.Channel.WSJSONRPC, blockchain) + builder.start(wsRequest) + builder.onRequest(request) + this.builder = builder + } + + override fun onResponse(msgSize: Long) { + builder + ?.onReply(msgSize) + ?.let(accessLogWriter::submit) + } } } diff --git a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/Events.kt b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/Events.kt index 793fb820..d3d95268 100644 --- a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/Events.kt +++ b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/Events.kt @@ -28,7 +28,7 @@ class Events { } enum class Channel { - GRPC, JSONRPC + GRPC, JSONRPC, WSJSONRPC } abstract class Base( diff --git a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/EventsBuilder.kt b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/EventsBuilder.kt index ed2a5de7..375d1386 100644 --- a/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/EventsBuilder.kt +++ b/src/main/kotlin/io/emeraldpay/dshackle/monitoring/accesslog/EventsBuilder.kt @@ -21,9 +21,11 @@ import io.emeraldpay.grpc.Chain import io.grpc.Attributes import io.grpc.Grpc import io.grpc.Metadata +import io.netty.handler.codec.http.HttpHeaders import org.apache.commons.lang3.StringUtils import org.slf4j.LoggerFactory import reactor.netty.http.server.HttpServerRequest +import reactor.netty.http.websocket.WebsocketInbound import java.net.InetAddress import java.net.InetSocketAddress import java.time.Instant @@ -44,12 +46,16 @@ class EventsBuilder { fun start(request: HttpServerRequest) } + interface StartingWsRequest { + fun start(request: WebsocketInbound) + } + interface RequestReply : StartingHttp2Request { fun onRequest(msg: Req) fun onReply(msg: Resp): E } - abstract class Base : StartingHttp2Request, StartingHttp1Request { + abstract class Base : StartingHttp2Request, StartingHttp1Request, StartingWsRequest { companion object { private val remoteIpHeaders = listOf( "x-real-ip", @@ -136,17 +142,9 @@ class EventsBuilder { override fun start(request: HttpServerRequest) { val headers = request.requestHeaders() - val userAgent = headers.get("user-agent") - ?.let(this@Base::clean) - ?: "" + val userAgent = getUserAgent(headers) val ips = ArrayList() - remoteIpHeaders.forEach { key -> - headers.get(key)?.let { - it.trim().ifEmpty { null } - ?.let(this@Base::toInetAddress) - ?.let(ips::add) - } - } + extractIps(headers, ips) request.remoteAddress()?.let { addr -> ips.add(addr.address) } @@ -161,6 +159,53 @@ class EventsBuilder { ) } + override fun start(request: WebsocketInbound) { + val headers = request.headers() + val userAgent = getUserAgent(headers) + val ips = ArrayList() + extractIps(headers, ips) + // class WebsocketServerOperations, which is an implementation for the Websocket server connection, has a remoteAddress method + // But the class, and it's parent HttpServerOperations, are both private and cannot be used directly, + // so we try to access the field via reflection when it's possible + val remoteAddress: InetSocketAddress? = request.javaClass.methods + .find { it.name == "remoteAddress" } + ?.let { + if (it.canAccess(request) || it.trySetAccessible()) { + it.invoke(request) as InetSocketAddress + } else { + null + } + } + remoteAddress?.let { addr -> + ips.add(addr.address) + } + val ip = findBestIp(ips)?.hostAddress ?: "" + this.requestDetails = this.requestDetails + .copy( + remote = Events.Remote( + ips = ips.map { it.hostAddress }, + ip = ip, + userAgent = userAgent + ) + ) + } + + fun getUserAgent(headers: HttpHeaders): String { + return headers.get("user-agent") + ?.let(this@Base::clean) + ?: "" + } + + fun extractIps(headers: HttpHeaders, ips: MutableList) { + remoteIpHeaders.forEach { key -> + headers.get(key)?.let { + it.trim().ifEmpty { null } + ?.let(this@Base::toInetAddress) + ?.let(ips::add) + } + } + } + fun withChain(chain: Int): T { this.chainId = chain this.chain = Chain.byId(chainId) @@ -301,7 +346,9 @@ class EventsBuilder { } } - class NativeSubscribe : + class NativeSubscribe( + val channel: Events.Channel + ) : Base(), RequestReply { var item: Events.NativeSubscribeItemDetails? = null @@ -331,6 +378,42 @@ class EventsBuilder { } } + class NativeSubscribeHttp( + val channel: Events.Channel, + chain: Chain, + ) : + Base(), + RequestReply, Long> { + var item: Events.NativeSubscribeItemDetails? = null + val replies = HashMap() + + init { + withChain(chain.id) + } + + override fun getT(): NativeSubscribeHttp { + return this + } + + override fun onRequest(msg: Pair) { + this.item = Events.NativeSubscribeItemDetails( + msg.first, + msg.second?.size?.toLong() ?: 0L + ) + } + + override fun onReply(msg: Long): Events.NativeSubscribe { + return Events.NativeSubscribe( + request = requestDetails, + blockchain = chain, + nativeSubscribe = item!!, + payloadSizeBytes = msg, + id = UUID.randomUUID(), + channel = channel + ) + } + } + class Describe : Base(), RequestReply { diff --git a/src/main/kotlin/io/emeraldpay/dshackle/proxy/ProxyServer.kt b/src/main/kotlin/io/emeraldpay/dshackle/proxy/ProxyServer.kt index 314a28ff..6b10b06a 100644 --- a/src/main/kotlin/io/emeraldpay/dshackle/proxy/ProxyServer.kt +++ b/src/main/kotlin/io/emeraldpay/dshackle/proxy/ProxyServer.kt @@ -82,7 +82,7 @@ class ProxyServer( private val httpHandler = HttpHandler(readRpcJson, writeRpcJson, nativeCall, accessHandler, requestMetrics) private val wsHandler: WebsocketHandler? = if (config.websocketEnabled) { - WebsocketHandler(readRpcJson, writeRpcJson, nativeCall, nativeSubscribe, requestMetrics) + WebsocketHandler(readRpcJson, writeRpcJson, nativeCall, nativeSubscribe, accessHandler, requestMetrics) } else null fun start() { diff --git a/src/main/kotlin/io/emeraldpay/dshackle/proxy/WebsocketHandler.kt b/src/main/kotlin/io/emeraldpay/dshackle/proxy/WebsocketHandler.kt index 567944fe..34385553 100644 --- a/src/main/kotlin/io/emeraldpay/dshackle/proxy/WebsocketHandler.kt +++ b/src/main/kotlin/io/emeraldpay/dshackle/proxy/WebsocketHandler.kt @@ -43,6 +43,7 @@ class WebsocketHandler( writeRpcJson: WriteRpcJson, nativeCall: NativeCall, private val nativeSubscribe: NativeSubscribe, + private val accessHandler: AccessHandlerHttp.HandlerFactory, requestMetrics: ProxyServer.RequestMetricsFactory, ) : BaseHandler(writeRpcJson, nativeCall, requestMetrics) { @@ -64,7 +65,8 @@ class WebsocketHandler( .map { ByteBufInputStream(it.content()).readAllBytes() } .flatMap(this@WebsocketHandler::parseRequest) - val eventHandler: AccessHandlerHttp.RequestHandler = AccessHandlerHttp.NoOpHandler() + val eventHandler = accessHandler.start(req, routeConfig.blockchain) + val responses = respond(routeConfig.blockchain, requests, eventHandler) .map { Unpooled.wrappedBuffer(it.toByteArray()) } @@ -95,13 +97,22 @@ class WebsocketHandler( } } - fun respond(blockchain: Chain, requests: Flux>, eventHandler: AccessHandlerHttp.RequestHandler): Flux { + fun respond(blockchain: Chain, requests: Flux>, eventHandlerFactory: AccessHandlerHttp.WsHandlerFactory): Flux { return requests.flatMap { call -> val method = call.method + if (method == "eth_subscribe") { val methodParams = splitMethodParams(call.params) if (methodParams != null) { + val eventHandler: AccessHandlerHttp.SubscriptionHandler = eventHandlerFactory.subscribe() val subscriptionId = nextSubscriptionId() + eventHandler.onRequest( + methodParams.let { mp -> + // TODO ineffective to encode the params each time just to get size, ideally should get a reference to the original JSON bytes + // but it doesn't happen very ofter, only on initial subscribe only for logs with filter + Pair(mp.first, mp.second?.let { Global.objectMapper.writeValueAsBytes(it) }) + } + ) // first need to respond with ID of the subscription, and the following responses would have it in "subscription" param val start = ResponseJson().also { it.id = call.id @@ -115,12 +126,20 @@ class WebsocketHandler( } Flux.concat(Mono.just(start), responses) .map { Global.objectMapper.writeValueAsString(it) } + .doOnNext { + eventHandler.onResponse(it.length.toLong()) + } } else { + // TODO should it produce a 404 to the AccessLog? Mono.empty() } } else { + val eventHandler: AccessHandlerHttp.RequestHandler = eventHandlerFactory.call() val proxyCall = readRpcJson.convertToNativeCall(ProxyCall.RpcType.SINGLE, listOf(call)) - execute(blockchain, proxyCall, eventHandler) + Mono.from(execute(blockchain, proxyCall, eventHandler)) + // thought the event handler is used in execute + // it still needs to be closed at the end, so it can render the logs + .doFinally { eventHandler.close() } } } } diff --git a/src/test/groovy/io/emeraldpay/dshackle/proxy/WebsocketHandlerSpec.groovy b/src/test/groovy/io/emeraldpay/dshackle/proxy/WebsocketHandlerSpec.groovy index 35027283..1ca8415d 100644 --- a/src/test/groovy/io/emeraldpay/dshackle/proxy/WebsocketHandlerSpec.groovy +++ b/src/test/groovy/io/emeraldpay/dshackle/proxy/WebsocketHandlerSpec.groovy @@ -27,12 +27,13 @@ import java.time.Duration class WebsocketHandlerSpec extends Specification { + def requestHandlerFactory = new AccessHandlerHttp.NoOpFactory() def requestHandler = new AccessHandlerHttp.NoOpHandler() def "Parse standard RPC request"() { setup: def handler = new WebsocketHandler( - new ReadRpcJson(), Stub(WriteRpcJson), Stub(NativeCall), Stub(NativeSubscribe), Stub(ProxyServer.RequestMetricsFactory) + new ReadRpcJson(), Stub(WriteRpcJson), Stub(NativeCall), Stub(NativeSubscribe), requestHandlerFactory, Stub(ProxyServer.RequestMetricsFactory) ) when: def act = handler.parseRequest('{"id": 5, "jsonrpc": "2.0", "method": "eth_getBlockByNumber", "params": ["0x100001", false]}'.bytes) @@ -47,7 +48,7 @@ class WebsocketHandlerSpec extends Specification { def "Parse to empty an invalid request"() { setup: def handler = new WebsocketHandler( - new ReadRpcJson(), Stub(WriteRpcJson), Stub(NativeCall), Stub(NativeSubscribe), Stub(ProxyServer.RequestMetricsFactory) + new ReadRpcJson(), Stub(WriteRpcJson), Stub(NativeCall), Stub(NativeSubscribe), requestHandlerFactory, Stub(ProxyServer.RequestMetricsFactory) ) when: def act = handler.parseRequest('hello world'.bytes) @@ -61,7 +62,7 @@ class WebsocketHandlerSpec extends Specification { setup: def req1 = '{"id": 5, "jsonrpc": "2.0", "method": "eth_getBlockByNumber", "params": ["0x100001", false]}' def handler = new WebsocketHandler( - new ReadRpcJson(), Stub(WriteRpcJson), Stub(NativeCall), Stub(NativeSubscribe), Stub(ProxyServer.RequestMetricsFactory) + new ReadRpcJson(), Stub(WriteRpcJson), Stub(NativeCall), Stub(NativeSubscribe), requestHandlerFactory, Stub(ProxyServer.RequestMetricsFactory) ) when: def act = handler.parseRequest("[$req1]".bytes) @@ -79,7 +80,7 @@ class WebsocketHandlerSpec extends Specification { 1 * it.nativeCallResult(_) >> Flux.fromIterable([response]) } def handler = new WebsocketHandler( - new ReadRpcJson(), new WriteRpcJson(), nativeCall, Stub(NativeSubscribe), Stub(ProxyServer.RequestMetricsFactory) + new ReadRpcJson(), new WriteRpcJson(), nativeCall, Stub(NativeSubscribe), requestHandlerFactory, Stub(ProxyServer.RequestMetricsFactory) ) def request = new RequestJson("foo_test", [], 2) @@ -100,7 +101,7 @@ class WebsocketHandlerSpec extends Specification { 1 * it.subscribe(Chain.ETHEREUM, "foo_test", null) >> Flux.fromIterable([response1, response2]) } def handler = new WebsocketHandler( - new ReadRpcJson(), new WriteRpcJson(), Stub(NativeCall), nativeSubscribe, Stub(ProxyServer.RequestMetricsFactory) + new ReadRpcJson(), new WriteRpcJson(), Stub(NativeCall), nativeSubscribe, requestHandlerFactory, Stub(ProxyServer.RequestMetricsFactory) ) def request = new RequestJson("eth_subscribe", ["foo_test"], 2)