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:
Cuong Manh Le
2026-03-03 14:22:32 +07:00
committed by Cuong Manh Le
parent e8d1a4604e
commit 1f4c47318e
3 changed files with 54 additions and 114 deletions
+43 -76
View File
@@ -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
View File
@@ -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)
} }
+1 -1
View File
@@ -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