solution: AccessLog for Websocket requests

This commit is contained in:
Igor Artamonov
2021-10-28 23:46:32 -04:00
parent 92657cf635
commit e001913acb
7 changed files with 258 additions and 52 deletions

View File

@@ -130,7 +130,7 @@ class AccessHandlerGrpc(
): ServerCall.Listener<ReqT> { ): ServerCall.Listener<ReqT> {
return process( return process(
call, headers, next, call, headers, next,
EventsBuilder.NativeSubscribe() as EventsBuilder.RequestReply<*, ReqT, RespT> EventsBuilder.NativeSubscribe(Events.Channel.GRPC) as EventsBuilder.RequestReply<*, ReqT, RespT>
) )
} }

View File

@@ -8,6 +8,7 @@ import org.slf4j.LoggerFactory
import org.springframework.beans.factory.annotation.Autowired import org.springframework.beans.factory.annotation.Autowired
import org.springframework.stereotype.Service import org.springframework.stereotype.Service
import reactor.netty.http.server.HttpServerRequest import reactor.netty.http.server.HttpServerRequest
import reactor.netty.http.websocket.WebsocketInbound
import java.time.Instant import java.time.Instant
import java.util.concurrent.locks.ReentrantLock import java.util.concurrent.locks.ReentrantLock
import kotlin.concurrent.withLock import kotlin.concurrent.withLock
@@ -25,6 +26,9 @@ class AccessHandlerHttp(
companion object { companion object {
private val log = LoggerFactory.getLogger(AccessHandlerHttp::class.java) 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 { interface HandlerFactory {
fun create(req: HttpServerRequest, blockchain: Chain): RequestHandler 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 { class NoOpFactory : HandlerFactory {
override fun create(req: HttpServerRequest, blockchain: Chain): RequestHandler { 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 { override fun create(req: HttpServerRequest, blockchain: Chain): RequestHandler {
return StandardHandler(accessLogWriter, req, blockchain) return StandardHandler(accessLogWriter, req, blockchain)
} }
override fun start(req: WebsocketInbound, blockchain: Chain): WsHandlerFactory {
return StandardWsHandlerFactory(accessLogWriter, req, blockchain)
}
} }
interface RequestHandler { interface RequestHandler {
@@ -58,7 +76,12 @@ class AccessHandlerHttp(
fun onResponse(callResult: NativeCall.CallResult) fun onResponse(callResult: NativeCall.CallResult)
} }
class NoOpHandler : RequestHandler { interface SubscriptionHandler {
fun onRequest(request: Pair<String, ByteArray?>)
fun onResponse(msgSize: Long)
}
class NoOpHandler : RequestHandler, WsHandlerFactory {
override fun close() { override fun close() {
} }
@@ -67,37 +90,47 @@ class AccessHandlerHttp(
override fun onResponse(callResult: NativeCall.CallResult) { override fun onResponse(callResult: NativeCall.CallResult) {
} }
override fun call(): RequestHandler {
return this
}
override fun subscribe(): SubscriptionHandler {
return NO_SUBSCRIBE
}
} }
class StandardHandler( class NoOnSubscriptionHandler : SubscriptionHandler {
private val accessLogWriter: AccessLogWriter, override fun onRequest(request: Pair<String, ByteArray?>) {
private val httpRequest: HttpServerRequest,
private val blockchain: Chain
) : RequestHandler {
private var request: BlockchainOuterClass.NativeCallRequest? = null
private val responses = ArrayList<NativeCall.CallResult>()
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)
} }
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<NativeCall.CallResult>()
protected val updateLock = ReentrantLock()
override fun onRequest(request: BlockchainOuterClass.NativeCallRequest) { override fun onRequest(request: BlockchainOuterClass.NativeCallRequest) {
this.request = request this.request = request
} }
@@ -107,5 +140,75 @@ class AccessHandlerHttp(
responses.add(callResult) 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<String, ByteArray?>) {
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)
}
} }
} }

View File

@@ -28,7 +28,7 @@ class Events {
} }
enum class Channel { enum class Channel {
GRPC, JSONRPC GRPC, JSONRPC, WSJSONRPC
} }
abstract class Base( abstract class Base(

View File

@@ -21,9 +21,11 @@ import io.emeraldpay.grpc.Chain
import io.grpc.Attributes import io.grpc.Attributes
import io.grpc.Grpc import io.grpc.Grpc
import io.grpc.Metadata import io.grpc.Metadata
import io.netty.handler.codec.http.HttpHeaders
import org.apache.commons.lang3.StringUtils import org.apache.commons.lang3.StringUtils
import org.slf4j.LoggerFactory import org.slf4j.LoggerFactory
import reactor.netty.http.server.HttpServerRequest import reactor.netty.http.server.HttpServerRequest
import reactor.netty.http.websocket.WebsocketInbound
import java.net.InetAddress import java.net.InetAddress
import java.net.InetSocketAddress import java.net.InetSocketAddress
import java.time.Instant import java.time.Instant
@@ -44,12 +46,16 @@ class EventsBuilder {
fun start(request: HttpServerRequest) fun start(request: HttpServerRequest)
} }
interface StartingWsRequest {
fun start(request: WebsocketInbound)
}
interface RequestReply<E, Req, Resp> : StartingHttp2Request { interface RequestReply<E, Req, Resp> : StartingHttp2Request {
fun onRequest(msg: Req) fun onRequest(msg: Req)
fun onReply(msg: Resp): E fun onReply(msg: Resp): E
} }
abstract class Base<T> : StartingHttp2Request, StartingHttp1Request { abstract class Base<T> : StartingHttp2Request, StartingHttp1Request, StartingWsRequest {
companion object { companion object {
private val remoteIpHeaders = listOf( private val remoteIpHeaders = listOf(
"x-real-ip", "x-real-ip",
@@ -136,17 +142,9 @@ class EventsBuilder {
override fun start(request: HttpServerRequest) { override fun start(request: HttpServerRequest) {
val headers = request.requestHeaders() val headers = request.requestHeaders()
val userAgent = headers.get("user-agent") val userAgent = getUserAgent(headers)
?.let(this@Base::clean)
?: ""
val ips = ArrayList<InetAddress>() val ips = ArrayList<InetAddress>()
remoteIpHeaders.forEach { key -> extractIps(headers, ips)
headers.get(key)?.let {
it.trim().ifEmpty { null }
?.let(this@Base::toInetAddress)
?.let(ips::add)
}
}
request.remoteAddress()?.let { addr -> request.remoteAddress()?.let { addr ->
ips.add(addr.address) 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<InetAddress>()
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<InetAddress>) {
remoteIpHeaders.forEach { key ->
headers.get(key)?.let {
it.trim().ifEmpty { null }
?.let(this@Base::toInetAddress)
?.let(ips::add)
}
}
}
fun withChain(chain: Int): T { fun withChain(chain: Int): T {
this.chainId = chain this.chainId = chain
this.chain = Chain.byId(chainId) this.chain = Chain.byId(chainId)
@@ -301,7 +346,9 @@ class EventsBuilder {
} }
} }
class NativeSubscribe : class NativeSubscribe(
val channel: Events.Channel
) :
Base<NativeSubscribe>(), Base<NativeSubscribe>(),
RequestReply<Events.NativeSubscribe, BlockchainOuterClass.NativeSubscribeRequest, BlockchainOuterClass.NativeSubscribeReplyItem> { RequestReply<Events.NativeSubscribe, BlockchainOuterClass.NativeSubscribeRequest, BlockchainOuterClass.NativeSubscribeReplyItem> {
var item: Events.NativeSubscribeItemDetails? = null var item: Events.NativeSubscribeItemDetails? = null
@@ -331,6 +378,42 @@ class EventsBuilder {
} }
} }
class NativeSubscribeHttp(
val channel: Events.Channel,
chain: Chain,
) :
Base<NativeSubscribeHttp>(),
RequestReply<Events.NativeSubscribe, Pair<String, ByteArray?>, Long> {
var item: Events.NativeSubscribeItemDetails? = null
val replies = HashMap<Int, Events.NativeSubscribeReplyDetails>()
init {
withChain(chain.id)
}
override fun getT(): NativeSubscribeHttp {
return this
}
override fun onRequest(msg: Pair<String, ByteArray?>) {
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 : class Describe :
Base<Describe>(), Base<Describe>(),
RequestReply<Events.Describe, BlockchainOuterClass.DescribeRequest, BlockchainOuterClass.DescribeResponse> { RequestReply<Events.Describe, BlockchainOuterClass.DescribeRequest, BlockchainOuterClass.DescribeResponse> {

View File

@@ -82,7 +82,7 @@ class ProxyServer(
private val httpHandler = HttpHandler(readRpcJson, writeRpcJson, nativeCall, accessHandler, requestMetrics) private val httpHandler = HttpHandler(readRpcJson, writeRpcJson, nativeCall, accessHandler, requestMetrics)
private val wsHandler: WebsocketHandler? = if (config.websocketEnabled) { private val wsHandler: WebsocketHandler? = if (config.websocketEnabled) {
WebsocketHandler(readRpcJson, writeRpcJson, nativeCall, nativeSubscribe, requestMetrics) WebsocketHandler(readRpcJson, writeRpcJson, nativeCall, nativeSubscribe, accessHandler, requestMetrics)
} else null } else null
fun start() { fun start() {

View File

@@ -43,6 +43,7 @@ class WebsocketHandler(
writeRpcJson: WriteRpcJson, writeRpcJson: WriteRpcJson,
nativeCall: NativeCall, nativeCall: NativeCall,
private val nativeSubscribe: NativeSubscribe, private val nativeSubscribe: NativeSubscribe,
private val accessHandler: AccessHandlerHttp.HandlerFactory,
requestMetrics: ProxyServer.RequestMetricsFactory, requestMetrics: ProxyServer.RequestMetricsFactory,
) : BaseHandler(writeRpcJson, nativeCall, requestMetrics) { ) : BaseHandler(writeRpcJson, nativeCall, requestMetrics) {
@@ -64,7 +65,8 @@ class WebsocketHandler(
.map { ByteBufInputStream(it.content()).readAllBytes() } .map { ByteBufInputStream(it.content()).readAllBytes() }
.flatMap(this@WebsocketHandler::parseRequest) .flatMap(this@WebsocketHandler::parseRequest)
val eventHandler: AccessHandlerHttp.RequestHandler = AccessHandlerHttp.NoOpHandler() val eventHandler = accessHandler.start(req, routeConfig.blockchain)
val responses = respond(routeConfig.blockchain, requests, eventHandler) val responses = respond(routeConfig.blockchain, requests, eventHandler)
.map { Unpooled.wrappedBuffer(it.toByteArray()) } .map { Unpooled.wrappedBuffer(it.toByteArray()) }
@@ -95,13 +97,22 @@ class WebsocketHandler(
} }
} }
fun respond(blockchain: Chain, requests: Flux<RequestJson<Any>>, eventHandler: AccessHandlerHttp.RequestHandler): Flux<String> { fun respond(blockchain: Chain, requests: Flux<RequestJson<Any>>, eventHandlerFactory: AccessHandlerHttp.WsHandlerFactory): Flux<String> {
return requests.flatMap { call -> return requests.flatMap { call ->
val method = call.method val method = call.method
if (method == "eth_subscribe") { if (method == "eth_subscribe") {
val methodParams = splitMethodParams(call.params) val methodParams = splitMethodParams(call.params)
if (methodParams != null) { if (methodParams != null) {
val eventHandler: AccessHandlerHttp.SubscriptionHandler = eventHandlerFactory.subscribe()
val subscriptionId = nextSubscriptionId() 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 // first need to respond with ID of the subscription, and the following responses would have it in "subscription" param
val start = ResponseJson<String, Any>().also { val start = ResponseJson<String, Any>().also {
it.id = call.id it.id = call.id
@@ -115,12 +126,20 @@ class WebsocketHandler(
} }
Flux.concat(Mono.just(start), responses) Flux.concat(Mono.just(start), responses)
.map { Global.objectMapper.writeValueAsString(it) } .map { Global.objectMapper.writeValueAsString(it) }
.doOnNext {
eventHandler.onResponse(it.length.toLong())
}
} else { } else {
// TODO should it produce a 404 to the AccessLog?
Mono.empty() Mono.empty()
} }
} else { } else {
val eventHandler: AccessHandlerHttp.RequestHandler = eventHandlerFactory.call()
val proxyCall = readRpcJson.convertToNativeCall(ProxyCall.RpcType.SINGLE, listOf(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() }
} }
} }
} }

View File

@@ -27,12 +27,13 @@ import java.time.Duration
class WebsocketHandlerSpec extends Specification { class WebsocketHandlerSpec extends Specification {
def requestHandlerFactory = new AccessHandlerHttp.NoOpFactory()
def requestHandler = new AccessHandlerHttp.NoOpHandler() def requestHandler = new AccessHandlerHttp.NoOpHandler()
def "Parse standard RPC request"() { def "Parse standard RPC request"() {
setup: setup:
def handler = new WebsocketHandler( 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: when:
def act = handler.parseRequest('{"id": 5, "jsonrpc": "2.0", "method": "eth_getBlockByNumber", "params": ["0x100001", false]}'.bytes) 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"() { def "Parse to empty an invalid request"() {
setup: setup:
def handler = new WebsocketHandler( 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: when:
def act = handler.parseRequest('hello world'.bytes) def act = handler.parseRequest('hello world'.bytes)
@@ -61,7 +62,7 @@ class WebsocketHandlerSpec extends Specification {
setup: setup:
def req1 = '{"id": 5, "jsonrpc": "2.0", "method": "eth_getBlockByNumber", "params": ["0x100001", false]}' def req1 = '{"id": 5, "jsonrpc": "2.0", "method": "eth_getBlockByNumber", "params": ["0x100001", false]}'
def handler = new WebsocketHandler( 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: when:
def act = handler.parseRequest("[$req1]".bytes) def act = handler.parseRequest("[$req1]".bytes)
@@ -79,7 +80,7 @@ class WebsocketHandlerSpec extends Specification {
1 * it.nativeCallResult(_) >> Flux.fromIterable([response]) 1 * it.nativeCallResult(_) >> Flux.fromIterable([response])
} }
def handler = new WebsocketHandler( 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) 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]) 1 * it.subscribe(Chain.ETHEREUM, "foo_test", null) >> Flux.fromIterable([response1, response2])
} }
def handler = new WebsocketHandler( 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) def request = new RequestJson("eth_subscribe", ["foo_test"], 2)