Improve OpenVPN & OpenConnect interoperability

This commit is contained in:
世界
2026-07-21 09:38:21 +08:00
parent 5cad5ad42d
commit 2ff294c4f4
47 changed files with 3915 additions and 440 deletions
+100 -32
View File
@@ -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 {
+373
View File
@@ -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
}
+92
View File
@@ -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
View File
@@ -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)
}
+427
View File
@@ -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
View File
@@ -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
View File
@@ -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{})