add custom headers as upstream settings (#761)

* add custom headers

* add docs
This commit is contained in:
Vyacheslav
2025-12-18 20:19:02 +02:00
committed by GitHub
parent 3e845bcdf0
commit b944aa8ac6
16 changed files with 196 additions and 16 deletions

View File

@@ -174,6 +174,41 @@ In case of rpc and ws connection we can specify different modes of works togethe
You can specify this modes through `connector-mode` parameter in connection config. You can specify this modes through `connector-mode` parameter in connection config.
=== Custom Headers
Dshackle allows you to add custom HTTP headers to all requests sent to an upstream.
This is useful when connecting to providers that require API keys, authentication tokens, or other custom headers.
Custom headers are configured at the upstream level and will be added to:
- All HTTP JSON RPC requests
- WebSocket connection handshakes (initial HTTP upgrade request)
NOTE: Custom headers are supported for all connection types except gRPC connections.
==== Configuration
Custom headers are specified using the `custom-headers` parameter in your upstream configuration:
[source,yaml]
----
version: v1
upstreams:
- id: my-ethereum-node
chain: ethereum
custom-headers:
X-API-Key: "your-api-key-here"
X-Custom-Header: "custom-value"
Authorization: "Bearer your-token"
connection:
generic:
rpc:
url: "https://api.example.com/rpc"
ws:
url: "wss://api.example.com/ws"
----
=== Bitcoin Methods === Bitcoin Methods
.By default an ethereum upstream allows call to the following JSON RPC methods: .By default an ethereum upstream allows call to the following JSON RPC methods:

View File

@@ -41,6 +41,7 @@ data class UpstreamsConfig(
var methods: Methods? = null, var methods: Methods? = null,
var methodGroups: MethodGroups? = null, var methodGroups: MethodGroups? = null,
var role: UpstreamRole = UpstreamRole.PRIMARY, var role: UpstreamRole = UpstreamRole.PRIMARY,
var customHeaders: Map<String, String> = emptyMap(),
) { ) {
@Suppress("UNCHECKED_CAST") @Suppress("UNCHECKED_CAST")

View File

@@ -300,6 +300,15 @@ class UpstreamsConfigReader(
} }
} }
} }
if (hasAny(upNode, "custom-headers")) {
getMapping(upNode, "custom-headers")?.let { headers ->
val headersMap = headers.value
.map { it.keyNode.valueAsString() to it.valueNode.valueAsString() }
.filter { StringUtils.isNotBlank(it.first) && StringUtils.isNotBlank(it.second) }
.associate { it.first!!.trim() to it.second!!.trim() }
upstream.customHeaders = headersMap
}
}
} }
private fun readUpstreamGrpc( private fun readUpstreamGrpc(

View File

@@ -38,7 +38,7 @@ class BitcoinUpstreamCreator(
): UpstreamCreationData { ): UpstreamCreationData {
val config = upstreamsConfig.cast(UpstreamsConfig.BitcoinConnection::class.java) val config = upstreamsConfig.cast(UpstreamsConfig.BitcoinConnection::class.java)
val conn = config.connection!! val conn = config.connection!!
val httpFactory = genericConnectorFactoryCreator.buildHttpFactory(conn.rpc) val httpFactory = genericConnectorFactoryCreator.buildHttpFactory(conn.rpc, customHeaders = config.customHeaders)
if (httpFactory == null) { if (httpFactory == null) {
log.warn("Upstream doesn't have API configuration") log.warn("Upstream doesn't have API configuration")
return UpstreamCreationData.default() return UpstreamCreationData.default()

View File

@@ -19,9 +19,14 @@ interface ConnectorFactoryCreator {
forkChoice: ForkChoice, forkChoice: ForkChoice,
blockValidator: BlockValidator, blockValidator: BlockValidator,
chainsConf: ChainsConfig.ChainConfig, chainsConf: ChainsConfig.ChainConfig,
customHeaders: Map<String, String> = emptyMap(),
): ConnectorFactory? ): ConnectorFactory?
fun buildHttpFactory(conn: UpstreamsConfig.HttpEndpoint?, urls: ArrayList<URI>? = null): HttpFactory? fun buildHttpFactory(
conn: UpstreamsConfig.HttpEndpoint?,
urls: ArrayList<URI>? = null,
customHeaders: Map<String, String> = emptyMap(),
): HttpFactory?
} }
@Component @Component

View File

@@ -38,10 +38,11 @@ open class GenericConnectorFactoryCreator(
forkChoice: ForkChoice, forkChoice: ForkChoice,
blockValidator: BlockValidator, blockValidator: BlockValidator,
chainsConf: ChainsConfig.ChainConfig, chainsConf: ChainsConfig.ChainConfig,
customHeaders: Map<String, String>,
): ConnectorFactory? { ): ConnectorFactory? {
val urls = ArrayList<URI>() val urls = ArrayList<URI>()
val wsFactoryApi = buildWsFactory(id, chain, conn, urls) val wsFactoryApi = buildWsFactory(id, chain, conn, urls, customHeaders)
val httpFactory = buildHttpFactory(conn.rpc, urls) val httpFactory = buildHttpFactory(conn.rpc, urls, customHeaders)
log.info("Using ${chain.chainName} upstream, at ${urls.joinToString()}") log.info("Using ${chain.chainName} upstream, at ${urls.joinToString()}")
val connectorFactory = val connectorFactory =
GenericConnectorFactory( GenericConnectorFactory(
@@ -62,7 +63,11 @@ open class GenericConnectorFactoryCreator(
return connectorFactory return connectorFactory
} }
override fun buildHttpFactory(conn: UpstreamsConfig.HttpEndpoint?, urls: ArrayList<URI>?): HttpFactory? { override fun buildHttpFactory(
conn: UpstreamsConfig.HttpEndpoint?,
urls: ArrayList<URI>?,
customHeaders: Map<String, String>,
): HttpFactory? {
return conn?.let { endpoint -> return conn?.let { endpoint ->
val tls = conn.tls?.let { tls -> val tls = conn.tls?.let { tls ->
tls.ca?.let { ca -> tls.ca?.let { ca ->
@@ -78,6 +83,7 @@ open class GenericConnectorFactoryCreator(
tls, tls,
monitoringCfg.nettyMetricsConfig.enabled, monitoringCfg.nettyMetricsConfig.enabled,
httpScheduler, httpScheduler,
customHeaders,
) )
} }
} }
@@ -87,6 +93,7 @@ open class GenericConnectorFactoryCreator(
chain: Chain, chain: Chain,
conn: UpstreamsConfig.RpcConnection, conn: UpstreamsConfig.RpcConnection,
urls: ArrayList<URI>? = null, urls: ArrayList<URI>? = null,
customHeaders: Map<String, String> = emptyMap(),
): WsConnectionPoolFactory? { ): WsConnectionPoolFactory? {
return conn.ws?.let { endpoint -> return conn.ws?.let { endpoint ->
val wsConnectionFactory = WsConnectionFactory( val wsConnectionFactory = WsConnectionFactory(
@@ -99,6 +106,7 @@ open class GenericConnectorFactoryCreator(
).apply { ).apply {
config = endpoint config = endpoint
basicAuth = endpoint.basicAuth basicAuth = endpoint.basicAuth
this.customHeaders = customHeaders
} }
val wsApi = WsConnectionPoolFactory( val wsApi = WsConnectionPoolFactory(
id, id,

View File

@@ -64,6 +64,7 @@ open class GenericUpstreamCreator(
NoChoiceWithPriorityForkChoice(nodeRating, config.id!!), NoChoiceWithPriorityForkChoice(nodeRating, config.id!!),
BlockValidator.ALWAYS_VALID, BlockValidator.ALWAYS_VALID,
chainConfig, chainConfig,
config.customHeaders,
) ?: return UpstreamCreationData.default() ) ?: return UpstreamCreationData.default()
val hashUrl = connection.let { val hashUrl = connection.let {

View File

@@ -41,10 +41,11 @@ class RestConnectorFactoryCreator(
forkChoice: ForkChoice, forkChoice: ForkChoice,
blockValidator: BlockValidator, blockValidator: BlockValidator,
chainsConf: ChainsConfig.ChainConfig, chainsConf: ChainsConfig.ChainConfig,
customHeaders: Map<String, String>,
): ConnectorFactory? { ): ConnectorFactory? {
val urls = ArrayList<URI>() val urls = ArrayList<URI>()
val httpFactory = buildHttpFactory(conn.rpc, urls) val httpFactory = buildHttpFactory(conn.rpc, urls, customHeaders)
val tonV3HttpFactory = buildHttpFactory(conn.getEndpointByTag("ton_v3")?.rpc, urls) val tonV3HttpFactory = buildHttpFactory(conn.getEndpointByTag("ton_v3")?.rpc, urls, customHeaders)
val upstreamHttpFactory = if (httpFactory != null && chain.type == BlockchainType.TON) { val upstreamHttpFactory = if (httpFactory != null && chain.type == BlockchainType.TON) {
TonCompoundHttpFactory(httpFactory, tonV3HttpFactory) TonCompoundHttpFactory(httpFactory, tonV3HttpFactory)
} else { } else {

View File

@@ -20,6 +20,7 @@ class BasicHttpFactory(
private val tls: ByteArray?, private val tls: ByteArray?,
private val nettyMetricsEnabled: Boolean, private val nettyMetricsEnabled: Boolean,
private val httpScheduler: Scheduler, private val httpScheduler: Scheduler,
private val customHeaders: Map<String, String> = emptyMap(),
) : HttpFactory { ) : HttpFactory {
private val log = LoggerFactory.getLogger(this::class.java) private val log = LoggerFactory.getLogger(this::class.java)
@@ -46,8 +47,8 @@ class BasicHttpFactory(
) )
if (chain.type.apiType == ApiType.REST) { if (chain.type.apiType == ApiType.REST) {
return RestHttpReader(url, maxConnections, queueSize, metrics, httpScheduler, chain, basicAuth, tls) return RestHttpReader(url, maxConnections, queueSize, metrics, httpScheduler, chain, basicAuth, tls, customHeaders)
} }
return JsonRpcHttpReader(url, maxConnections, queueSize, metrics, httpScheduler, basicAuth, tls) return JsonRpcHttpReader(url, maxConnections, queueSize, metrics, httpScheduler, basicAuth, tls, customHeaders)
} }
} }

View File

@@ -27,6 +27,7 @@ abstract class HttpReader(
protected val metrics: RequestMetrics?, protected val metrics: RequestMetrics?,
basicAuth: AuthConfig.ClientBasicAuth? = null, basicAuth: AuthConfig.ClientBasicAuth? = null,
tlsCAAuth: ByteArray? = null, tlsCAAuth: ByteArray? = null,
customHeaders: Map<String, String> = emptyMap(),
) : ChainReader { ) : ChainReader {
constructor() : this("", 1500, 1000, null) constructor() : this("", 1500, 1000, null)
@@ -65,6 +66,15 @@ abstract class HttpReader(
build = build.headers(headers) build = build.headers(headers)
} }
if (customHeaders.isNotEmpty()) {
val headers = Consumer { h: HttpHeaders ->
customHeaders.forEach { (key, value) ->
h.add(key, value)
}
}
build = build.headers(headers)
}
tlsCAAuth?.let { auth -> tlsCAAuth?.let { auth ->
val cf = CertificateFactory.getInstance("X.509") val cf = CertificateFactory.getInstance("X.509")
val cert = cf.generateCertificate(ByteArrayInputStream(auth)) as X509Certificate val cert = cf.generateCertificate(ByteArrayInputStream(auth)) as X509Certificate

View File

@@ -22,6 +22,7 @@ open class WsConnectionFactory(
var basicAuth: AuthConfig.ClientBasicAuth? = null var basicAuth: AuthConfig.ClientBasicAuth? = null
var config: UpstreamsConfig.WsEndpoint? = null var config: UpstreamsConfig.WsEndpoint? = null
var customHeaders: Map<String, String> = emptyMap()
private fun metrics(connIndex: Int): RequestMetrics { private fun metrics(connIndex: Int): RequestMetrics {
val metricsTags = listOf( val metricsTags = listOf(
@@ -46,7 +47,7 @@ open class WsConnectionFactory(
} }
open fun createWsConnection(connIndex: Int = 0): WsConnection = open fun createWsConnection(connIndex: Int = 0): WsConnection =
WsConnectionImpl(uri, origin, basicAuth, metrics(connIndex), scheduler, eventsScheduler).also { ws -> WsConnectionImpl(uri, origin, basicAuth, metrics(connIndex), scheduler, eventsScheduler, customHeaders).also { ws ->
config?.frameSize?.let { config?.frameSize?.let {
ws.frameSize = it ws.frameSize = it
} }

View File

@@ -66,6 +66,7 @@ open class WsConnectionImpl(
private val requestMetrics: RequestMetrics?, private val requestMetrics: RequestMetrics?,
private val scheduler: Scheduler, private val scheduler: Scheduler,
private val eventsScheduler: Scheduler, private val eventsScheduler: Scheduler,
private val customHeaders: Map<String, String> = emptyMap(),
) : AutoCloseable, WsConnection, Cloneable { ) : AutoCloseable, WsConnection, Cloneable {
companion object { companion object {
@@ -226,6 +227,9 @@ open class WsConnectionImpl(
val base64password = Base64.getEncoder().encodeToString(tmp.toByteArray()) val base64password = Base64.getEncoder().encodeToString(tmp.toByteArray())
headers.add(HttpHeaderNames.AUTHORIZATION, "Basic $base64password") headers.add(HttpHeaderNames.AUTHORIZATION, "Basic $base64password")
} }
customHeaders.forEach { (key, value) ->
headers.add(key, value)
}
} }
.let { .let {
if (uri.scheme == "wss") it.secure() else it if (uri.scheme == "wss") it.secure() else it

View File

@@ -33,7 +33,8 @@ class RestHttpReader(
private val chain: Chain, private val chain: Chain,
basicAuth: AuthConfig.ClientBasicAuth? = null, basicAuth: AuthConfig.ClientBasicAuth? = null,
tlsCAAuth: ByteArray? = null, tlsCAAuth: ByteArray? = null,
) : HttpReader(target, maxConnections, queueSize, metrics, basicAuth, tlsCAAuth) { customHeaders: Map<String, String> = emptyMap(),
) : HttpReader(target, maxConnections, queueSize, metrics, basicAuth, tlsCAAuth, customHeaders) {
private val parser = ResponseRpcParser() private val parser = ResponseRpcParser()
private val requestParser = RestRequestParser private val requestParser = RestRequestParser

View File

@@ -44,7 +44,8 @@ class JsonRpcHttpReader(
private val httpScheduler: Scheduler, private val httpScheduler: Scheduler,
basicAuth: AuthConfig.ClientBasicAuth? = null, basicAuth: AuthConfig.ClientBasicAuth? = null,
tlsCAAuth: ByteArray? = null, tlsCAAuth: ByteArray? = null,
) : HttpReader(target, maxConnections, queueSize, metrics, basicAuth, tlsCAAuth) { customHeaders: Map<String, String> = emptyMap(),
) : HttpReader(target, maxConnections, queueSize, metrics, basicAuth, tlsCAAuth, customHeaders) {
private val parser = ResponseRpcParser() private val parser = ResponseRpcParser()
private val streamParser = JsonRpcStreamParser() private val streamParser = JsonRpcStreamParser()

View File

@@ -53,7 +53,7 @@ class JsonRpcHttpReaderSpec extends Specification {
def "Make a request"() { def "Make a request"() {
setup: setup:
JsonRpcHttpReader client = new JsonRpcHttpReader("localhost:${port}", 50, 50, metrics, Schedulers.boundedElastic(),null, null) JsonRpcHttpReader client = new JsonRpcHttpReader("localhost:${port}", 50, 50, metrics, Schedulers.boundedElastic(),null, null, [:])
def resp = '{' + def resp = '{' +
' "jsonrpc": "2.0",' + ' "jsonrpc": "2.0",' +
' "result": "0x98de45",' + ' "result": "0x98de45",' +
@@ -74,8 +74,7 @@ class JsonRpcHttpReaderSpec extends Specification {
def "Produces RPC Exception on error status code"() { def "Produces RPC Exception on error status code"() {
setup: setup:
def client = new JsonRpcHttpReader("localhost:${port}", 50, 50, metrics, Schedulers.boundedElastic(), null, null) def client = new JsonRpcHttpReader("localhost:${port}", 50, 50, metrics, Schedulers.boundedElastic(), null, null, [:])
mockServer.when( mockServer.when(
HttpRequest.request() HttpRequest.request()
).respond( ).respond(
@@ -98,7 +97,7 @@ class JsonRpcHttpReaderSpec extends Specification {
def "Tries to extract message if HTTP error if it still contains a JSON RPC message"() { def "Tries to extract message if HTTP error if it still contains a JSON RPC message"() {
setup: setup:
def client = new JsonRpcHttpReader("localhost:${port}", 50, 50, metrics, Schedulers.boundedElastic(), null, null) def client = new JsonRpcHttpReader("localhost:${port}", 50, 50, metrics, Schedulers.boundedElastic(), null, null, [:])
mockServer.when( mockServer.when(
HttpRequest.request() HttpRequest.request()

View File

@@ -0,0 +1,103 @@
package io.emeraldpay.dshackle.config
import io.emeraldpay.dshackle.FileResolver
import io.emeraldpay.dshackle.foundation.ChainOptionsReader
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertNotNull
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
import java.io.File
class UpstreamsConfigReaderTest {
@Test
fun `should parse customHeaders from YAML`() {
val yaml = """
version: v1
upstreams:
- id: test-upstream
chain: ethereum
custom-headers:
X-Custom-Header: "custom-value"
Authorization: "Bearer token"
X-Another-Header: "another-value"
connection:
ethereum:
rpc:
url: "http://localhost:8545"
""".trimIndent()
val reader = UpstreamsConfigReader(
FileResolver(File(".")),
ChainOptionsReader(),
)
val config = reader.readInternal(yaml.byteInputStream())
assertNotNull(config)
assertEquals(1, config.upstreams.size)
val upstream = config.upstreams[0]
assertEquals("test-upstream", upstream.id)
assertEquals(3, upstream.customHeaders.size)
assertEquals("custom-value", upstream.customHeaders["X-Custom-Header"])
assertEquals("Bearer token", upstream.customHeaders["Authorization"])
assertEquals("another-value", upstream.customHeaders["X-Another-Header"])
}
@Test
fun `should work without customHeaders`() {
val yaml = """
version: v1
upstreams:
- id: test-upstream
chain: ethereum
connection:
ethereum:
rpc:
url: "http://localhost:8545"
""".trimIndent()
val reader = UpstreamsConfigReader(
FileResolver(File(".")),
ChainOptionsReader(),
)
val config = reader.readInternal(yaml.byteInputStream())
assertNotNull(config)
assertEquals(1, config.upstreams.size)
val upstream = config.upstreams[0]
assertEquals("test-upstream", upstream.id)
assertTrue(upstream.customHeaders.isEmpty())
}
@Test
fun `should trim header names and values`() {
val yaml = """
version: v1
upstreams:
- id: test-upstream
chain: ethereum
custom-headers:
" X-Header ": " value "
connection:
ethereum:
rpc:
url: "http://localhost:8545"
""".trimIndent()
val reader = UpstreamsConfigReader(
FileResolver(File(".")),
ChainOptionsReader(),
)
val config = reader.readInternal(yaml.byteInputStream())
assertNotNull(config)
val upstream = config.upstreams[0]
assertEquals(1, upstream.customHeaders.size)
assertEquals("value", upstream.customHeaders["X-Header"])
}
}