mirror of
https://github.com/Control-D-Inc/ctrld.git
synced 2026-08-23 01:27:13 +02:00
refactor(config): consolidate transport setup and eliminate duplication
Consolidate DoH/DoH3/DoQ transport initialization into a single SetupTransport method and introduce generic helper functions to eliminate duplicated IP stack selection logic across transport getters. This reduces code duplication by ~77 lines while maintaining the same functionality.
This commit is contained in:
@@ -9,7 +9,6 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"math/rand"
|
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -509,54 +508,49 @@ func (uc *UpstreamConfig) ReBootstrap() {
|
|||||||
// For now, only DoH upstream is supported.
|
// For now, only DoH upstream is supported.
|
||||||
func (uc *UpstreamConfig) SetupTransport() {
|
func (uc *UpstreamConfig) SetupTransport() {
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeDOH:
|
case ResolverTypeDOH, ResolverTypeDOH3, ResolverTypeDOQ:
|
||||||
uc.setupDOHTransport()
|
default:
|
||||||
case ResolverTypeDOH3:
|
return
|
||||||
uc.setupDOH3Transport()
|
|
||||||
case ResolverTypeDOQ:
|
|
||||||
uc.setupDOQTransport()
|
|
||||||
}
|
}
|
||||||
}
|
ips := uc.bootstrapIPs
|
||||||
|
|
||||||
func (uc *UpstreamConfig) setupDOQTransport() {
|
|
||||||
switch uc.IPStack {
|
switch uc.IPStack {
|
||||||
case IpStackBoth, "":
|
|
||||||
uc.doqConnPool = uc.newDOQConnPool(uc.bootstrapIPs)
|
|
||||||
case IpStackV4:
|
case IpStackV4:
|
||||||
uc.doqConnPool = uc.newDOQConnPool(uc.bootstrapIPs4)
|
ips = uc.bootstrapIPs4
|
||||||
case IpStackV6:
|
case IpStackV6:
|
||||||
uc.doqConnPool = uc.newDOQConnPool(uc.bootstrapIPs6)
|
ips = uc.bootstrapIPs6
|
||||||
case IpStackSplit:
|
}
|
||||||
|
uc.transport = uc.newDOHTransport(ips)
|
||||||
|
uc.http3RoundTripper = uc.newDOH3Transport(ips)
|
||||||
|
uc.doqConnPool = uc.newDOQConnPool(ips)
|
||||||
|
if uc.IPStack == IpStackSplit {
|
||||||
|
uc.transport4 = uc.newDOHTransport(uc.bootstrapIPs4)
|
||||||
|
uc.http3RoundTripper4 = uc.newDOH3Transport(uc.bootstrapIPs4)
|
||||||
uc.doqConnPool4 = uc.newDOQConnPool(uc.bootstrapIPs4)
|
uc.doqConnPool4 = uc.newDOQConnPool(uc.bootstrapIPs4)
|
||||||
if HasIPv6() {
|
if HasIPv6() {
|
||||||
|
uc.transport6 = uc.newDOHTransport(uc.bootstrapIPs6)
|
||||||
|
uc.http3RoundTripper6 = uc.newDOH3Transport(uc.bootstrapIPs6)
|
||||||
uc.doqConnPool6 = uc.newDOQConnPool(uc.bootstrapIPs6)
|
uc.doqConnPool6 = uc.newDOQConnPool(uc.bootstrapIPs6)
|
||||||
} else {
|
} else {
|
||||||
|
uc.transport6 = uc.transport4
|
||||||
|
uc.http3RoundTripper6 = uc.http3RoundTripper4
|
||||||
uc.doqConnPool6 = uc.doqConnPool4
|
uc.doqConnPool6 = uc.doqConnPool4
|
||||||
}
|
}
|
||||||
uc.doqConnPool = uc.newDOQConnPool(uc.bootstrapIPs)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) setupDOHTransport() {
|
func (uc *UpstreamConfig) ensureSetupTransport() {
|
||||||
switch uc.IPStack {
|
uc.transportOnce.Do(func() {
|
||||||
case IpStackBoth, "":
|
uc.SetupTransport()
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs)
|
})
|
||||||
case IpStackV4:
|
if uc.rebootstrap.CompareAndSwap(true, false) {
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs4)
|
uc.SetupTransport()
|
||||||
case IpStackV6:
|
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs6)
|
|
||||||
case IpStackSplit:
|
|
||||||
uc.transport4 = uc.newDOHTransport(uc.bootstrapIPs4)
|
|
||||||
if HasIPv6() {
|
|
||||||
uc.transport6 = uc.newDOHTransport(uc.bootstrapIPs6)
|
|
||||||
} else {
|
|
||||||
uc.transport6 = uc.transport4
|
|
||||||
}
|
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
|
func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
|
||||||
|
if uc.Type != ResolverTypeDOH {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
transport := http.DefaultTransport.(*http.Transport).Clone()
|
transport := http.DefaultTransport.(*http.Transport).Clone()
|
||||||
transport.MaxIdleConnsPerHost = 100
|
transport.MaxIdleConnsPerHost = 100
|
||||||
transport.TLSClientConfig = &tls.Config{
|
transport.TLSClientConfig = &tls.Config{
|
||||||
@@ -690,46 +684,8 @@ func (uc *UpstreamConfig) isNextDNS() bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) dohTransport(dnsType uint16) http.RoundTripper {
|
func (uc *UpstreamConfig) dohTransport(dnsType uint16) http.RoundTripper {
|
||||||
uc.transportOnce.Do(func() {
|
uc.ensureSetupTransport()
|
||||||
uc.SetupTransport()
|
return transportByIpStack(uc.IPStack, dnsType, uc.transport, uc.transport4, uc.transport6)
|
||||||
})
|
|
||||||
if uc.rebootstrap.CompareAndSwap(true, false) {
|
|
||||||
uc.SetupTransport()
|
|
||||||
}
|
|
||||||
switch uc.IPStack {
|
|
||||||
case IpStackBoth, IpStackV4, IpStackV6:
|
|
||||||
return uc.transport
|
|
||||||
case IpStackSplit:
|
|
||||||
switch dnsType {
|
|
||||||
case dns.TypeA:
|
|
||||||
return uc.transport4
|
|
||||||
default:
|
|
||||||
return uc.transport6
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return uc.transport
|
|
||||||
}
|
|
||||||
|
|
||||||
func (uc *UpstreamConfig) bootstrapIPForDNSType(dnsType uint16) string {
|
|
||||||
switch uc.IPStack {
|
|
||||||
case IpStackBoth:
|
|
||||||
return pick(uc.bootstrapIPs)
|
|
||||||
case IpStackV4:
|
|
||||||
return pick(uc.bootstrapIPs4)
|
|
||||||
case IpStackV6:
|
|
||||||
return pick(uc.bootstrapIPs6)
|
|
||||||
case IpStackSplit:
|
|
||||||
switch dnsType {
|
|
||||||
case dns.TypeA:
|
|
||||||
return pick(uc.bootstrapIPs4)
|
|
||||||
default:
|
|
||||||
if HasIPv6() {
|
|
||||||
return pick(uc.bootstrapIPs6)
|
|
||||||
}
|
|
||||||
return pick(uc.bootstrapIPs4)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pick(uc.bootstrapIPs)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) {
|
func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) {
|
||||||
@@ -974,10 +930,6 @@ func ResolverTypeFromEndpoint(endpoint string) string {
|
|||||||
return ResolverTypeDOT
|
return ResolverTypeDOT
|
||||||
}
|
}
|
||||||
|
|
||||||
func pick(s []string) string {
|
|
||||||
return s[rand.Intn(len(s))]
|
|
||||||
}
|
|
||||||
|
|
||||||
// upstreamUID generates an unique identifier for an upstream.
|
// upstreamUID generates an unique identifier for an upstream.
|
||||||
func upstreamUID() string {
|
func upstreamUID() string {
|
||||||
b := make([]byte, 4)
|
b := make([]byte, 4)
|
||||||
@@ -1013,3 +965,18 @@ func bootstrapIPsFromControlDDomain(domain string) []string {
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func transportByIpStack[T any](ipStack string, dnsType uint16, transport, transport4, transport6 T) T {
|
||||||
|
switch ipStack {
|
||||||
|
case IpStackBoth, IpStackV4, IpStackV6:
|
||||||
|
return transport
|
||||||
|
case IpStackSplit:
|
||||||
|
switch dnsType {
|
||||||
|
case dns.TypeA:
|
||||||
|
return transport4
|
||||||
|
default:
|
||||||
|
return transport6
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return transport
|
||||||
|
}
|
||||||
|
|||||||
+10
-37
@@ -9,7 +9,6 @@ import (
|
|||||||
"runtime"
|
"runtime"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/miekg/dns"
|
|
||||||
"github.com/quic-go/quic-go"
|
"github.com/quic-go/quic-go"
|
||||||
"github.com/quic-go/quic-go/http3"
|
"github.com/quic-go/quic-go/http3"
|
||||||
)
|
)
|
||||||
@@ -34,6 +33,9 @@ func (uc *UpstreamConfig) setupDOH3Transport() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
|
func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
|
||||||
|
if uc.Type != ResolverTypeDOH3 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
rt := &http3.Transport{}
|
rt := &http3.Transport{}
|
||||||
rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool}
|
rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool}
|
||||||
rt.Dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
rt.Dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||||
@@ -71,45 +73,13 @@ func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) doh3Transport(dnsType uint16) http.RoundTripper {
|
func (uc *UpstreamConfig) doh3Transport(dnsType uint16) http.RoundTripper {
|
||||||
uc.transportOnce.Do(func() {
|
uc.ensureSetupTransport()
|
||||||
uc.SetupTransport()
|
return transportByIpStack(uc.IPStack, dnsType, uc.http3RoundTripper, uc.http3RoundTripper4, uc.http3RoundTripper6)
|
||||||
})
|
|
||||||
if uc.rebootstrap.CompareAndSwap(true, false) {
|
|
||||||
uc.SetupTransport()
|
|
||||||
}
|
|
||||||
switch uc.IPStack {
|
|
||||||
case IpStackBoth, IpStackV4, IpStackV6:
|
|
||||||
return uc.http3RoundTripper
|
|
||||||
case IpStackSplit:
|
|
||||||
switch dnsType {
|
|
||||||
case dns.TypeA:
|
|
||||||
return uc.http3RoundTripper4
|
|
||||||
default:
|
|
||||||
return uc.http3RoundTripper6
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return uc.http3RoundTripper
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) doqTransport(dnsType uint16) *doqConnPool {
|
func (uc *UpstreamConfig) doqTransport(dnsType uint16) *doqConnPool {
|
||||||
uc.transportOnce.Do(func() {
|
uc.ensureSetupTransport()
|
||||||
uc.SetupTransport()
|
return transportByIpStack(uc.IPStack, dnsType, uc.doqConnPool, uc.doqConnPool4, uc.doqConnPool6)
|
||||||
})
|
|
||||||
if uc.rebootstrap.CompareAndSwap(true, false) {
|
|
||||||
uc.SetupTransport()
|
|
||||||
}
|
|
||||||
switch uc.IPStack {
|
|
||||||
case IpStackBoth, IpStackV4, IpStackV6:
|
|
||||||
return uc.doqConnPool
|
|
||||||
case IpStackSplit:
|
|
||||||
switch dnsType {
|
|
||||||
case dns.TypeA:
|
|
||||||
return uc.doqConnPool4
|
|
||||||
default:
|
|
||||||
return uc.doqConnPool6
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return uc.doqConnPool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Putting the code for quic parallel dialer here:
|
// Putting the code for quic parallel dialer here:
|
||||||
@@ -181,5 +151,8 @@ func (d *quicParallelDialer) Dial(ctx context.Context, addrs []string, tlsCfg *t
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) newDOQConnPool(addrs []string) *doqConnPool {
|
func (uc *UpstreamConfig) newDOQConnPool(addrs []string) *doqConnPool {
|
||||||
|
if uc.Type != ResolverTypeDOQ {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
return newDOQConnPool(uc, addrs)
|
return newDOQConnPool(uc, addrs)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -86,7 +86,7 @@ func newDOQConnPool(uc *UpstreamConfig, addrs []string) *doqConnPool {
|
|||||||
// Resolve performs a DNS query using a pooled QUIC connection.
|
// Resolve performs a DNS query using a pooled QUIC connection.
|
||||||
func (p *doqConnPool) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error) {
|
func (p *doqConnPool) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error) {
|
||||||
// Retry logic for io.EOF errors (as per original implementation)
|
// Retry logic for io.EOF errors (as per original implementation)
|
||||||
for i := 0; i < 5; i++ {
|
for range 5 {
|
||||||
answer, err := p.doResolve(ctx, msg)
|
answer, err := p.doResolve(ctx, msg)
|
||||||
if err == io.EOF {
|
if err == io.EOF {
|
||||||
continue
|
continue
|
||||||
|
|||||||
Reference in New Issue
Block a user