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
|
||||
}
|
||||
|
||||
+279
-68
@@ -4,6 +4,8 @@ import (
|
||||
"context"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
@@ -49,6 +51,7 @@ type ClientEndpoint struct {
|
||||
device ovpntransport.Device
|
||||
stateAccess sync.Mutex
|
||||
state atomic.Pointer[clientState]
|
||||
dnsTransport *DNSTransport
|
||||
deviceStarted bool
|
||||
readLoopDone chan struct{}
|
||||
statusAccess sync.Mutex
|
||||
@@ -63,6 +66,8 @@ type clientState struct {
|
||||
localAddresses []netip.Prefix
|
||||
routeSet *netipx.IPSet
|
||||
blockIPv6 bool
|
||||
configuration ovpntransport.Configuration
|
||||
preferredDomains []string
|
||||
tunnelInfo adapter.OpenVPNTunnelInfo
|
||||
}
|
||||
|
||||
@@ -149,8 +154,14 @@ func NewClientEndpoint(ctx context.Context, router adapter.Router, logger log.Co
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpointOptions) (ovpn.ClientOptions, error) {
|
||||
if options.TLS == nil {
|
||||
return ovpn.ClientOptions{}, E.New("missing `tls` options")
|
||||
mode := options.Mode
|
||||
if mode == "" {
|
||||
mode = ovpn.ModeTLS
|
||||
}
|
||||
switch mode {
|
||||
case ovpn.ModeTLS, ovpn.ModeStaticKey:
|
||||
default:
|
||||
return ovpn.ClientOptions{}, E.New("unsupported mode: ", mode, " (expected \"tls\" or \"static_key\")")
|
||||
}
|
||||
if options.Server != "" && len(options.Servers) > 0 {
|
||||
return ovpn.ClientOptions{}, E.New("`server` is conflict with `servers`")
|
||||
@@ -158,6 +169,26 @@ func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpoint
|
||||
if options.Server == "" && len(options.Servers) == 0 {
|
||||
return ovpn.ClientOptions{}, E.New("missing `server` or `servers`")
|
||||
}
|
||||
protocol, remotes := buildClientRemoteOptions(options)
|
||||
tunnelOptions, err := buildClientTunnelOptions(options, mode == ovpn.ModeStaticKey)
|
||||
if err != nil {
|
||||
return ovpn.ClientOptions{}, err
|
||||
}
|
||||
if mode == ovpn.ModeStaticKey {
|
||||
return c.buildStaticKeyClientOptions(options, protocol, remotes, tunnelOptions)
|
||||
}
|
||||
if options.TLS == nil {
|
||||
return ovpn.ClientOptions{}, E.New("missing `tls` options")
|
||||
}
|
||||
if len(options.StaticKey) > 0 || options.StaticKeyPath != "" {
|
||||
return ovpn.ClientOptions{}, E.New("`static_key` and `static_key_path` are only supported in `static_key` mode")
|
||||
}
|
||||
if options.KeyDirection != "" {
|
||||
return ovpn.ClientOptions{}, E.New("`key_direction` is only supported in `static_key` mode; use `tls.control_wrap.direction` for `tls_auth`")
|
||||
}
|
||||
if options.Cipher != "" {
|
||||
return ovpn.ClientOptions{}, E.New("`cipher` is only supported in `static_key` mode; use `data_ciphers` or `data_ciphers_fallback` in TLS mode")
|
||||
}
|
||||
certificateAuthority, err := materialSource("tls.certificate", options.TLS.Certificate, options.TLS.CertificatePath)
|
||||
if err != nil {
|
||||
return ovpn.ClientOptions{}, err
|
||||
@@ -198,34 +229,9 @@ func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpoint
|
||||
}
|
||||
controlCryptV2 = controlKey
|
||||
case "":
|
||||
return ovpn.ClientOptions{}, E.New("missing OpenVPN control wrap type")
|
||||
return ovpn.ClientOptions{}, E.New("missing control wrap type")
|
||||
default:
|
||||
return ovpn.ClientOptions{}, E.New("unknown OpenVPN control wrap type: ", controlWrap.Type)
|
||||
}
|
||||
}
|
||||
protocol := options.Network
|
||||
if protocol == "" {
|
||||
protocol = N.NetworkUDP
|
||||
}
|
||||
var remotes []ovpn.Remote
|
||||
if options.Server != "" {
|
||||
remotes = append(remotes, ovpn.Remote{
|
||||
Host: options.Server,
|
||||
Port: options.ServerPort,
|
||||
Protocol: protocol,
|
||||
})
|
||||
} else {
|
||||
remotes = make([]ovpn.Remote, 0, len(options.Servers))
|
||||
for _, remoteOptions := range options.Servers {
|
||||
remoteProtocol := remoteOptions.Network
|
||||
if remoteProtocol == "" {
|
||||
remoteProtocol = protocol
|
||||
}
|
||||
remotes = append(remotes, ovpn.Remote{
|
||||
Host: remoteOptions.Server,
|
||||
Port: remoteOptions.ServerPort,
|
||||
Protocol: remoteProtocol,
|
||||
})
|
||||
return ovpn.ClientOptions{}, E.New("unknown control wrap type: ", controlWrap.Type)
|
||||
}
|
||||
}
|
||||
pullFilters := common.Map(options.PullFilters, func(filterOptions option.OpenVPNPullFilterOptions) ovpn.PullFilter {
|
||||
@@ -234,9 +240,6 @@ func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpoint
|
||||
Text: filterOptions.Text,
|
||||
}
|
||||
})
|
||||
tunnelRoutes := common.Map(options.Routes, func(route netip.Prefix) ovpn.TunnelRoute {
|
||||
return ovpn.TunnelRoute{Prefix: route}
|
||||
})
|
||||
remoteCertificateTLS := options.TLS.RemoteCertificateTLS
|
||||
switch remoteCertificateTLS {
|
||||
case "", "server", "client", "none":
|
||||
@@ -264,6 +267,7 @@ func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpoint
|
||||
RemoteCertificateKU: options.TLS.RemoteCertificateKU,
|
||||
RemoteCertificateEKU: options.TLS.RemoteCertificateEKU,
|
||||
RemoteCertificateTLS: remoteCertificateTLS,
|
||||
NSCertificateType: options.TLS.NSCertificateType,
|
||||
VersionMin: options.TLS.VersionMin,
|
||||
VersionMax: options.TLS.VersionMax,
|
||||
CertificateProfile: options.TLS.CertificateProfile,
|
||||
@@ -278,7 +282,7 @@ func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpoint
|
||||
}
|
||||
return ovpn.ClientOptions{
|
||||
Context: c.loopContext,
|
||||
Mode: ovpn.ModeTLS,
|
||||
Mode: mode,
|
||||
Transport: ovpn.ClientTransportOptions{
|
||||
Remotes: remotes,
|
||||
RemoteRandom: options.RemoteRandom,
|
||||
@@ -286,19 +290,8 @@ func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpoint
|
||||
ExplicitExitNotify: options.ExplicitExitNotify,
|
||||
DialContextWithAddressIndex: c.transportDialContextWithAddressIndex,
|
||||
},
|
||||
DataChannel: ovpn.ClientDataChannelOptions{
|
||||
MTU: options.MTU,
|
||||
MSSFix: options.MSSFix,
|
||||
Fragment: options.Fragment,
|
||||
Ciphers: options.DataCiphers,
|
||||
FallbackCipher: options.DataCiphersFallback,
|
||||
Auth: options.Auth,
|
||||
Compression: options.Compression,
|
||||
CompressionLZO: options.CompressionLZO,
|
||||
AllowCompression: options.AllowCompression,
|
||||
PacketHeadroom: ovpntransport.PacketHeadroom,
|
||||
},
|
||||
TLS: clientTLSOptions,
|
||||
DataChannel: buildClientDataChannelOptions(options),
|
||||
TLS: clientTLSOptions,
|
||||
Authentication: ovpn.ClientAuthenticationOptions{
|
||||
Username: options.Username,
|
||||
Password: options.Password,
|
||||
@@ -311,25 +304,176 @@ func (c *ClientEndpoint) buildClientOptions(options option.OpenVPNClientEndpoint
|
||||
Filters: pullFilters,
|
||||
RouteNoPull: options.RouteNoPull,
|
||||
},
|
||||
Tunnel: ovpn.ClientTunnelOptions{
|
||||
DevType: "tun",
|
||||
RedirectGateway: options.RedirectGateway,
|
||||
RedirectGatewayFlags: options.RedirectGatewayFlags,
|
||||
RouteMetric: options.RouteMetric,
|
||||
RouteGateway: options.RouteGateway.Build(netip.Addr{}),
|
||||
Routes: tunnelRoutes,
|
||||
},
|
||||
Timing: ovpn.ClientTimingOptions{
|
||||
RenegotiationInterval: time.Duration(options.RenegotiateInterval),
|
||||
PingInterval: time.Duration(options.PingInterval),
|
||||
PingRestart: time.Duration(options.PingRestart),
|
||||
},
|
||||
Tunnel: tunnelOptions,
|
||||
Timing: buildClientTimingOptions(options),
|
||||
KeyDirection: keyDirection,
|
||||
OnTunnelConfiguration: c.handleTunnelConfiguration,
|
||||
Logger: c.logger,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) buildStaticKeyClientOptions(options option.OpenVPNClientEndpointOptions, protocol string, remotes []ovpn.Remote, tunnelOptions ovpn.ClientTunnelOptions) (ovpn.ClientOptions, error) {
|
||||
if options.TLS != nil {
|
||||
return ovpn.ClientOptions{}, E.New("`tls` options are not supported in `static_key` mode")
|
||||
}
|
||||
if options.Username != "" || options.Password != "" || (options.AuthRetry != "" && options.AuthRetry != "none") || options.StaticChallenge != "" || options.StaticChallengeEcho {
|
||||
return ovpn.ClientOptions{}, E.New("username/password authentication is not supported in `static_key` mode")
|
||||
}
|
||||
if options.RouteNoPull || len(options.PullFilters) > 0 {
|
||||
return ovpn.ClientOptions{}, E.New("pull options are not supported in `static_key` mode")
|
||||
}
|
||||
if options.RenegotiateInterval != 0 || options.RenegotiateDisabled || options.RenegotiateBytes != 0 || options.RenegotiatePackets != 0 || options.TLSTimeout != 0 || options.HandshakeWindow != 0 {
|
||||
return ovpn.ClientOptions{}, E.New("TLS timing and renegotiation options are not supported in `static_key` mode")
|
||||
}
|
||||
if len(options.DataCiphers) > 0 || options.DataCiphersFallback != "" {
|
||||
return ovpn.ClientOptions{}, E.New("`data_ciphers` and `data_ciphers_fallback` are not supported in `static_key` mode; use `cipher`")
|
||||
}
|
||||
staticKey, err := requiredMaterialSource("static_key", options.StaticKey, options.StaticKeyPath)
|
||||
if err != nil {
|
||||
return ovpn.ClientOptions{}, err
|
||||
}
|
||||
keyDirection, err := keyDirectionValue(options.KeyDirection)
|
||||
if err != nil {
|
||||
return ovpn.ClientOptions{}, err
|
||||
}
|
||||
return ovpn.ClientOptions{
|
||||
Context: c.loopContext,
|
||||
Mode: ovpn.ModeStaticKey,
|
||||
Transport: ovpn.ClientTransportOptions{
|
||||
Remotes: remotes,
|
||||
RemoteRandom: options.RemoteRandom,
|
||||
Protocol: protocol,
|
||||
ExplicitExitNotify: options.ExplicitExitNotify,
|
||||
DialContextWithAddressIndex: c.transportDialContextWithAddressIndex,
|
||||
},
|
||||
DataChannel: buildClientDataChannelOptions(options),
|
||||
Tunnel: tunnelOptions,
|
||||
Timing: buildClientTimingOptions(options),
|
||||
StaticKey: staticKey,
|
||||
KeyDirection: keyDirection,
|
||||
OnTunnelConfiguration: c.handleTunnelConfiguration,
|
||||
Logger: c.logger,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildClientRemoteOptions(options option.OpenVPNClientEndpointOptions) (string, []ovpn.Remote) {
|
||||
protocol := options.Network
|
||||
if protocol == "" {
|
||||
protocol = N.NetworkUDP
|
||||
}
|
||||
if options.Server != "" {
|
||||
return protocol, []ovpn.Remote{{
|
||||
Host: options.Server,
|
||||
Port: options.ServerPort,
|
||||
Protocol: protocol,
|
||||
}}
|
||||
}
|
||||
remotes := make([]ovpn.Remote, 0, len(options.Servers))
|
||||
for _, remoteOptions := range options.Servers {
|
||||
remoteProtocol := remoteOptions.Network
|
||||
if remoteProtocol == "" {
|
||||
remoteProtocol = protocol
|
||||
}
|
||||
remotes = append(remotes, ovpn.Remote{
|
||||
Host: remoteOptions.Server,
|
||||
Port: remoteOptions.ServerPort,
|
||||
Protocol: remoteProtocol,
|
||||
})
|
||||
}
|
||||
return protocol, remotes
|
||||
}
|
||||
|
||||
func buildClientDataChannelOptions(options option.OpenVPNClientEndpointOptions) ovpn.ClientDataChannelOptions {
|
||||
return ovpn.ClientDataChannelOptions{
|
||||
MTU: options.MTU,
|
||||
MSSFix: options.MSSFix,
|
||||
MSSFixDisabled: options.MSSFixDisabled,
|
||||
MSSFixMode: options.MSSFixMode,
|
||||
Fragment: options.Fragment,
|
||||
Cipher: options.Cipher,
|
||||
Ciphers: options.DataCiphers,
|
||||
FallbackCipher: options.DataCiphersFallback,
|
||||
Auth: options.Auth,
|
||||
Compression: options.Compression,
|
||||
CompressionLZO: options.CompressionLZO,
|
||||
AllowCompression: options.AllowCompression,
|
||||
ReplayWindow: options.ReplayWindow,
|
||||
ReplayWindowTime: time.Duration(options.ReplayWindowTime),
|
||||
PacketHeadroom: ovpntransport.PacketHeadroom,
|
||||
}
|
||||
}
|
||||
|
||||
func buildClientTunnelOptions(options option.OpenVPNClientEndpointOptions, requirePeerAddress bool) (ovpn.ClientTunnelOptions, error) {
|
||||
vpnGateway := netip.Addr(options.PeerAddress)
|
||||
if vpnGateway.IsValid() && !vpnGateway.Is4() {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("`peer_address` must be an IPv4 address")
|
||||
}
|
||||
vpnGatewayIPv6 := netip.Addr(options.PeerAddressIPv6)
|
||||
if vpnGatewayIPv6.IsValid() && !vpnGatewayIPv6.Is6() {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("`peer_address_ipv6` must be an IPv6 address")
|
||||
}
|
||||
var hasIPv4 bool
|
||||
var hasIPv6 bool
|
||||
for addressIndex, address := range options.Address {
|
||||
if !address.IsValid() {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("`address[", addressIndex, "]` is invalid")
|
||||
}
|
||||
if address.Addr().Is4() {
|
||||
hasIPv4 = true
|
||||
} else {
|
||||
hasIPv6 = true
|
||||
}
|
||||
}
|
||||
if requirePeerAddress {
|
||||
if len(options.Address) == 0 {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("missing `address` in `static_key` mode")
|
||||
}
|
||||
if hasIPv4 && !vpnGateway.IsValid() {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("missing `peer_address` for the IPv4 tunnel address in `static_key` mode")
|
||||
}
|
||||
if hasIPv6 && !vpnGatewayIPv6.IsValid() {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("missing `peer_address_ipv6` for the IPv6 tunnel address in `static_key` mode")
|
||||
}
|
||||
if vpnGateway.IsValid() && !hasIPv4 {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("`peer_address` requires an IPv4 tunnel `address` in `static_key` mode")
|
||||
}
|
||||
if vpnGatewayIPv6.IsValid() && !hasIPv6 {
|
||||
return ovpn.ClientTunnelOptions{}, E.New("`peer_address_ipv6` requires an IPv6 tunnel `address` in `static_key` mode")
|
||||
}
|
||||
}
|
||||
tunnelRoutes := common.Map(options.Routes, func(route netip.Prefix) ovpn.TunnelRoute {
|
||||
return ovpn.TunnelRoute{Prefix: route}
|
||||
})
|
||||
return ovpn.ClientTunnelOptions{
|
||||
DevType: "tun",
|
||||
Topology: options.Topology,
|
||||
RedirectGateway: options.RedirectGateway,
|
||||
RedirectGatewayFlags: options.RedirectGatewayFlags,
|
||||
RedirectPrivate: options.RedirectPrivate,
|
||||
BlockIPv6: options.BlockIPv6,
|
||||
RouteMetric: options.RouteMetric,
|
||||
RouteGateway: options.RouteGateway.Build(netip.Addr{}),
|
||||
Routes: tunnelRoutes,
|
||||
LocalAddress: options.Address,
|
||||
VPNGateway: vpnGateway,
|
||||
VPNGatewayIPv6: vpnGatewayIPv6,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildClientTimingOptions(options option.OpenVPNClientEndpointOptions) ovpn.ClientTimingOptions {
|
||||
return ovpn.ClientTimingOptions{
|
||||
RenegotiationInterval: time.Duration(options.RenegotiateInterval),
|
||||
RenegotiationDisabled: options.RenegotiateDisabled,
|
||||
RenegotiationBytes: options.RenegotiateBytes,
|
||||
RenegotiationPackets: options.RenegotiatePackets,
|
||||
PingInterval: time.Duration(options.PingInterval),
|
||||
PingRestart: time.Duration(options.PingRestart),
|
||||
PingRestartDisabled: options.PingRestartDisabled,
|
||||
TLSTimeout: time.Duration(options.TLSTimeout),
|
||||
HandWindow: time.Duration(options.HandshakeWindow),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) transportDialContextWithAddressIndex(ctx context.Context, network string, address string, addressIndex int) (net.Conn, error) {
|
||||
destination := M.ParseSocksaddr(address)
|
||||
if destination.IsDomain() {
|
||||
@@ -361,33 +505,51 @@ func (c *ClientEndpoint) transportDialContextWithAddressIndex(ctx context.Contex
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) handleTunnelConfiguration(event ovpn.TunnelConfigurationEvent) error {
|
||||
configuration := configurationFromClientEvent(event, c.logger)
|
||||
defer c.notifyStatusUpdated()
|
||||
c.stateAccess.Lock()
|
||||
defer c.stateAccess.Unlock()
|
||||
configuration := configurationFromClientEvent(event, c.logger)
|
||||
c.updateState(func(state *clientState) {
|
||||
state.tunnelConfigured = false
|
||||
})
|
||||
err := c.device.UpdateConfiguration(configuration)
|
||||
deviceConfiguration := ovpntransport.Configuration{
|
||||
MTU: configuration.MTU,
|
||||
Address: configuration.Address,
|
||||
BlockIPv6: configuration.BlockIPv6,
|
||||
}
|
||||
err := c.device.UpdateConfiguration(deviceConfiguration)
|
||||
if err != nil {
|
||||
c.stateAccess.Unlock()
|
||||
return E.Cause(err, "update device configuration")
|
||||
}
|
||||
if !c.deviceStarted {
|
||||
err = c.device.Start()
|
||||
if err != nil {
|
||||
c.stateAccess.Unlock()
|
||||
return E.Cause(err, "start device")
|
||||
}
|
||||
c.deviceStarted = true
|
||||
}
|
||||
routeSet, err := buildIPSet(configuration.Routes)
|
||||
routeSet, err := buildIPSet(configuration.Routes, configuration.ExcludedRoutes)
|
||||
if err != nil {
|
||||
c.stateAccess.Unlock()
|
||||
return E.Cause(err, "build route set")
|
||||
}
|
||||
preferredDomains := slices.Clone(configuration.DNSRoutes)
|
||||
preferredDomains = append(preferredDomains, configuration.SearchDomains...)
|
||||
if len(configuration.DNSServers) > 0 {
|
||||
servers := slices.Clone(configuration.DNSServers)
|
||||
slices.SortFunc(servers, func(left ovpntransport.DNSServer, right ovpntransport.DNSServer) int {
|
||||
return left.Priority - right.Priority
|
||||
})
|
||||
preferredDomains = append(preferredDomains, servers[0].ResolveDomains...)
|
||||
}
|
||||
c.updateState(func(state *clientState) {
|
||||
state.tunnelConfigured = true
|
||||
state.localAddresses = configuration.Address
|
||||
state.routeSet = routeSet
|
||||
state.blockIPv6 = configuration.BlockIPv6
|
||||
state.configuration = configuration
|
||||
state.preferredDomains = preferredDomains
|
||||
state.tunnelInfo.Cipher = event.Configuration.SelectedCipher
|
||||
state.tunnelInfo.IPv4 = event.Configuration.LocalIPv4
|
||||
state.tunnelInfo.IPv6 = event.Configuration.LocalIPv6
|
||||
@@ -397,6 +559,11 @@ func (c *ClientEndpoint) handleTunnelConfiguration(event ovpn.TunnelConfiguratio
|
||||
state.tunnelInfo.ConnectedSince = time.Now()
|
||||
}
|
||||
})
|
||||
dnsTransport := c.dnsTransport
|
||||
c.stateAccess.Unlock()
|
||||
if dnsTransport != nil {
|
||||
dnsTransport.onReconfiguration(configuration)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -406,6 +573,28 @@ func (c *ClientEndpoint) updateState(update func(state *clientState)) {
|
||||
c.state.Store(&newState)
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) installDNSTransport(dnsTransport *DNSTransport) error {
|
||||
c.stateAccess.Lock()
|
||||
defer c.stateAccess.Unlock()
|
||||
if c.dnsTransport != nil && c.dnsTransport != dnsTransport && c.dnsTransport.Tag() != dnsTransport.Tag() {
|
||||
return E.New("only one DNS server is allowed for an endpoint")
|
||||
}
|
||||
err := dnsTransport.updateResolvers(c.state.Load().configuration)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.dnsTransport = dnsTransport
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) uninstallDNSTransport(dnsTransport *DNSTransport) {
|
||||
c.stateAccess.Lock()
|
||||
if c.dnsTransport == dnsTransport {
|
||||
c.dnsTransport = nil
|
||||
}
|
||||
c.stateAccess.Unlock()
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) Start(stage adapter.StartStage) error {
|
||||
if stage != adapter.StartStatePostStart {
|
||||
return nil
|
||||
@@ -434,7 +623,7 @@ func (c *ClientEndpoint) readLoop() {
|
||||
if E.IsClosedOrCanceled(err) || c.loopContext.Err() != nil {
|
||||
return
|
||||
}
|
||||
c.logger.Error(E.Cause(err, "OpenVPN client terminated"))
|
||||
c.logger.Error(E.Cause(err, "client terminated"))
|
||||
c.setTerminalError(err)
|
||||
return
|
||||
}
|
||||
@@ -509,7 +698,7 @@ func (c *ClientEndpoint) ready() bool {
|
||||
func (c *ClientEndpoint) WritePackets(packets [][]byte) error {
|
||||
state := c.state.Load()
|
||||
if !state.started || !state.tunnelConfigured {
|
||||
return E.New("OpenVPN client is not ready yet")
|
||||
return E.New("endpoint is not ready yet")
|
||||
}
|
||||
if state.blockIPv6 {
|
||||
outboundPackets := packets[:0]
|
||||
@@ -529,7 +718,7 @@ func (c *ClientEndpoint) WritePackets(packets [][]byte) error {
|
||||
}
|
||||
err := c.client.WriteDataPacketBuffers(packetBuffers)
|
||||
if E.IsMulti(err, ovpn.ErrDataChannelNotReady) {
|
||||
return E.New("OpenVPN client is not ready yet")
|
||||
return E.New("endpoint is not ready yet")
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -577,7 +766,7 @@ func (c *ClientEndpoint) DialContext(ctx context.Context, network string, destin
|
||||
c.logger.InfoContext(ctx, "outbound packet connection to ", destination)
|
||||
}
|
||||
if !c.ready() || !c.client.Ready() {
|
||||
return nil, E.New("OpenVPN client is not ready yet")
|
||||
return nil, E.New("endpoint is not ready yet")
|
||||
}
|
||||
if destination.IsDomain() {
|
||||
destinationAddresses, err := c.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
|
||||
@@ -595,7 +784,7 @@ func (c *ClientEndpoint) DialContext(ctx context.Context, network string, destin
|
||||
func (c *ClientEndpoint) ListenPacketWithDestination(ctx context.Context, destination M.Socksaddr) (net.PacketConn, netip.Addr, error) {
|
||||
c.logger.InfoContext(ctx, "outbound packet connection to ", destination)
|
||||
if !c.ready() || !c.client.Ready() {
|
||||
return nil, netip.Addr{}, E.New("OpenVPN client is not ready yet")
|
||||
return nil, netip.Addr{}, E.New("endpoint is not ready yet")
|
||||
}
|
||||
if destination.IsDomain() {
|
||||
destinationAddresses, err := c.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
|
||||
@@ -626,6 +815,15 @@ func (c *ClientEndpoint) ListenPacket(ctx context.Context, destination M.Socksad
|
||||
}
|
||||
|
||||
func (c *ClientEndpoint) PreferredDomain(metadata *adapter.InboundContext, domain string) bool {
|
||||
state := c.state.Load()
|
||||
if !state.started || !state.tunnelConfigured || !c.client.Ready() {
|
||||
return false
|
||||
}
|
||||
for _, preferredDomain := range state.preferredDomains {
|
||||
if openVPNDomainMatches(preferredDomain, domain) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -636,3 +834,16 @@ func (c *ClientEndpoint) PreferredAddress(metadata *adapter.InboundContext, addr
|
||||
}
|
||||
return state.routeSet.Contains(address)
|
||||
}
|
||||
|
||||
func openVPNDomainMatches(suffix string, domain string) bool {
|
||||
normalizedSuffix := strings.ToLower(strings.TrimSpace(suffix))
|
||||
if normalizedSuffix == "." {
|
||||
return true
|
||||
}
|
||||
normalizedSuffix = strings.TrimSuffix(normalizedSuffix, ".")
|
||||
normalizedDomain := strings.TrimSuffix(strings.ToLower(strings.TrimSpace(domain)), ".")
|
||||
if normalizedSuffix == "" {
|
||||
return false
|
||||
}
|
||||
return normalizedDomain == normalizedSuffix || strings.HasSuffix(normalizedDomain, "."+normalizedSuffix)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,427 @@
|
||||
package openvpn
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/sagernet/sing-box/adapter"
|
||||
boxTLS "github.com/sagernet/sing-box/common/tls"
|
||||
C "github.com/sagernet/sing-box/constant"
|
||||
boxDNS "github.com/sagernet/sing-box/dns"
|
||||
dnsTransport "github.com/sagernet/sing-box/dns/transport"
|
||||
"github.com/sagernet/sing-box/log"
|
||||
"github.com/sagernet/sing-box/option"
|
||||
ovpntransport "github.com/sagernet/sing-box/transport/openvpn"
|
||||
"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"
|
||||
"golang.org/x/net/http2"
|
||||
)
|
||||
|
||||
func RegisterDNSTransport(registry *boxDNS.TransportRegistry) {
|
||||
boxDNS.RegisterTransport[option.OpenVPNDNSServerOptions](registry, C.DNSTypeOpenVPN, NewDNSTransport)
|
||||
}
|
||||
|
||||
type DNSTransport struct {
|
||||
boxDNS.TransportAdapter
|
||||
ctx context.Context
|
||||
logger logger.ContextLogger
|
||||
endpointTag string
|
||||
acceptDefaultResolvers bool
|
||||
acceptSearchDomain bool
|
||||
endpointManager adapter.EndpointManager
|
||||
endpoint *ClientEndpoint
|
||||
dialer N.Dialer
|
||||
updateAccess sync.Mutex
|
||||
access sync.RWMutex
|
||||
closed bool
|
||||
routes map[string][]adapter.DNSTransport
|
||||
searchDomains []string
|
||||
defaultResolvers []adapter.DNSTransport
|
||||
}
|
||||
|
||||
func NewDNSTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.OpenVPNDNSServerOptions) (adapter.DNSTransport, error) {
|
||||
if options.Endpoint == "" {
|
||||
return nil, E.New("missing endpoint tag")
|
||||
}
|
||||
return &DNSTransport{
|
||||
TransportAdapter: boxDNS.NewTransportAdapter(C.DNSTypeOpenVPN, tag, nil),
|
||||
ctx: ctx,
|
||||
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)
|
||||
}
|
||||
endpoint, isOpenVPN := rawEndpoint.(*ClientEndpoint)
|
||||
if !isOpenVPN {
|
||||
return E.New("endpoint is not an OpenVPN client: ", t.endpointTag)
|
||||
}
|
||||
t.endpoint = endpoint
|
||||
t.dialer = endpoint
|
||||
err := endpoint.installDNSTransport(t)
|
||||
if err != nil {
|
||||
t.endpoint = nil
|
||||
t.dialer = nil
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *DNSTransport) onReconfiguration(configuration ovpntransport.Configuration) {
|
||||
err := t.updateResolvers(configuration)
|
||||
if err != nil && !E.IsClosed(err) {
|
||||
t.logger.Error(E.Cause(err, "update DNS resolvers"))
|
||||
}
|
||||
}
|
||||
|
||||
func (t *DNSTransport) updateResolvers(configuration ovpntransport.Configuration) error {
|
||||
t.updateAccess.Lock()
|
||||
defer t.updateAccess.Unlock()
|
||||
t.access.RLock()
|
||||
closed := t.closed
|
||||
t.access.RUnlock()
|
||||
if closed {
|
||||
return net.ErrClosed
|
||||
}
|
||||
routes := make(map[string][]adapter.DNSTransport)
|
||||
searchDomains := normalizeOpenVPNDomains(configuration.SearchDomains)
|
||||
var defaultResolvers []adapter.DNSTransport
|
||||
var newResolvers []adapter.DNSTransport
|
||||
servers := slices.Clone(configuration.DNSServers)
|
||||
slices.SortFunc(servers, func(left ovpntransport.DNSServer, right ovpntransport.DNSServer) int {
|
||||
return left.Priority - right.Priority
|
||||
})
|
||||
var selectedResolvers []adapter.DNSTransport
|
||||
if len(servers) > 0 {
|
||||
server := servers[0]
|
||||
if server.DNSSEC == "yes" {
|
||||
return t.failResolverUpdate(newResolvers, E.New("DNSSEC validation is required but is not supported"))
|
||||
}
|
||||
for _, address := range server.Addresses {
|
||||
resolver, err := t.createResolver(server, address)
|
||||
if err != nil {
|
||||
return t.failResolverUpdate(newResolvers, err)
|
||||
}
|
||||
selectedResolvers = append(selectedResolvers, resolver)
|
||||
newResolvers = append(newResolvers, resolver)
|
||||
}
|
||||
if len(selectedResolvers) == 0 {
|
||||
return t.failResolverUpdate(newResolvers, E.New("DNS server ", server.Priority, " has no addresses"))
|
||||
}
|
||||
if len(server.ResolveDomains) == 0 {
|
||||
defaultResolvers = slices.Clone(selectedResolvers)
|
||||
} else {
|
||||
for _, domain := range server.ResolveDomains {
|
||||
normalizedDomain := normalizeOpenVPNDomain(domain)
|
||||
if normalizedDomain != "" {
|
||||
routes[normalizedDomain] = slices.Clone(selectedResolvers)
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for _, address := range configuration.DNS {
|
||||
resolver := dnsTransport.NewUDPRaw(t.logger, t.TransportAdapter, t.dialer, M.SocksaddrFrom(address, 53))
|
||||
selectedResolvers = append(selectedResolvers, resolver)
|
||||
newResolvers = append(newResolvers, resolver)
|
||||
}
|
||||
if len(configuration.DNSRoutes) > 0 {
|
||||
if len(selectedResolvers) == 0 {
|
||||
return t.failResolverUpdate(newResolvers, E.New("DOMAIN-ROUTE requires traditional pushed DNS servers"))
|
||||
}
|
||||
for _, domain := range configuration.DNSRoutes {
|
||||
normalizedDomain := normalizeOpenVPNDomain(domain)
|
||||
if normalizedDomain != "" {
|
||||
routes[normalizedDomain] = slices.Clone(selectedResolvers)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
defaultResolvers = slices.Clone(selectedResolvers)
|
||||
}
|
||||
}
|
||||
if len(searchDomains) > 0 && len(selectedResolvers) == 0 {
|
||||
return t.failResolverUpdate(newResolvers, E.New("search domains require pushed DNS servers"))
|
||||
}
|
||||
for _, searchDomain := range searchDomains {
|
||||
routes[searchDomain] = slices.Clone(selectedResolvers)
|
||||
}
|
||||
|
||||
t.access.Lock()
|
||||
oldResolvers := t.collectResolversLocked()
|
||||
t.routes = routes
|
||||
t.searchDomains = searchDomains
|
||||
t.defaultResolvers = defaultResolvers
|
||||
t.access.Unlock()
|
||||
closeErr := closeDNSTransports(oldResolvers)
|
||||
t.logger.Info("updated ", len(routes), " DNS routes, ", len(searchDomains), " search domains and ", len(defaultResolvers), " default resolvers")
|
||||
return closeErr
|
||||
}
|
||||
|
||||
func (t *DNSTransport) failResolverUpdate(newResolvers []adapter.DNSTransport, updateErr error) error {
|
||||
newCloseErr := closeDNSTransports(newResolvers)
|
||||
t.access.Lock()
|
||||
oldResolvers := t.collectResolversLocked()
|
||||
t.routes = nil
|
||||
t.searchDomains = nil
|
||||
t.defaultResolvers = nil
|
||||
t.access.Unlock()
|
||||
oldCloseErr := closeDNSTransports(oldResolvers)
|
||||
return E.Errors(updateErr, newCloseErr, oldCloseErr)
|
||||
}
|
||||
|
||||
func (t *DNSTransport) createResolver(server ovpntransport.DNSServer, address netip.AddrPort) (adapter.DNSTransport, error) {
|
||||
transportType := strings.ToLower(server.Transport)
|
||||
if transportType == "" {
|
||||
transportType = "plain"
|
||||
}
|
||||
port := address.Port()
|
||||
switch transportType {
|
||||
case "plain":
|
||||
if port == 0 {
|
||||
port = 53
|
||||
}
|
||||
return dnsTransport.NewUDPRaw(t.logger, t.TransportAdapter, t.dialer, M.SocksaddrFrom(address.Addr(), port)), nil
|
||||
case "dot", "doh":
|
||||
default:
|
||||
return nil, E.New("unsupported DNS transport: ", server.Transport)
|
||||
}
|
||||
serverName := server.SNI
|
||||
if serverName == "" {
|
||||
serverName = address.Addr().String()
|
||||
}
|
||||
if transportType == "dot" {
|
||||
if port == 0 {
|
||||
port = 853
|
||||
}
|
||||
tlsConfig, err := boxTLS.NewClient(t.ctx, t.logger, serverName, option.OutboundTLSOptions{
|
||||
Enabled: true,
|
||||
ServerName: serverName,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return dnsTransport.NewTLSRaw(t.logger, t.TransportAdapter, t.dialer, M.SocksaddrFrom(address.Addr(), port), tlsConfig), nil
|
||||
}
|
||||
if port == 0 {
|
||||
port = 443
|
||||
}
|
||||
tlsConfig, err := boxTLS.NewClient(t.ctx, t.logger, serverName, option.OutboundTLSOptions{
|
||||
Enabled: true,
|
||||
ServerName: serverName,
|
||||
ALPN: []string{http2.NextProtoTLS, "http/1.1"},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
host := serverName
|
||||
if port != 443 {
|
||||
host = net.JoinHostPort(host, strconv.Itoa(int(port)))
|
||||
} else if strings.Contains(host, ":") {
|
||||
host = "[" + host + "]"
|
||||
}
|
||||
destination := &url.URL{Scheme: "https", Host: host, Path: "/dns-query"}
|
||||
return dnsTransport.NewHTTPSRaw(t.TransportAdapter, t.logger, t.dialer, destination, http.Header{}, M.SocksaddrFrom(address.Addr(), port), tlsConfig), nil
|
||||
}
|
||||
|
||||
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.uninstallDNSTransport(t)
|
||||
}
|
||||
t.updateAccess.Lock()
|
||||
t.access.Lock()
|
||||
resolvers := t.collectResolversLocked()
|
||||
t.closed = true
|
||||
t.routes = nil
|
||||
t.searchDomains = nil
|
||||
t.defaultResolvers = nil
|
||||
t.access.Unlock()
|
||||
t.endpoint = nil
|
||||
t.dialer = nil
|
||||
t.updateAccess.Unlock()
|
||||
return closeDNSTransports(resolvers)
|
||||
}
|
||||
|
||||
func (t *DNSTransport) Raw() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (t *DNSTransport) PreferredDomain(domain string) bool {
|
||||
t.access.RLock()
|
||||
defer t.access.RUnlock()
|
||||
for route := range t.routes {
|
||||
if openVPNDomainMatches(route, domain) {
|
||||
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 := slices.Clone(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, t.acceptDefaultResolvers, 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([]dnsTransport.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, false, func(response *mDNS.Msg, err error) {
|
||||
if err == nil {
|
||||
restoreOpenVPNOriginalQuestion(response, expandedName, originalQuestion)
|
||||
}
|
||||
exchangeCallback(response, err)
|
||||
})
|
||||
})
|
||||
}
|
||||
exchangers = append(exchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
|
||||
t.exchangeOnce(exchangeCtx, message, t.acceptDefaultResolvers, exchangeCallback)
|
||||
})
|
||||
dnsTransport.ExchangeSequential(ctx, exchangers, func(response *mDNS.Msg, err error) bool {
|
||||
return err == nil && response.Rcode != mDNS.RcodeNameError
|
||||
}, callback)
|
||||
}
|
||||
|
||||
func (t *DNSTransport) exchangeOnce(ctx context.Context, message *mDNS.Msg, allowDefaultResolvers bool, callback func(response *mDNS.Msg, err error)) {
|
||||
question := message.Question[0]
|
||||
t.access.RLock()
|
||||
var matchedResolvers []adapter.DNSTransport
|
||||
matchedLength := -1
|
||||
for route, resolvers := range t.routes {
|
||||
if openVPNDomainMatches(route, question.Name) && len(route) > matchedLength {
|
||||
matchedLength = len(route)
|
||||
matchedResolvers = resolvers
|
||||
}
|
||||
}
|
||||
defaultResolvers := slices.Clone(t.defaultResolvers)
|
||||
t.access.RUnlock()
|
||||
if len(matchedResolvers) > 0 {
|
||||
dnsTransport.ExchangeSequential(ctx, openVPNResolverExchangers(matchedResolvers, message), nil, callback)
|
||||
return
|
||||
}
|
||||
if allowDefaultResolvers && len(defaultResolvers) > 0 {
|
||||
dnsTransport.ExchangeSequential(ctx, openVPNResolverExchangers(defaultResolvers, message), nil, callback)
|
||||
return
|
||||
}
|
||||
callback(nil, boxDNS.RcodeNameError)
|
||||
}
|
||||
|
||||
func openVPNResolverExchangers(resolvers []adapter.DNSTransport, message *mDNS.Msg) []dnsTransport.AsyncExchanger {
|
||||
return common.Map(resolvers, func(resolver adapter.DNSTransport) dnsTransport.AsyncExchanger {
|
||||
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
|
||||
resolver.ExchangeAsync(ctx, message, callback)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func (t *DNSTransport) collectResolversLocked() []adapter.DNSTransport {
|
||||
var resolvers []adapter.DNSTransport
|
||||
for _, routeResolvers := range t.routes {
|
||||
resolvers = append(resolvers, routeResolvers...)
|
||||
}
|
||||
resolvers = append(resolvers, t.defaultResolvers...)
|
||||
return common.Uniq(resolvers)
|
||||
}
|
||||
|
||||
func closeDNSTransports(resolvers []adapter.DNSTransport) error {
|
||||
var err error
|
||||
for _, resolver := range common.Uniq(resolvers) {
|
||||
err = E.Append(err, resolver.Close(), func(closeErr error) error {
|
||||
return E.Cause(closeErr, "close DNS resolver")
|
||||
})
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func normalizeOpenVPNDomain(domain string) string {
|
||||
normalized := strings.TrimSpace(strings.ToLower(domain))
|
||||
if normalized == "." {
|
||||
return normalized
|
||||
}
|
||||
normalized = strings.TrimSuffix(normalized, ".")
|
||||
if normalized == "" {
|
||||
return ""
|
||||
}
|
||||
return normalized + "."
|
||||
}
|
||||
|
||||
func normalizeOpenVPNDomains(domains []string) []string {
|
||||
normalized := make([]string, 0, len(domains))
|
||||
for _, domain := range domains {
|
||||
normalizedDomain := normalizeOpenVPNDomain(domain)
|
||||
if normalizedDomain != "" && normalizedDomain != "." && !slices.Contains(normalized, normalizedDomain) {
|
||||
normalized = append(normalized, normalizedDomain)
|
||||
}
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func restoreOpenVPNOriginalQuestion(response *mDNS.Msg, expandedName string, originalQuestion mDNS.Question) {
|
||||
response.Question = []mDNS.Question{originalQuestion}
|
||||
for _, resourceRecord := range response.Answer {
|
||||
if strings.EqualFold(resourceRecord.Header().Name, expandedName) {
|
||||
resourceRecord.Header().Name = originalQuestion.Name
|
||||
}
|
||||
}
|
||||
}
|
||||
+115
-25
@@ -17,6 +17,7 @@ import (
|
||||
ovpn "github.com/sagernet/sing-openvpn"
|
||||
"github.com/sagernet/sing-tun"
|
||||
"github.com/sagernet/sing-tun/gtcpip/header"
|
||||
"github.com/sagernet/sing/common"
|
||||
"github.com/sagernet/sing/common/auth"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
@@ -118,7 +119,7 @@ func keyDirectionValue(direction string) (int, error) {
|
||||
case "client":
|
||||
return 1, nil
|
||||
default:
|
||||
return 0, E.New("unsupported OpenVPN key direction: ", direction, " (expected \"server\" or \"client\")")
|
||||
return 0, E.New("unsupported key direction: ", direction, " (expected \"server\" or \"client\")")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -187,21 +188,43 @@ func configurationFromClientEvent(event ovpn.TunnelConfigurationEvent, logger lo
|
||||
hasInet6DefaultRoute = true
|
||||
}
|
||||
}
|
||||
var excludedRoutes []ovpntransport.Route
|
||||
for _, route := range configuration.ExcludedIPv4Routes {
|
||||
excludedRoutes = append(excludedRoutes, ovpntransport.Route{Prefix: route.Prefix, Gateway: route.Gateway, Metric: route.Metric})
|
||||
}
|
||||
for _, route := range configuration.ExcludedIPv6Routes {
|
||||
excludedRoutes = append(excludedRoutes, ovpntransport.Route{Prefix: route.Prefix, Gateway: route.Gateway, Metric: route.Metric})
|
||||
}
|
||||
if configuration.RedirectGateway {
|
||||
if !hasOpenVPNFlag(configuration.RedirectGatewayFlags, "!ipv4") && !hasInet4DefaultRoute {
|
||||
routes = append(routes, ovpntransport.Route{
|
||||
Prefix: inet4DefaultRoute,
|
||||
Gateway: configuration.VPNGateway,
|
||||
Metric: configuration.RouteMetric,
|
||||
})
|
||||
if hasOpenVPNFlag(configuration.RedirectGatewayFlags, "def1") {
|
||||
for _, prefix := range []netip.Prefix{
|
||||
netip.PrefixFrom(netip.IPv4Unspecified(), 1),
|
||||
netip.MustParsePrefix("128.0.0.0/1"),
|
||||
} {
|
||||
if !openVPNRoutesContainPrefix(routes, prefix) {
|
||||
routes = append(routes, ovpntransport.Route{Prefix: prefix, Gateway: configuration.VPNGateway, Metric: configuration.RouteMetric})
|
||||
}
|
||||
}
|
||||
} else {
|
||||
routes = append(routes, ovpntransport.Route{
|
||||
Prefix: inet4DefaultRoute,
|
||||
Gateway: configuration.VPNGateway,
|
||||
Metric: configuration.RouteMetric,
|
||||
})
|
||||
}
|
||||
}
|
||||
if hasOpenVPNFlag(configuration.RedirectGatewayFlags, "ipv6") && !hasInet6DefaultRoute {
|
||||
routes = append(routes, ovpntransport.Route{
|
||||
Prefix: inet6DefaultRoute,
|
||||
Gateway: configuration.VPNGatewayIPv6,
|
||||
Metric: configuration.RouteMetric,
|
||||
})
|
||||
hasInet6DefaultRoute = true
|
||||
for _, prefix := range []netip.Prefix{
|
||||
netip.MustParsePrefix("::/3"),
|
||||
netip.MustParsePrefix("2000::/4"),
|
||||
netip.MustParsePrefix("3000::/4"),
|
||||
netip.MustParsePrefix("fc00::/7"),
|
||||
} {
|
||||
if !openVPNRoutesContainPrefix(routes, prefix) {
|
||||
routes = append(routes, ovpntransport.Route{Prefix: prefix, Gateway: configuration.VPNGatewayIPv6, Metric: configuration.RouteMetric})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if configuration.BlockIPv6 && !hasInet6DefaultRoute {
|
||||
@@ -211,47 +234,114 @@ func configurationFromClientEvent(event ovpn.TunnelConfigurationEvent, logger lo
|
||||
Metric: configuration.RouteMetric,
|
||||
})
|
||||
}
|
||||
var dnsAddresses []netip.Addr
|
||||
if len(configuration.DNSServers) > 0 {
|
||||
servers := slices.Clone(configuration.DNSServers)
|
||||
slices.SortFunc(servers, func(left ovpn.TunnelDNSServer, right ovpn.TunnelDNSServer) int {
|
||||
return left.Priority - right.Priority
|
||||
})
|
||||
for _, address := range servers[0].Addresses {
|
||||
if address.Addr().IsValid() && !slices.Contains(dnsAddresses, address.Addr()) {
|
||||
dnsAddresses = append(dnsAddresses, address.Addr())
|
||||
}
|
||||
}
|
||||
} else {
|
||||
dnsAddresses = slices.Clone(configuration.DNS)
|
||||
}
|
||||
for _, dnsAddress := range dnsAddresses {
|
||||
if openVPNRoutesContainAddress(routes, dnsAddress) || openVPNRoutesContainAddress(excludedRoutes, dnsAddress) {
|
||||
continue
|
||||
}
|
||||
gateway := configuration.VPNGateway
|
||||
if dnsAddress.Is6() {
|
||||
gateway = configuration.VPNGatewayIPv6
|
||||
}
|
||||
routes = append(routes, ovpntransport.Route{
|
||||
Prefix: netip.PrefixFrom(dnsAddress, dnsAddress.BitLen()),
|
||||
Gateway: gateway,
|
||||
Metric: configuration.RouteMetric,
|
||||
})
|
||||
}
|
||||
var ignoredOptions []string
|
||||
var notApplicableOptions []string
|
||||
for _, flag := range configuration.RedirectGatewayFlags {
|
||||
switch strings.ToLower(flag) {
|
||||
case "!ipv4", "ipv6":
|
||||
case "!ipv4", "ipv6", "def1", "local", "autolocal":
|
||||
case "bypass-dhcp", "bypass-dns":
|
||||
notApplicableOptions = append(notApplicableOptions, "redirect-gateway "+flag)
|
||||
default:
|
||||
if flag != "" {
|
||||
ignoredOptions = append(ignoredOptions, "redirect-gateway "+flag)
|
||||
}
|
||||
}
|
||||
}
|
||||
if configuration.RedirectPrivate {
|
||||
ignoredOptions = append(ignoredOptions, "redirect-private")
|
||||
}
|
||||
if configuration.BlockOutsideDNS {
|
||||
ignoredOptions = append(ignoredOptions, "block-outside-dns")
|
||||
}
|
||||
for _, dhcpOption := range configuration.DHCPOptions {
|
||||
fields := strings.Fields(dhcpOption)
|
||||
if len(fields) == 0 || strings.EqualFold(fields[0], "DNS") || strings.EqualFold(fields[0], "DNS6") {
|
||||
if len(fields) == 0 || slices.ContainsFunc([]string{"DNS", "DNS6", "DOMAIN", "ADAPTER_DOMAIN_SUFFIX", "DOMAIN-SEARCH", "DOMAIN-ROUTE"}, func(optionName string) bool {
|
||||
return strings.EqualFold(fields[0], optionName)
|
||||
}) {
|
||||
continue
|
||||
}
|
||||
ignoredOptions = append(ignoredOptions, "dhcp-option "+strings.TrimSpace(dhcpOption))
|
||||
}
|
||||
if len(ignoredOptions) > 0 && logger != nil {
|
||||
logger.Debug("ignored pushed OpenVPN options: ", strings.Join(ignoredOptions, ", "))
|
||||
logger.Debug("ignored pushed options: ", strings.Join(ignoredOptions, ", "))
|
||||
}
|
||||
if len(notApplicableOptions) > 0 && logger != nil {
|
||||
logger.Debug("pushed options are not applicable: ", strings.Join(notApplicableOptions, ", "))
|
||||
}
|
||||
return ovpntransport.Configuration{
|
||||
MTU: mtu,
|
||||
Address: addresses,
|
||||
Routes: routes,
|
||||
DNS: configuration.DNS,
|
||||
Topology: configuration.Topology,
|
||||
BlockIPv6: configuration.BlockIPv6,
|
||||
MTU: mtu,
|
||||
Address: addresses,
|
||||
Routes: routes,
|
||||
ExcludedRoutes: excludedRoutes,
|
||||
DNS: configuration.DNS,
|
||||
DNSServers: common.Map(configuration.DNSServers, func(server ovpn.TunnelDNSServer) ovpntransport.DNSServer {
|
||||
return ovpntransport.DNSServer{
|
||||
Priority: server.Priority,
|
||||
Addresses: slices.Clone(server.Addresses),
|
||||
ResolveDomains: slices.Clone(server.ResolveDomains),
|
||||
DNSSEC: server.DNSSEC,
|
||||
Transport: server.Transport,
|
||||
SNI: server.SNI,
|
||||
}
|
||||
}),
|
||||
SearchDomains: slices.Clone(configuration.SearchDomains),
|
||||
DNSRoutes: slices.Clone(configuration.DNSRoutes),
|
||||
Topology: configuration.Topology,
|
||||
BlockIPv6: configuration.BlockIPv6,
|
||||
}
|
||||
}
|
||||
|
||||
func buildIPSet(routes []ovpntransport.Route) (*netipx.IPSet, error) {
|
||||
func openVPNRoutesContainAddress(routes []ovpntransport.Route, address netip.Addr) bool {
|
||||
for _, route := range routes {
|
||||
if route.Prefix.Contains(address) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func openVPNRoutesContainPrefix(routes []ovpntransport.Route, prefix netip.Prefix) bool {
|
||||
for _, route := range routes {
|
||||
if route.Prefix == prefix {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func buildIPSet(routes []ovpntransport.Route, excludedRoutes []ovpntransport.Route) (*netipx.IPSet, error) {
|
||||
var builder netipx.IPSetBuilder
|
||||
for _, route := range routes {
|
||||
builder.AddPrefix(route.Prefix)
|
||||
}
|
||||
for _, route := range excludedRoutes {
|
||||
builder.RemovePrefix(route.Prefix)
|
||||
}
|
||||
return builder.IPSet()
|
||||
}
|
||||
|
||||
|
||||
+211
-32
@@ -5,6 +5,7 @@ import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strconv"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
@@ -90,18 +91,16 @@ func NewServerEndpoint(ctx context.Context, router adapter.Router, logger log.Co
|
||||
return nil, err
|
||||
}
|
||||
serverOptions.Context = loopContext
|
||||
serverOptions.Authentication.Authenticator = authenticatorFromUsers(options.Users)
|
||||
serverOptions.Authentication.DuplicateCN = options.DuplicateCN
|
||||
if serverOptions.Mode == ovpn.ModeTLS {
|
||||
serverOptions.Authentication.Authenticator = authenticatorFromUsers(options.Users)
|
||||
serverOptions.Authentication.DuplicateCN = options.DuplicateCN
|
||||
}
|
||||
serverOptions.Logger = logger
|
||||
serverEndpoint.serverOptions = serverOptions
|
||||
udpTimeout := C.UDPTimeout
|
||||
if options.UDPTimeout != 0 {
|
||||
udpTimeout = time.Duration(options.UDPTimeout)
|
||||
}
|
||||
deviceRoutes := make([]ovpntransport.Route, 0, len(options.Address))
|
||||
for _, prefix := range options.Address {
|
||||
deviceRoutes = append(deviceRoutes, ovpntransport.Route{Prefix: prefix.Masked()})
|
||||
}
|
||||
device, err := ovpntransport.NewDevice(ovpntransport.DeviceOptions{
|
||||
Context: ctx,
|
||||
Logger: logger,
|
||||
@@ -118,7 +117,6 @@ func NewServerEndpoint(ctx context.Context, router adapter.Router, logger log.Co
|
||||
Configuration: ovpntransport.Configuration{
|
||||
MTU: options.MTU,
|
||||
Address: options.Address,
|
||||
Routes: deviceRoutes,
|
||||
Topology: options.Topology,
|
||||
},
|
||||
})
|
||||
@@ -134,15 +132,18 @@ func NewServerEndpoint(ctx context.Context, router adapter.Router, logger log.Co
|
||||
func validateServerAddresses(addresses []netip.Prefix) error {
|
||||
var hasIPv4 bool
|
||||
var hasIPv6 bool
|
||||
for _, prefix := range addresses {
|
||||
for addressIndex, prefix := range addresses {
|
||||
if !prefix.IsValid() {
|
||||
return E.New("server address[", addressIndex, "] is invalid")
|
||||
}
|
||||
if prefix.Addr().Is4() {
|
||||
if hasIPv4 {
|
||||
return E.New("multiple IPv4 OpenVPN server address pools are not supported")
|
||||
return E.New("multiple IPv4 server address pools are not supported")
|
||||
}
|
||||
hasIPv4 = true
|
||||
} else {
|
||||
if hasIPv6 {
|
||||
return E.New("multiple IPv6 OpenVPN server address pools are not supported")
|
||||
return E.New("multiple IPv6 server address pools are not supported")
|
||||
}
|
||||
hasIPv6 = true
|
||||
}
|
||||
@@ -155,7 +156,7 @@ func validateServerTopology(topology string) error {
|
||||
case "", "subnet", "p2p", "net30":
|
||||
return nil
|
||||
default:
|
||||
return E.New("invalid OpenVPN topology ", topology, ", allowed values: subnet, p2p, net30")
|
||||
return E.New("invalid topology ", topology, ", allowed values: subnet, p2p, net30")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -258,11 +259,17 @@ func (s *ServerEndpoint) Start(stage adapter.StartStage) error {
|
||||
}
|
||||
|
||||
func buildServerOptions(options option.OpenVPNServerEndpointOptions) (ovpn.ServerOptions, error) {
|
||||
if len(options.Address) == 0 {
|
||||
return ovpn.ServerOptions{}, E.New("missing OpenVPN server address")
|
||||
mode := options.Mode
|
||||
if mode == "" {
|
||||
mode = ovpn.ModeTLS
|
||||
}
|
||||
if options.TLS == nil {
|
||||
return ovpn.ServerOptions{}, E.New("missing `tls` options")
|
||||
switch mode {
|
||||
case ovpn.ModeTLS, ovpn.ModeStaticKey:
|
||||
default:
|
||||
return ovpn.ServerOptions{}, E.New("unsupported mode: ", mode, " (expected \"tls\" or \"static_key\")")
|
||||
}
|
||||
if len(options.Address) == 0 {
|
||||
return ovpn.ServerOptions{}, E.New("missing server address")
|
||||
}
|
||||
err := validateServerAddresses(options.Address)
|
||||
if err != nil {
|
||||
@@ -279,7 +286,16 @@ func buildServerOptions(options option.OpenVPNServerEndpointOptions) (ovpn.Serve
|
||||
switch protocol {
|
||||
case N.NetworkTCP, N.NetworkUDP:
|
||||
default:
|
||||
return ovpn.ServerOptions{}, E.New("unsupported OpenVPN network: ", protocol)
|
||||
return ovpn.ServerOptions{}, E.New("unsupported network: ", protocol)
|
||||
}
|
||||
if mode == ovpn.ModeStaticKey {
|
||||
return buildStaticKeyServerOptions(options, protocol)
|
||||
}
|
||||
if options.TLS == nil {
|
||||
return ovpn.ServerOptions{}, E.New("missing `tls` options")
|
||||
}
|
||||
if len(options.StaticKey) > 0 || options.StaticKeyPath != "" || options.KeyDirection != "" || options.Cipher != "" || options.Remote != "" || options.RemotePort != 0 || netip.Addr(options.PeerAddress).IsValid() || netip.Addr(options.PeerAddressIPv6).IsValid() {
|
||||
return ovpn.ServerOptions{}, E.New("static-key server options require `mode: static_key`")
|
||||
}
|
||||
tlsOptions, keyDirection, err := buildServerTLSOptions(*options.TLS)
|
||||
if err != nil {
|
||||
@@ -295,29 +311,136 @@ func buildServerOptions(options option.OpenVPNServerEndpointOptions) (ovpn.Serve
|
||||
MaxClients: options.MaxClients,
|
||||
},
|
||||
DataChannel: ovpn.ServerDataChannelOptions{
|
||||
MTU: options.MTU,
|
||||
Ciphers: []string(options.DataCiphers),
|
||||
FallbackCipher: options.DataCiphersFallback,
|
||||
Auth: options.Auth,
|
||||
PacketHeadroom: ovpntransport.PacketHeadroom,
|
||||
MTU: options.MTU,
|
||||
MSSFix: options.MSSFix,
|
||||
MSSFixDisabled: options.MSSFixDisabled,
|
||||
MSSFixMode: options.MSSFixMode,
|
||||
Ciphers: []string(options.DataCiphers),
|
||||
FallbackCipher: options.DataCiphersFallback,
|
||||
Auth: options.Auth,
|
||||
ReplayWindow: options.ReplayWindow,
|
||||
ReplayWindowTime: time.Duration(options.ReplayWindowTime),
|
||||
PacketHeadroom: ovpntransport.PacketHeadroom,
|
||||
},
|
||||
TLS: tlsOptions,
|
||||
Timing: ovpn.ServerTimingOptions{
|
||||
RenegotiationInterval: time.Duration(options.RenegotiateInterval),
|
||||
RenegotiationDisabled: options.RenegotiateDisabled,
|
||||
RenegotiationBytes: options.RenegotiateBytes,
|
||||
RenegotiationPackets: options.RenegotiatePackets,
|
||||
HandWindow: time.Duration(options.HandshakeWindow),
|
||||
PingInterval: time.Duration(options.PingInterval),
|
||||
PingRestart: time.Duration(options.PingRestart),
|
||||
},
|
||||
}
|
||||
applyServerPushOptions(&serverOptions, options)
|
||||
err = applyServerPushOptions(&serverOptions, options)
|
||||
if err != nil {
|
||||
return ovpn.ServerOptions{}, err
|
||||
}
|
||||
return serverOptions, nil
|
||||
}
|
||||
|
||||
func buildStaticKeyServerOptions(options option.OpenVPNServerEndpointOptions, protocol string) (ovpn.ServerOptions, error) {
|
||||
if options.TLS != nil {
|
||||
return ovpn.ServerOptions{}, E.New("`tls` options are not supported in `static_key` mode")
|
||||
}
|
||||
if len(options.Users) > 0 || options.DuplicateCN {
|
||||
return ovpn.ServerOptions{}, E.New("user authentication is not supported in `static_key` mode")
|
||||
}
|
||||
if options.Push != nil {
|
||||
return ovpn.ServerOptions{}, E.New("push options are not supported in `static_key` mode")
|
||||
}
|
||||
if options.RenegotiateInterval != 0 || options.RenegotiateDisabled || options.RenegotiateBytes != 0 || options.RenegotiatePackets != 0 || options.HandshakeWindow != 0 {
|
||||
return ovpn.ServerOptions{}, E.New("TLS timing and renegotiation options are not supported in `static_key` mode")
|
||||
}
|
||||
if len(options.DataCiphers) > 0 || options.DataCiphersFallback != "" {
|
||||
return ovpn.ServerOptions{}, E.New("`data_ciphers` and `data_ciphers_fallback` are not supported in `static_key` mode; use `cipher`")
|
||||
}
|
||||
staticKey, err := requiredMaterialSource("static_key", options.StaticKey, options.StaticKeyPath)
|
||||
if err != nil {
|
||||
return ovpn.ServerOptions{}, err
|
||||
}
|
||||
keyDirection, err := keyDirectionValue(options.KeyDirection)
|
||||
if err != nil {
|
||||
return ovpn.ServerOptions{}, err
|
||||
}
|
||||
vpnGateway := netip.Addr(options.PeerAddress)
|
||||
if vpnGateway.IsValid() && !vpnGateway.Is4() {
|
||||
return ovpn.ServerOptions{}, E.New("`peer_address` must be an IPv4 address")
|
||||
}
|
||||
vpnGatewayIPv6 := netip.Addr(options.PeerAddressIPv6)
|
||||
if vpnGatewayIPv6.IsValid() && !vpnGatewayIPv6.Is6() {
|
||||
return ovpn.ServerOptions{}, E.New("`peer_address_ipv6` must be an IPv6 address")
|
||||
}
|
||||
var hasIPv4 bool
|
||||
var hasIPv6 bool
|
||||
for _, address := range options.Address {
|
||||
hasIPv4 = hasIPv4 || address.Addr().Is4()
|
||||
hasIPv6 = hasIPv6 || address.Addr().Is6()
|
||||
}
|
||||
if hasIPv4 && !vpnGateway.IsValid() {
|
||||
return ovpn.ServerOptions{}, E.New("missing `peer_address` for the IPv4 static-key tunnel")
|
||||
}
|
||||
if hasIPv6 && !vpnGatewayIPv6.IsValid() {
|
||||
return ovpn.ServerOptions{}, E.New("missing `peer_address_ipv6` for the IPv6 static-key tunnel")
|
||||
}
|
||||
if vpnGateway.IsValid() && !hasIPv4 {
|
||||
return ovpn.ServerOptions{}, E.New("`peer_address` requires an IPv4 tunnel `address` in `static_key` mode")
|
||||
}
|
||||
if vpnGatewayIPv6.IsValid() && !hasIPv6 {
|
||||
return ovpn.ServerOptions{}, E.New("`peer_address_ipv6` requires an IPv6 tunnel `address` in `static_key` mode")
|
||||
}
|
||||
remoteAddress := ""
|
||||
if protocol == N.NetworkUDP {
|
||||
if options.Remote == "" || options.RemotePort == 0 {
|
||||
return ovpn.ServerOptions{}, E.New("`remote` and `remote_port` are required for a UDP static-key server")
|
||||
}
|
||||
remoteAddress = net.JoinHostPort(options.Remote, strconv.Itoa(int(options.RemotePort)))
|
||||
} else if options.Remote != "" || options.RemotePort != 0 {
|
||||
return ovpn.ServerOptions{}, E.New("`remote` and `remote_port` are only used by a UDP static-key server")
|
||||
}
|
||||
topology := options.Topology
|
||||
if topology == "" {
|
||||
topology = "p2p"
|
||||
}
|
||||
return ovpn.ServerOptions{
|
||||
Mode: ovpn.ModeStaticKey,
|
||||
StaticKey: staticKey,
|
||||
KeyDirection: keyDirection,
|
||||
Transport: ovpn.ServerTransportOptions{
|
||||
Protocol: protocol,
|
||||
RemoteAddress: remoteAddress,
|
||||
},
|
||||
Resources: ovpn.ServerResourceOptions{MaxClients: options.MaxClients},
|
||||
DataChannel: ovpn.ServerDataChannelOptions{
|
||||
MTU: options.MTU,
|
||||
MSSFix: options.MSSFix,
|
||||
MSSFixDisabled: options.MSSFixDisabled,
|
||||
MSSFixMode: options.MSSFixMode,
|
||||
Cipher: options.Cipher,
|
||||
Auth: options.Auth,
|
||||
ReplayWindow: options.ReplayWindow,
|
||||
ReplayWindowTime: time.Duration(options.ReplayWindowTime),
|
||||
PacketHeadroom: ovpntransport.PacketHeadroom,
|
||||
},
|
||||
Timing: ovpn.ServerTimingOptions{
|
||||
PingInterval: time.Duration(options.PingInterval),
|
||||
PingRestart: time.Duration(options.PingRestart),
|
||||
},
|
||||
Tunnel: ovpn.ServerTunnelOptions{
|
||||
Topology: topology,
|
||||
LocalAddress: slices.Clone(options.Address),
|
||||
VPNGateway: vpnGateway,
|
||||
VPNGatewayIPv6: vpnGatewayIPv6,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildServerTLSOptions(options option.OpenVPNInboundTLSOptions) (ovpn.ServerTLSOptions, int, error) {
|
||||
switch options.VerifyClientCertificate {
|
||||
case "", "require", "optional", "none":
|
||||
default:
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("invalid OpenVPN client certificate policy ", options.VerifyClientCertificate, ", allowed values: require, optional, none")
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("invalid client certificate policy ", options.VerifyClientCertificate, ", allowed values: require, optional, none")
|
||||
}
|
||||
certificate, err := requiredMaterialSource("tls.certificate", options.Certificate, options.CertificatePath)
|
||||
if err != nil {
|
||||
@@ -327,16 +450,46 @@ func buildServerTLSOptions(options option.OpenVPNInboundTLSOptions) (ovpn.Server
|
||||
if err != nil {
|
||||
return ovpn.ServerTLSOptions{}, 0, err
|
||||
}
|
||||
certificateAuthority, err := requiredMaterialSource("tls.client_certificate", options.ClientCertificate, options.ClientCertificatePath)
|
||||
certificateAuthority, err := materialSource("tls.client_certificate", options.ClientCertificate, options.ClientCertificatePath)
|
||||
if err != nil {
|
||||
return ovpn.ServerTLSOptions{}, 0, err
|
||||
}
|
||||
remoteCertificateTLS := options.RemoteCertificateTLS
|
||||
switch remoteCertificateTLS {
|
||||
case "", "server", "client", "none":
|
||||
default:
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("invalid `tls.remote_certificate_tls`: ", remoteCertificateTLS)
|
||||
}
|
||||
if options.RemoteCertificateEKU != "" && remoteCertificateTLS != "" {
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("`tls.remote_certificate_eku` is conflict with `tls.remote_certificate_tls`")
|
||||
}
|
||||
if remoteCertificateTLS == "" && options.RemoteCertificateEKU == "" {
|
||||
remoteCertificateTLS = "client"
|
||||
} else if remoteCertificateTLS == "none" {
|
||||
remoteCertificateTLS = ""
|
||||
}
|
||||
clientNameType := options.ClientNameType
|
||||
if options.ClientName != "" && clientNameType == "" {
|
||||
clientNameType = "name"
|
||||
}
|
||||
tlsOptions := ovpn.ServerTLSOptions{
|
||||
CertificateAuthority: certificateAuthority,
|
||||
Certificate: certificate,
|
||||
Key: key,
|
||||
VerifyClientCertificate: options.VerifyClientCertificate,
|
||||
VerifyX509Name: options.ClientName,
|
||||
VerifyX509Type: clientNameType,
|
||||
PeerFingerprint: options.PeerFingerprint,
|
||||
CRLVerify: options.CRLPath,
|
||||
RemoteCertificateKU: options.RemoteCertificateKU,
|
||||
RemoteCertificateEKU: options.RemoteCertificateEKU,
|
||||
RemoteCertificateTLS: remoteCertificateTLS,
|
||||
NSCertificateType: options.NSCertificateType,
|
||||
CertificateProfile: options.CertificateProfile,
|
||||
VersionMin: options.VersionMin,
|
||||
VersionMax: options.VersionMax,
|
||||
Cipher: options.Cipher,
|
||||
Groups: options.Groups,
|
||||
}
|
||||
keyDirection := -1
|
||||
controlWrap := options.ControlWrap
|
||||
@@ -369,15 +522,15 @@ func buildServerTLSOptions(options option.OpenVPNInboundTLSOptions) (ovpn.Server
|
||||
tlsOptions.CryptV2ForceCookie = controlWrap.ForceCookie
|
||||
}
|
||||
case "":
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("missing OpenVPN control wrap type")
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("missing control wrap type")
|
||||
default:
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("unknown OpenVPN control wrap type: ", controlWrap.Type)
|
||||
return ovpn.ServerTLSOptions{}, 0, E.New("unknown control wrap type: ", controlWrap.Type)
|
||||
}
|
||||
}
|
||||
return tlsOptions, keyDirection, nil
|
||||
}
|
||||
|
||||
func applyServerPushOptions(serverOptions *ovpn.ServerOptions, options option.OpenVPNServerEndpointOptions) {
|
||||
func applyServerPushOptions(serverOptions *ovpn.ServerOptions, options option.OpenVPNServerEndpointOptions) error {
|
||||
topology := options.Topology
|
||||
if topology == "" {
|
||||
topology = "subnet"
|
||||
@@ -399,10 +552,35 @@ func applyServerPushOptions(serverOptions *ovpn.ServerOptions, options option.Op
|
||||
LocalAddress: localAddresses,
|
||||
}
|
||||
if options.Push == nil {
|
||||
return
|
||||
return nil
|
||||
}
|
||||
serverOptions.Push.Routes = slices.Clone(options.Push.Routes)
|
||||
serverOptions.Push.DNS = slices.Clone(options.Push.DNS)
|
||||
serverOptions.Push.SearchDomains = slices.Clone(options.Push.SearchDomains)
|
||||
serverOptions.Push.DHCPOptions = slices.Clone(options.Push.DHCPOptions)
|
||||
for serverIndex, server := range options.Push.DNSServers {
|
||||
addresses := make([]netip.AddrPort, 0, len(server.Addresses))
|
||||
for addressIndex, addressValue := range server.Addresses {
|
||||
address, err := netip.ParseAddr(addressValue)
|
||||
if err == nil {
|
||||
addresses = append(addresses, netip.AddrPortFrom(address, 0))
|
||||
continue
|
||||
}
|
||||
addressPort, addressPortErr := netip.ParseAddrPort(addressValue)
|
||||
if addressPortErr != nil || addressPort.Port() == 0 {
|
||||
return E.New("invalid push.dns_servers[", serverIndex, "].addresses[", addressIndex, "]: ", addressValue)
|
||||
}
|
||||
addresses = append(addresses, addressPort)
|
||||
}
|
||||
serverOptions.Push.DNSServers = append(serverOptions.Push.DNSServers, ovpn.TunnelDNSServer{
|
||||
Priority: server.Priority,
|
||||
Addresses: addresses,
|
||||
ResolveDomains: slices.Clone(server.ResolveDomains),
|
||||
DNSSEC: server.DNSSEC,
|
||||
Transport: server.Transport,
|
||||
SNI: server.SNI,
|
||||
})
|
||||
}
|
||||
serverOptions.Push.BlockOutsideDNS = options.Push.BlockOutsideDNS
|
||||
serverOptions.Push.PingInterval = time.Duration(options.Push.PingInterval)
|
||||
serverOptions.Push.PingRestart = time.Duration(options.Push.PingRestart)
|
||||
@@ -414,6 +592,7 @@ func applyServerPushOptions(serverOptions *ovpn.ServerOptions, options option.Op
|
||||
serverOptions.Push.RedirectGatewayFlags = []string{"def1"}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *ServerEndpoint) readLoop() {
|
||||
@@ -424,7 +603,7 @@ func (s *ServerEndpoint) readLoop() {
|
||||
if E.IsClosedOrCanceled(err) || s.loopContext.Err() != nil {
|
||||
return
|
||||
}
|
||||
s.logger.Error(E.Cause(err, "OpenVPN server terminated"))
|
||||
s.logger.Error(E.Cause(err, "server terminated"))
|
||||
return
|
||||
}
|
||||
packetBuffers := make([]*buf.Buffer, len(serverPacketBuffers))
|
||||
@@ -491,7 +670,7 @@ func (s *ServerEndpoint) NewDNSPacket(payload []byte, source M.Socksaddr, destin
|
||||
|
||||
func (s *ServerEndpoint) WritePackets(packets [][]byte) error {
|
||||
if !s.started.Load() {
|
||||
return E.New("OpenVPN server is not ready yet")
|
||||
return E.New("endpoint is not ready yet")
|
||||
}
|
||||
packetBuffers := make([]*buf.Buffer, len(packets))
|
||||
for i, packet := range packets {
|
||||
@@ -547,7 +726,7 @@ func (s *ServerEndpoint) DialContext(ctx context.Context, network string, destin
|
||||
s.logger.InfoContext(ctx, "outbound packet connection to ", destination)
|
||||
}
|
||||
if !s.started.Load() {
|
||||
return nil, E.New("OpenVPN server is not ready yet")
|
||||
return nil, E.New("endpoint is not ready yet")
|
||||
}
|
||||
if destination.IsDomain() {
|
||||
destinationAddresses, err := s.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
|
||||
@@ -565,7 +744,7 @@ func (s *ServerEndpoint) DialContext(ctx context.Context, network string, destin
|
||||
func (s *ServerEndpoint) ListenPacketWithDestination(ctx context.Context, destination M.Socksaddr) (net.PacketConn, netip.Addr, error) {
|
||||
s.logger.InfoContext(ctx, "outbound packet connection to ", destination)
|
||||
if !s.started.Load() {
|
||||
return nil, netip.Addr{}, E.New("OpenVPN server is not ready yet")
|
||||
return nil, netip.Addr{}, E.New("endpoint is not ready yet")
|
||||
}
|
||||
if destination.IsDomain() {
|
||||
destinationAddresses, err := s.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
|
||||
|
||||
Reference in New Issue
Block a user