Improve OpenVPN & OpenConnect interoperability
This commit is contained in:
+100
-32
@@ -2,6 +2,7 @@ package openconnect
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"net"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
@@ -50,6 +51,8 @@ type Endpoint struct {
|
||||
flavor string
|
||||
stateAccess sync.Mutex
|
||||
state atomic.Pointer[clientState]
|
||||
dnsTransportAccess sync.Mutex
|
||||
dnsTransport *DNSTransport
|
||||
deviceStarted bool
|
||||
readLoopDone chan struct{}
|
||||
statusAccess sync.Mutex
|
||||
@@ -65,10 +68,22 @@ type clientState struct {
|
||||
tunnelConfigured bool
|
||||
localAddresses []netip.Prefix
|
||||
routeSet *netipx.IPSet
|
||||
preferredDomains map[string]bool
|
||||
configuration openconnecttransport.Configuration
|
||||
tunnelInfo adapter.OpenConnectTunnelInfo
|
||||
}
|
||||
|
||||
func NewEndpoint(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.OpenConnectEndpointOptions) (adapter.Endpoint, error) {
|
||||
tcpKeepAliveEnabled := options.TCPKeepAliveEnabled || options.TCPKeepAlive != 0 || options.TCPKeepAliveInterval != 0
|
||||
if tcpKeepAliveEnabled && options.DisableTCPKeepAlive {
|
||||
return nil, E.New("tcp_keep_alive_enabled conflicts with disable_tcp_keep_alive")
|
||||
}
|
||||
if !tcpKeepAliveEnabled {
|
||||
options.DisableTCPKeepAlive = true
|
||||
} else if options.TCPKeepAlive == 0 && options.TCPKeepAliveInterval == 0 {
|
||||
options.TCPKeepAliveSystemDefaults = true
|
||||
}
|
||||
options.UDPBindPort = options.DTLSLocalPort
|
||||
loopContext, cancelLoop := context.WithCancel(ctx)
|
||||
openConnectEndpoint := &Endpoint{
|
||||
endpointBase: endpointBase{
|
||||
@@ -98,7 +113,7 @@ func NewEndpoint(ctx context.Context, router adapter.Router, logger log.ContextL
|
||||
}
|
||||
serverURL, err := url.Parse(server)
|
||||
if err != nil {
|
||||
return nil, E.Cause(err, "parse OpenConnect server")
|
||||
return nil, E.Cause(err, "parse server")
|
||||
}
|
||||
serverPort := serverURL.Port()
|
||||
if serverPort == "" {
|
||||
@@ -162,6 +177,10 @@ func NewEndpoint(ctx context.Context, router adapter.Router, logger log.ContextL
|
||||
}
|
||||
|
||||
func (e *Endpoint) buildClientOptions(options option.OpenConnectEndpointOptions, outboundDialer N.Dialer) (openconnect.ClientOptions, error) {
|
||||
var tlsConfig *tls.Config
|
||||
if options.TLS.Insecure {
|
||||
tlsConfig = &tls.Config{InsecureSkipVerify: true}
|
||||
}
|
||||
certificateAuthority, err := materialSource("tls.certificate_authority", options.TLS.CertificateAuthority, options.TLS.CertificateAuthorityPath)
|
||||
if err != nil {
|
||||
return openconnect.ClientOptions{}, err
|
||||
@@ -185,12 +204,13 @@ func (e *Endpoint) buildClientOptions(options option.OpenConnectEndpointOptions,
|
||||
var tokenOptions *openconnect.TokenOptions
|
||||
if options.Token != nil {
|
||||
tokenOptions = &openconnect.TokenOptions{
|
||||
Mode: options.Token.Mode,
|
||||
Secret: options.Token.Secret,
|
||||
PIN: options.Token.PIN,
|
||||
Password: options.Token.Password,
|
||||
DeviceID: options.Token.DeviceID,
|
||||
Counter: options.Token.Counter,
|
||||
Mode: options.Token.Mode,
|
||||
Secret: options.Token.Secret,
|
||||
SecretPath: options.Token.SecretPath,
|
||||
PIN: options.Token.PIN,
|
||||
Password: options.Token.Password,
|
||||
DeviceID: options.Token.DeviceID,
|
||||
Counter: options.Token.Counter,
|
||||
}
|
||||
if tokenOptions.Mode == openconnect.TokenModeHOTP {
|
||||
e.hotpCounter.Store(tokenOptions.Counter)
|
||||
@@ -201,6 +221,14 @@ func (e *Endpoint) buildClientOptions(options option.OpenConnectEndpointOptions,
|
||||
}
|
||||
}
|
||||
var csdOptions *openconnect.CSDOptions
|
||||
var mobileOptions *openconnect.MobileOptions
|
||||
if options.Mobile != nil {
|
||||
mobileOptions = &openconnect.MobileOptions{
|
||||
PlatformVersion: options.Mobile.PlatformVersion,
|
||||
DeviceType: options.Mobile.DeviceType,
|
||||
DeviceUniqueID: options.Mobile.DeviceUniqueID,
|
||||
}
|
||||
}
|
||||
if options.CSD != nil {
|
||||
csdOptions = &openconnect.CSDOptions{WrapperPath: options.CSD.WrapperPath}
|
||||
}
|
||||
@@ -236,21 +264,44 @@ func (e *Endpoint) buildClientOptions(options option.OpenConnectEndpointOptions,
|
||||
}
|
||||
})
|
||||
return openconnect.ClientOptions{
|
||||
Context: e.loopContext,
|
||||
Server: options.Server,
|
||||
Flavor: options.Flavor,
|
||||
Username: options.Username,
|
||||
Password: options.Password,
|
||||
AuthGroup: options.AuthGroup,
|
||||
Token: tokenOptions,
|
||||
ReportedOS: options.ReportedOS,
|
||||
UserAgent: options.UserAgent,
|
||||
CSD: csdOptions,
|
||||
HIP: hipOptions,
|
||||
TNCC: tnccOptions,
|
||||
NoUDP: options.NoUDP,
|
||||
AllowInsecureCrypto: options.AllowInsecureCrypto,
|
||||
Context: e.loopContext,
|
||||
Server: options.Server,
|
||||
Flavor: options.Flavor,
|
||||
Username: options.Username,
|
||||
Password: options.Password,
|
||||
AuthGroup: options.AuthGroup,
|
||||
Cookie: options.Cookie,
|
||||
Token: tokenOptions,
|
||||
ReportedOS: options.ReportedOS,
|
||||
UserAgent: options.UserAgent,
|
||||
Version: options.Version,
|
||||
LocalHostname: options.LocalHostname,
|
||||
Mobile: mobileOptions,
|
||||
CSD: csdOptions,
|
||||
HIP: hipOptions,
|
||||
TNCC: tnccOptions,
|
||||
NoUDP: options.NoUDP,
|
||||
DTLSLocalPort: options.DTLSLocalPort,
|
||||
CompressionDisabled: options.CompressionDisabled,
|
||||
CompressionMode: options.CompressionMode,
|
||||
IPv6Disabled: options.IPv6Disabled,
|
||||
HTTPKeepAliveDisabled: options.HTTPKeepAliveDisabled,
|
||||
XMLPostDisabled: options.XMLPostDisabled,
|
||||
ExternalAuthDisabled: options.ExternalAuthDisabled,
|
||||
PasswordAuthenticationDisabled: options.PasswordAuthenticationDisabled,
|
||||
PFS: options.PFS,
|
||||
MTU: options.MTU,
|
||||
BaseMTU: options.BaseMTU,
|
||||
DPDInterval: time.Duration(options.DPDInterval),
|
||||
ReconnectTimeout: time.Duration(options.ReconnectTimeout),
|
||||
TrojanInterval: time.Duration(options.TrojanInterval),
|
||||
QueueLength: options.QueueLength,
|
||||
AllowInsecureCrypto: options.AllowInsecureCrypto,
|
||||
TLSConfig: openconnect.ClientTLSOptions{
|
||||
Config: tlsConfig,
|
||||
ServerName: options.TLS.ServerName,
|
||||
PeerFingerprints: options.TLS.PeerFingerprint,
|
||||
SystemTrustDisabled: options.TLS.SystemTrustDisabled,
|
||||
CertificateAuthority: certificateAuthority,
|
||||
Certificate: clientCertificate,
|
||||
Key: clientKey,
|
||||
@@ -274,7 +325,14 @@ func (e *Endpoint) handleTunnelConfiguration(event openconnect.TunnelConfigurati
|
||||
e.updateState(func(state *clientState) {
|
||||
state.tunnelConfigured = false
|
||||
})
|
||||
err := e.device.UpdateConfiguration(configuration)
|
||||
routeSet, err := buildIPSet(configuration.Routes, configuration.ExcludedRoutes)
|
||||
if err != nil {
|
||||
return E.Cause(err, "build route set")
|
||||
}
|
||||
err = e.device.UpdateConfiguration(openconnecttransport.Configuration{
|
||||
MTU: configuration.MTU,
|
||||
Addresses: configuration.Addresses,
|
||||
})
|
||||
if err != nil {
|
||||
return E.Cause(err, "update device configuration")
|
||||
}
|
||||
@@ -285,10 +343,7 @@ func (e *Endpoint) handleTunnelConfiguration(event openconnect.TunnelConfigurati
|
||||
}
|
||||
e.deviceStarted = true
|
||||
}
|
||||
routeSet, err := buildIPSet(configuration.Routes, configuration.ExcludedRoutes)
|
||||
if err != nil {
|
||||
return E.Cause(err, "build route set")
|
||||
}
|
||||
preferredDomains := buildPreferredDomains(configuration)
|
||||
var ipv4Addresses []netip.Prefix
|
||||
var ipv6Addresses []netip.Prefix
|
||||
for _, address := range configuration.Addresses {
|
||||
@@ -308,6 +363,8 @@ func (e *Endpoint) handleTunnelConfiguration(event openconnect.TunnelConfigurati
|
||||
state.tunnelConfigured = true
|
||||
state.localAddresses = configuration.Addresses
|
||||
state.routeSet = routeSet
|
||||
state.preferredDomains = preferredDomains
|
||||
state.configuration = configuration
|
||||
state.tunnelInfo = adapter.OpenConnectTunnelInfo{
|
||||
Server: e.server,
|
||||
Flavor: e.flavor,
|
||||
@@ -319,6 +376,12 @@ func (e *Endpoint) handleTunnelConfiguration(event openconnect.TunnelConfigurati
|
||||
ConnectedSince: connectedSince,
|
||||
}
|
||||
})
|
||||
e.dnsTransportAccess.Lock()
|
||||
dnsTransport := e.dnsTransport
|
||||
e.dnsTransportAccess.Unlock()
|
||||
if dnsTransport != nil {
|
||||
dnsTransport.updateConfiguration(configuration)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -358,7 +421,7 @@ func (e *Endpoint) readLoop() {
|
||||
if E.IsClosedOrCanceled(err) || e.loopContext.Err() != nil {
|
||||
return
|
||||
}
|
||||
e.logger.Error(E.Cause(err, "OpenConnect client terminated"))
|
||||
e.logger.Error(E.Cause(err, "client terminated"))
|
||||
e.setTerminalError(err)
|
||||
return
|
||||
}
|
||||
@@ -436,11 +499,11 @@ func (e *Endpoint) ready() bool {
|
||||
|
||||
func (e *Endpoint) WritePackets(packets [][]byte) error {
|
||||
if !e.ready() {
|
||||
return E.New("OpenConnect client is not ready yet")
|
||||
return E.New("endpoint is not ready yet")
|
||||
}
|
||||
err := e.client.WriteDataPackets(packets)
|
||||
if E.IsMulti(err, openconnect.ErrDataChannelNotReady) {
|
||||
return E.New("OpenConnect client is not ready yet")
|
||||
return E.New("endpoint is not ready yet")
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -473,7 +536,7 @@ func (e *Endpoint) DialContext(ctx context.Context, network string, destination
|
||||
e.logger.InfoContext(ctx, "outbound packet connection to ", destination)
|
||||
}
|
||||
if !e.ready() || !e.client.Ready() {
|
||||
return nil, E.New("OpenConnect client is not ready yet")
|
||||
return nil, E.New("endpoint is not ready yet")
|
||||
}
|
||||
if destination.IsDomain() {
|
||||
destinationAddresses, err := e.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
|
||||
@@ -491,7 +554,7 @@ func (e *Endpoint) DialContext(ctx context.Context, network string, destination
|
||||
func (e *Endpoint) ListenPacketWithDestination(ctx context.Context, destination M.Socksaddr) (net.PacketConn, netip.Addr, error) {
|
||||
e.logger.InfoContext(ctx, "outbound packet connection to ", destination)
|
||||
if !e.ready() || !e.client.Ready() {
|
||||
return nil, netip.Addr{}, E.New("OpenConnect client is not ready yet")
|
||||
return nil, netip.Addr{}, E.New("endpoint is not ready yet")
|
||||
}
|
||||
if destination.IsDomain() {
|
||||
destinationAddresses, err := e.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
|
||||
@@ -522,7 +585,12 @@ func (e *Endpoint) ListenPacket(ctx context.Context, destination M.Socksaddr) (n
|
||||
}
|
||||
|
||||
func (e *Endpoint) PreferredDomain(metadata *adapter.InboundContext, domain string) bool {
|
||||
return false
|
||||
state := e.state.Load()
|
||||
if !state.started || !state.tunnelConfigured || !e.client.Ready() {
|
||||
return false
|
||||
}
|
||||
canonicalDomain := canonicalOpenConnectDomain(domain)
|
||||
return openConnectDomainMatchesAny(canonicalDomain, state.preferredDomains)
|
||||
}
|
||||
|
||||
func (e *Endpoint) PreferredAddress(metadata *adapter.InboundContext, address netip.Addr) bool {
|
||||
|
||||
@@ -0,0 +1,373 @@
|
||||
package openconnect
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/sagernet/sing-box/adapter"
|
||||
C "github.com/sagernet/sing-box/constant"
|
||||
"github.com/sagernet/sing-box/dns"
|
||||
"github.com/sagernet/sing-box/dns/transport"
|
||||
"github.com/sagernet/sing-box/log"
|
||||
"github.com/sagernet/sing-box/option"
|
||||
openconnecttransport "github.com/sagernet/sing-box/transport/openconnect"
|
||||
"github.com/sagernet/sing/common"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/sagernet/sing/service"
|
||||
|
||||
mDNS "github.com/miekg/dns"
|
||||
)
|
||||
|
||||
func RegisterDNSTransport(registry *dns.TransportRegistry) {
|
||||
dns.RegisterTransport[option.OpenConnectDNSServerOptions](registry, C.DNSTypeOpenConnect, NewDNSTransport)
|
||||
}
|
||||
|
||||
type DNSTransport struct {
|
||||
dns.TransportAdapter
|
||||
logger logger.ContextLogger
|
||||
endpointTag string
|
||||
acceptDefaultResolvers bool
|
||||
acceptSearchDomain bool
|
||||
endpointManager adapter.EndpointManager
|
||||
endpoint *Endpoint
|
||||
dialer N.Dialer
|
||||
access sync.RWMutex
|
||||
closed bool
|
||||
routes []openConnectDNSRoute
|
||||
searchDomains []string
|
||||
defaultResolvers []adapter.DNSTransport
|
||||
}
|
||||
|
||||
type openConnectDNSRoute struct {
|
||||
domain string
|
||||
resolvers []adapter.DNSTransport
|
||||
}
|
||||
|
||||
func NewDNSTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.OpenConnectDNSServerOptions) (adapter.DNSTransport, error) {
|
||||
if options.Endpoint == "" {
|
||||
return nil, E.New("missing endpoint tag")
|
||||
}
|
||||
return &DNSTransport{
|
||||
TransportAdapter: dns.NewTransportAdapter(C.DNSTypeOpenConnect, tag, nil),
|
||||
logger: logger,
|
||||
endpointTag: options.Endpoint,
|
||||
acceptDefaultResolvers: options.AcceptDefaultResolvers,
|
||||
acceptSearchDomain: options.AcceptSearchDomain,
|
||||
endpointManager: service.FromContext[adapter.EndpointManager](ctx),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (t *DNSTransport) Start(stage adapter.StartStage) error {
|
||||
if stage != adapter.StartStateInitialize {
|
||||
return nil
|
||||
}
|
||||
rawEndpoint, loaded := t.endpointManager.Get(t.endpointTag)
|
||||
if !loaded {
|
||||
return E.New("endpoint not found: ", t.endpointTag)
|
||||
}
|
||||
openConnectEndpoint, isOpenConnect := rawEndpoint.(*Endpoint)
|
||||
if !isOpenConnect {
|
||||
return E.New("endpoint is not OpenConnect: ", t.endpointTag)
|
||||
}
|
||||
openConnectEndpoint.dnsTransportAccess.Lock()
|
||||
if openConnectEndpoint.dnsTransport != nil && openConnectEndpoint.dnsTransport.Tag() != t.Tag() {
|
||||
openConnectEndpoint.dnsTransportAccess.Unlock()
|
||||
return E.New("only one DNS server is allowed for an endpoint")
|
||||
}
|
||||
openConnectEndpoint.dnsTransport = t
|
||||
t.endpoint = openConnectEndpoint
|
||||
t.dialer = openConnectEndpoint
|
||||
state := openConnectEndpoint.state.Load()
|
||||
if state.started && state.tunnelConfigured && openConnectEndpoint.client.Ready() {
|
||||
t.updateConfiguration(state.configuration)
|
||||
}
|
||||
openConnectEndpoint.dnsTransportAccess.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *DNSTransport) updateConfiguration(configuration openconnecttransport.Configuration) {
|
||||
resolverByAddress := make(map[netip.Addr]adapter.DNSTransport)
|
||||
resolverFor := func(address netip.Addr) adapter.DNSTransport {
|
||||
if !address.IsValid() {
|
||||
return nil
|
||||
}
|
||||
resolver, loaded := resolverByAddress[address]
|
||||
if loaded {
|
||||
return resolver
|
||||
}
|
||||
resolver = transport.NewUDPRaw(
|
||||
t.logger,
|
||||
dns.NewTransportAdapter(C.DNSTypeUDP, t.Tag()+"/"+address.String(), nil),
|
||||
t.dialer,
|
||||
M.SocksaddrFrom(address, 53),
|
||||
)
|
||||
resolverByAddress[address] = resolver
|
||||
return resolver
|
||||
}
|
||||
resolversFor := func(addresses []netip.Addr) []adapter.DNSTransport {
|
||||
resolvers := make([]adapter.DNSTransport, 0, len(addresses))
|
||||
resolverSet := make(map[adapter.DNSTransport]bool)
|
||||
for _, address := range addresses {
|
||||
resolver := resolverFor(address)
|
||||
if resolver != nil && !resolverSet[resolver] {
|
||||
resolverSet[resolver] = true
|
||||
resolvers = append(resolvers, resolver)
|
||||
}
|
||||
}
|
||||
return resolvers
|
||||
}
|
||||
defaultResolvers := resolversFor(configuration.DNS)
|
||||
routes := make([]openConnectDNSRoute, 0, len(configuration.SplitDNS)+len(configuration.SearchDomains)+len(configuration.SplitDNSRules))
|
||||
routeIndex := make(map[string]int)
|
||||
for _, rule := range configuration.SplitDNSRules {
|
||||
resolvers := resolversFor(rule.Servers)
|
||||
for _, domain := range rule.Domains {
|
||||
canonicalDomain := canonicalOpenConnectDomain(domain)
|
||||
if canonicalDomain != "" {
|
||||
fqdn := mDNS.Fqdn(canonicalDomain)
|
||||
index, loaded := routeIndex[fqdn]
|
||||
if loaded {
|
||||
resolverSet := make(map[adapter.DNSTransport]bool)
|
||||
for _, resolver := range routes[index].resolvers {
|
||||
resolverSet[resolver] = true
|
||||
}
|
||||
for _, resolver := range resolvers {
|
||||
if !resolverSet[resolver] {
|
||||
routes[index].resolvers = append(routes[index].resolvers, resolver)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
routeIndex[fqdn] = len(routes)
|
||||
routes = append(routes, openConnectDNSRoute{domain: fqdn, resolvers: resolvers})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, domain := range append(append([]string(nil), configuration.SplitDNS...), configuration.SearchDomains...) {
|
||||
canonicalDomain := canonicalOpenConnectDomain(domain)
|
||||
if canonicalDomain != "" {
|
||||
fqdn := mDNS.Fqdn(canonicalDomain)
|
||||
_, loaded := routeIndex[fqdn]
|
||||
if !loaded {
|
||||
routeIndex[fqdn] = len(routes)
|
||||
routes = append(routes, openConnectDNSRoute{domain: fqdn, resolvers: defaultResolvers})
|
||||
}
|
||||
}
|
||||
}
|
||||
searchDomains := make([]string, 0, len(configuration.SearchDomains))
|
||||
searchDomainSet := make(map[string]bool)
|
||||
for _, domain := range configuration.SearchDomains {
|
||||
canonicalDomain := canonicalOpenConnectDomain(domain)
|
||||
if canonicalDomain != "" {
|
||||
fqdn := mDNS.Fqdn(canonicalDomain)
|
||||
if !searchDomainSet[fqdn] {
|
||||
searchDomainSet[fqdn] = true
|
||||
searchDomains = append(searchDomains, fqdn)
|
||||
}
|
||||
}
|
||||
}
|
||||
if !t.acceptDefaultResolvers || !configuration.TunnelAllDNS && (len(configuration.SplitDNS) > 0 || len(configuration.SplitDNSRules) > 0) {
|
||||
defaultResolvers = nil
|
||||
}
|
||||
|
||||
t.access.Lock()
|
||||
if t.closed {
|
||||
t.access.Unlock()
|
||||
for _, resolver := range resolverByAddress {
|
||||
_ = resolver.Close()
|
||||
}
|
||||
return
|
||||
}
|
||||
oldResolvers := t.collectResolversLocked()
|
||||
t.routes = routes
|
||||
t.searchDomains = searchDomains
|
||||
t.defaultResolvers = defaultResolvers
|
||||
activeResolvers := t.collectResolversLocked()
|
||||
t.access.Unlock()
|
||||
|
||||
for _, resolver := range oldResolvers {
|
||||
_ = resolver.Close()
|
||||
}
|
||||
activeResolverSet := make(map[adapter.DNSTransport]bool, len(activeResolvers))
|
||||
for _, resolver := range activeResolvers {
|
||||
activeResolverSet[resolver] = true
|
||||
}
|
||||
for _, resolver := range resolverByAddress {
|
||||
if !activeResolverSet[resolver] {
|
||||
_ = resolver.Close()
|
||||
}
|
||||
}
|
||||
if len(resolverByAddress) > 0 {
|
||||
t.logger.Info("updated ", len(routes), " DNS routes and ", len(resolverByAddress), " resolvers")
|
||||
} else {
|
||||
t.logger.Info("cleared DNS configuration")
|
||||
}
|
||||
}
|
||||
|
||||
func (t *DNSTransport) Reset() {
|
||||
t.access.RLock()
|
||||
resolvers := t.collectResolversLocked()
|
||||
t.access.RUnlock()
|
||||
for _, resolver := range resolvers {
|
||||
resolver.Reset()
|
||||
}
|
||||
}
|
||||
|
||||
func (t *DNSTransport) Close() error {
|
||||
if t.endpoint != nil {
|
||||
t.endpoint.dnsTransportAccess.Lock()
|
||||
if t.endpoint.dnsTransport == t {
|
||||
t.endpoint.dnsTransport = nil
|
||||
}
|
||||
t.endpoint.dnsTransportAccess.Unlock()
|
||||
}
|
||||
t.access.Lock()
|
||||
resolvers := t.collectResolversLocked()
|
||||
t.closed = true
|
||||
t.routes = nil
|
||||
t.searchDomains = nil
|
||||
t.defaultResolvers = nil
|
||||
t.access.Unlock()
|
||||
var closeErr error
|
||||
for _, resolver := range resolvers {
|
||||
closeErr = E.Errors(closeErr, resolver.Close())
|
||||
}
|
||||
return closeErr
|
||||
}
|
||||
|
||||
func (t *DNSTransport) PreferredDomain(domain string) bool {
|
||||
canonicalDomain := mDNS.Fqdn(canonicalOpenConnectDomain(domain))
|
||||
t.access.RLock()
|
||||
routes := t.routes
|
||||
t.access.RUnlock()
|
||||
for _, route := range routes {
|
||||
if mDNS.IsSubDomain(route.domain, canonicalDomain) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *DNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
|
||||
done := make(chan struct{})
|
||||
var response *mDNS.Msg
|
||||
var err error
|
||||
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
|
||||
response = callbackResponse
|
||||
err = callbackErr
|
||||
close(done)
|
||||
})
|
||||
<-done
|
||||
return response, err
|
||||
}
|
||||
|
||||
func (t *DNSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||
if len(message.Question) != 1 {
|
||||
callback(nil, os.ErrInvalid)
|
||||
return
|
||||
}
|
||||
t.access.RLock()
|
||||
searchDomains := append([]string(nil), t.searchDomains...)
|
||||
t.access.RUnlock()
|
||||
if t.acceptSearchDomain && len(searchDomains) > 0 && mDNS.CountLabel(message.Question[0].Name) == 1 {
|
||||
t.exchangeWithSearchDomains(ctx, message, searchDomains, callback)
|
||||
return
|
||||
}
|
||||
t.exchangeOnce(ctx, message, callback)
|
||||
}
|
||||
|
||||
func (t *DNSTransport) exchangeWithSearchDomains(ctx context.Context, message *mDNS.Msg, searchDomains []string, callback func(response *mDNS.Msg, err error)) {
|
||||
originalQuestion := message.Question[0]
|
||||
singleLabel := strings.TrimSuffix(originalQuestion.Name, ".")
|
||||
exchangers := make([]transport.AsyncExchanger, 0, len(searchDomains)+1)
|
||||
for _, searchDomain := range searchDomains {
|
||||
expandedName := singleLabel + "." + searchDomain
|
||||
exchangers = append(exchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
|
||||
question := originalQuestion
|
||||
question.Name = expandedName
|
||||
rewritten := *message
|
||||
rewritten.Question = []mDNS.Question{question}
|
||||
t.exchangeOnce(exchangeCtx, &rewritten, func(response *mDNS.Msg, err error) {
|
||||
if err == nil {
|
||||
restoreOpenConnectDNSQuestion(response, expandedName, originalQuestion)
|
||||
}
|
||||
exchangeCallback(response, err)
|
||||
})
|
||||
})
|
||||
}
|
||||
exchangers = append(exchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
|
||||
t.exchangeOnce(exchangeCtx, message, exchangeCallback)
|
||||
})
|
||||
transport.ExchangeSequential(ctx, exchangers, func(response *mDNS.Msg, err error) bool {
|
||||
return err == nil && response.Rcode != mDNS.RcodeNameError
|
||||
}, callback)
|
||||
}
|
||||
|
||||
func restoreOpenConnectDNSQuestion(response *mDNS.Msg, expandedName string, originalQuestion mDNS.Question) {
|
||||
response.Question = []mDNS.Question{originalQuestion}
|
||||
for _, record := range response.Answer {
|
||||
if strings.EqualFold(record.Header().Name, expandedName) {
|
||||
record.Header().Name = originalQuestion.Name
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
|
||||
question := message.Question[0]
|
||||
t.access.RLock()
|
||||
routes := t.routes
|
||||
defaultResolvers := t.defaultResolvers
|
||||
t.access.RUnlock()
|
||||
var matchedResolvers []adapter.DNSTransport
|
||||
matchedDomainLength := -1
|
||||
for _, route := range routes {
|
||||
if len(route.domain) > matchedDomainLength && mDNS.IsSubDomain(route.domain, question.Name) {
|
||||
matchedDomainLength = len(route.domain)
|
||||
matchedResolvers = route.resolvers
|
||||
}
|
||||
}
|
||||
if matchedDomainLength != -1 {
|
||||
if len(matchedResolvers) == 0 {
|
||||
callback(nil, dns.RcodeNameError)
|
||||
return
|
||||
}
|
||||
transport.ExchangeSequential(ctx, openConnectDNSExchangers(matchedResolvers, message), nil, callback)
|
||||
return
|
||||
}
|
||||
if len(defaultResolvers) == 0 {
|
||||
callback(nil, dns.RcodeNameError)
|
||||
return
|
||||
}
|
||||
transport.ExchangeSequential(ctx, openConnectDNSExchangers(defaultResolvers, message), nil, callback)
|
||||
}
|
||||
|
||||
func openConnectDNSExchangers(resolvers []adapter.DNSTransport, message *mDNS.Msg) []transport.AsyncExchanger {
|
||||
return common.Map(resolvers, func(resolver adapter.DNSTransport) transport.AsyncExchanger {
|
||||
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||
resolver.ExchangeAsync(ctx, message, callback)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func (t *DNSTransport) collectResolversLocked() []adapter.DNSTransport {
|
||||
resolverSet := make(map[adapter.DNSTransport]bool)
|
||||
for _, route := range t.routes {
|
||||
for _, resolver := range route.resolvers {
|
||||
resolverSet[resolver] = true
|
||||
}
|
||||
}
|
||||
for _, resolver := range t.defaultResolvers {
|
||||
resolverSet[resolver] = true
|
||||
}
|
||||
resolvers := make([]adapter.DNSTransport, 0, len(resolverSet))
|
||||
for resolver := range resolverSet {
|
||||
resolvers = append(resolvers, resolver)
|
||||
}
|
||||
return resolvers
|
||||
}
|
||||
@@ -133,6 +133,55 @@ func configurationFromClientEvent(event openconnect.TunnelConfigurationEvent) op
|
||||
Metric: route.Metric,
|
||||
}
|
||||
})
|
||||
if configuration.RemoteAddress.IsValid() {
|
||||
remoteAddress := configuration.RemoteAddress.Unmap()
|
||||
if remoteAddress.Is6() {
|
||||
remoteAddress = remoteAddress.WithZone("")
|
||||
}
|
||||
remoteAddressExcluded := false
|
||||
for _, route := range excludedRoutes {
|
||||
if route.Prefix.Contains(remoteAddress) {
|
||||
remoteAddressExcluded = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !remoteAddressExcluded {
|
||||
excludedRoutes = append(excludedRoutes, openconnecttransport.Route{
|
||||
Prefix: netip.PrefixFrom(remoteAddress, remoteAddress.BitLen()),
|
||||
})
|
||||
}
|
||||
}
|
||||
dnsAddresses := append([]netip.Addr(nil), configuration.DNS...)
|
||||
for _, rule := range configuration.SplitDNSRules {
|
||||
dnsAddresses = append(dnsAddresses, rule.Servers...)
|
||||
}
|
||||
for _, dnsAddress := range dnsAddresses {
|
||||
if !dnsAddress.IsValid() {
|
||||
continue
|
||||
}
|
||||
dnsAddressExcluded := false
|
||||
for _, route := range excludedRoutes {
|
||||
if route.Prefix.Contains(dnsAddress) {
|
||||
dnsAddressExcluded = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if dnsAddressExcluded {
|
||||
continue
|
||||
}
|
||||
dnsAddressIncluded := false
|
||||
for _, route := range routes {
|
||||
if route.Prefix.Contains(dnsAddress) {
|
||||
dnsAddressIncluded = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !dnsAddressIncluded {
|
||||
routes = append(routes, openconnecttransport.Route{
|
||||
Prefix: netip.PrefixFrom(dnsAddress, dnsAddress.BitLen()),
|
||||
})
|
||||
}
|
||||
}
|
||||
splitDNSRules := common.Map(configuration.SplitDNSRules, func(rule openconnect.TunnelSplitDNSRule) openconnecttransport.SplitDNSRule {
|
||||
return openconnecttransport.SplitDNSRule{
|
||||
Domains: rule.Domains,
|
||||
@@ -168,3 +217,46 @@ func buildIPSet(routes []openconnecttransport.Route, excludedRoutes []openconnec
|
||||
}
|
||||
return builder.IPSet()
|
||||
}
|
||||
|
||||
func buildPreferredDomains(configuration openconnecttransport.Configuration) map[string]bool {
|
||||
preferredDomains := make(map[string]bool)
|
||||
for _, domain := range configuration.SearchDomains {
|
||||
canonicalDomain := canonicalOpenConnectDomain(domain)
|
||||
if canonicalDomain != "" {
|
||||
preferredDomains[canonicalDomain] = true
|
||||
}
|
||||
}
|
||||
for _, domain := range configuration.SplitDNS {
|
||||
canonicalDomain := canonicalOpenConnectDomain(domain)
|
||||
if canonicalDomain != "" {
|
||||
preferredDomains[canonicalDomain] = true
|
||||
}
|
||||
}
|
||||
for _, rule := range configuration.SplitDNSRules {
|
||||
for _, domain := range rule.Domains {
|
||||
canonicalDomain := canonicalOpenConnectDomain(domain)
|
||||
if canonicalDomain != "" {
|
||||
preferredDomains[canonicalDomain] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
return preferredDomains
|
||||
}
|
||||
|
||||
func canonicalOpenConnectDomain(domain string) string {
|
||||
return strings.ToLower(strings.Trim(strings.TrimSpace(domain), "."))
|
||||
}
|
||||
|
||||
func openConnectDomainMatchesAny(domain string, suffixes map[string]bool) bool {
|
||||
for domain != "" {
|
||||
if suffixes[domain] {
|
||||
return true
|
||||
}
|
||||
dotIndex := strings.IndexByte(domain, '.')
|
||||
if dotIndex == -1 {
|
||||
break
|
||||
}
|
||||
domain = domain[dotIndex+1:]
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user