374 lines
12 KiB
Go
374 lines
12 KiB
Go
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
|
|
}
|