mirror of
https://github.com/AdguardTeam/AdGuardHome.git
synced 2026-08-04 15:28:58 +00:00
Pull request 2718: AGDNS-3720-refactor-tls-vol.1
Some checks are pending
build / test (macOS-latest) (push) Waiting to run
build / test (ubuntu-latest) (push) Waiting to run
build / test (windows-latest) (push) Waiting to run
build / build-release (push) Blocked by required conditions
build / notify (push) Blocked by required conditions
lint / go-lint (push) Waiting to run
lint / eslint (push) Waiting to run
lint / notify (push) Blocked by required conditions
Some checks are pending
build / test (macOS-latest) (push) Waiting to run
build / test (ubuntu-latest) (push) Waiting to run
build / test (windows-latest) (push) Waiting to run
build / build-release (push) Blocked by required conditions
build / notify (push) Blocked by required conditions
lint / go-lint (push) Waiting to run
lint / eslint (push) Waiting to run
lint / notify (push) Blocked by required conditions
Squashed commit of the following: commit e4d7f8a690f7e0731b86f41c05b42f816a4f4c38 Merge: d4c296f79336e9c9dfAuthor: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Mon Jul 27 17:03:08 2026 +0300 Merge branch 'master' into AGDNS-3720-refactor-tls-vol.1 commit d4c296f799db3683d51b1eebc5092621ce412104 Author: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Mon Jul 27 11:50:20 2026 +0300 all: upd chlog; commit 403f0c9c0593befb6d03ea28baa0d424ac7ea8bc Author: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Mon Jul 27 11:41:52 2026 +0300 all: upd chlog; commitba44e52537Author: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Fri Jul 24 18:50:02 2026 +0300 home: rm supports certificate; commit4ebec73863Author: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Wed Jul 22 13:04:17 2026 +0300 internal: imp code; imp tests; commit715baa5bfaAuthor: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Tue Jul 21 17:24:22 2026 +0300 internal: upd interface; do not check certificates any more; commit92592a017dAuthor: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Mon Jul 20 14:25:16 2026 +0300 dnsforward: dont ignore error from origGetCert; commitafc6b1ba76Author: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Mon Jul 20 13:46:32 2026 +0300 dnsforward: add common name fallback; home: imp onGetCertificate; commit30b71c50ceAuthor: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Thu Jul 16 18:48:15 2026 +0300 dnsforward: imp code; fix bugs; home: store tls cert in tls manager; commit3b55b20cb2Author: Maksim Kazantsev <m.kazantsev@adguard.com> Date: Wed Jul 15 17:57:16 2026 +0300 internal: impl tls config provider; impl automatic tls certificate replacement; get tls config from tls config provider;
This commit is contained in:
parent
336e9c9df4
commit
254e17dabf
21 changed files with 435 additions and 224 deletions
|
|
@ -30,6 +30,10 @@ NOTE: Add new changes BELOW THIS COMMENT.
|
|||
|
||||
- The `edge` channel has been switched to the new UI and versioning scheme.
|
||||
|
||||
### Deprecated
|
||||
|
||||
- `strict_sni_check` is now deprecated.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Multiple inaccuracies in the OpenAPI specification:
|
||||
|
|
|
|||
|
|
@ -204,10 +204,10 @@ func (m *Registrar) Register(method, path string, h http.HandlerFunc) {
|
|||
|
||||
// TLSConfigProvider is a fake [aghtls.TLSConfigProvider] implementation for
|
||||
// tests.
|
||||
// TODO(m.kazantsev): Use in tests.
|
||||
type TLSConfigProvider struct {
|
||||
OnTLSConfig func() (conf *tls.Config)
|
||||
OnRootCAs func() (cert *x509.CertPool)
|
||||
OnTLSConfig func() (conf *tls.Config)
|
||||
OnRootCAs func() (cert *x509.CertPool)
|
||||
OnHasIPAddrs func() (ok bool)
|
||||
}
|
||||
|
||||
// type check
|
||||
|
|
@ -224,3 +224,9 @@ func (t *TLSConfigProvider) TLSConfig() (conf *tls.Config) {
|
|||
func (t *TLSConfigProvider) RootCAs() (pool *x509.CertPool) {
|
||||
return t.OnRootCAs()
|
||||
}
|
||||
|
||||
// HasIPAddrs implements the [aghtls.TLSConfigProvider] interface for
|
||||
// *TLSConfigProvider.
|
||||
func (t *TLSConfigProvider) HasIPAddrs() (ok bool) {
|
||||
return t.OnHasIPAddrs()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import (
|
|||
// must be safe for concurrent use.
|
||||
//
|
||||
// TODO(m.kazantsev): Merge with the Manager interface.
|
||||
// TODO(m.kazantsev): Add at least one real implementation.
|
||||
type TLSConfigProvider interface {
|
||||
// TLSConfig returns a clone of the current TLS configuration. conf
|
||||
// provides its certificates via GetConfigForClient method.
|
||||
|
|
@ -17,6 +16,10 @@ type TLSConfigProvider interface {
|
|||
|
||||
// RootCAs returns the current root CA pool.
|
||||
RootCAs() (root *x509.CertPool)
|
||||
|
||||
// HasIPAddrs returns true if the current TLS configuration has at least one
|
||||
// certificate with an IP address in its SAN extension.
|
||||
HasIPAddrs() (ok bool)
|
||||
}
|
||||
|
||||
// type check
|
||||
|
|
@ -37,3 +40,9 @@ func (EmptyTLSConfigProvider) TLSConfig() (conf *tls.Config) {
|
|||
func (EmptyTLSConfigProvider) RootCAs() (root *x509.CertPool) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// HasIPAddrs implements the [TLSConfigProvider] interface for
|
||||
// EmptyTLSConfigProvider. It always returns false.
|
||||
func (EmptyTLSConfigProvider) HasIPAddrs() (ok bool) {
|
||||
return false
|
||||
}
|
||||
|
|
|
|||
|
|
@ -17,7 +17,6 @@ import (
|
|||
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghslog"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghtls"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/client"
|
||||
"github.com/AdguardTeam/dnscrypt"
|
||||
"github.com/AdguardTeam/dnsproxy/proxy"
|
||||
|
|
@ -191,16 +190,14 @@ type TLSConfig struct {
|
|||
// It is nil if the DNSCrypt server is disabled.
|
||||
DNSCryptConf *DNSCryptConfig
|
||||
|
||||
// Cert is the TLS certificate used for TLS connections. It is nil if
|
||||
// encryption is disabled.
|
||||
Cert *tls.Certificate
|
||||
|
||||
// TLSListenAddrs are the addresses to listen on for DoT connections. Each
|
||||
// item in the list must be non-nil if Cert is not nil.
|
||||
// item in the list must be non-nil if TLS has at least one valid
|
||||
// certificate.
|
||||
TLSListenAddrs []*net.TCPAddr
|
||||
|
||||
// QUICListenAddrs are the addresses to listen on for DoQ connections. Each
|
||||
// item in the list must be non-nil if Cert is not nil.
|
||||
// item in the list must be non-nil if TLS has at least one valid
|
||||
// certificate.
|
||||
QUICListenAddrs []*net.UDPAddr
|
||||
|
||||
// HTTPSListenAddrs should be the addresses AdGuard Home is listening on for
|
||||
|
|
@ -326,6 +323,7 @@ const (
|
|||
)
|
||||
|
||||
// newProxyConfig creates and validates configuration for the main proxy.
|
||||
// s.serverLock must be locked.
|
||||
//
|
||||
// TODO(d.kolyshev): Improve maintainability.
|
||||
func (s *Server) newProxyConfig(ctx context.Context) (conf *proxy.Config, err error) {
|
||||
|
|
@ -385,10 +383,7 @@ func (s *Server) newProxyConfig(ctx context.Context) (conf *proxy.Config, err er
|
|||
return nil, fmt.Errorf("bogus_nxdomain: %w", err)
|
||||
}
|
||||
|
||||
err = s.prepareTLS(ctx, conf)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("validating tls: %w", err)
|
||||
}
|
||||
s.prepareTLS(ctx, conf)
|
||||
|
||||
err = s.preparePlain(ctx, conf)
|
||||
if err != nil {
|
||||
|
|
@ -707,55 +702,28 @@ func (s *Server) prepareDNSCrypt(proxyConf *proxy.Config) {
|
|||
proxyConf.DNSCryptResolverCert = dnsCryptConf.ResolverCert
|
||||
}
|
||||
|
||||
// prepareTLS sets up the TLS configuration for the DNS proxy.
|
||||
func (s *Server) prepareTLS(ctx context.Context, proxyConf *proxy.Config) (err error) {
|
||||
// prepareTLS sets up the TLS configuration for the DNS proxy. s.serverLock
|
||||
// must be locked. proxyConf must be non-nil and valid.
|
||||
func (s *Server) prepareTLS(ctx context.Context, proxyConf *proxy.Config) {
|
||||
s.prepareDNSCrypt(proxyConf)
|
||||
|
||||
if s.conf.TLSConf.Cert == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if s.conf.TLSConf.TLSListenAddrs == nil && s.conf.TLSConf.QUICListenAddrs == nil {
|
||||
return nil
|
||||
return
|
||||
}
|
||||
|
||||
proxyConf.TLSListenAddr = s.conf.TLSConf.TLSListenAddrs
|
||||
proxyConf.QUICListenAddr = s.conf.TLSConf.QUICListenAddrs
|
||||
|
||||
cert, err := x509.ParseCertificate(s.conf.TLSConf.Cert.Certificate[0])
|
||||
if err != nil {
|
||||
return fmt.Errorf("x509.ParseCertificate(): %w", err)
|
||||
proxyConf.TLSConfig = s.tlsConfigProvider.TLSConfig()
|
||||
if proxyConf.TLSConfig == nil {
|
||||
s.logger.WarnContext(ctx, "tls configuration is not set")
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
s.hasIPAddrs = aghtls.CertificateHasIP(cert)
|
||||
|
||||
if s.conf.TLSConf.StrictSNICheck {
|
||||
if len(cert.DNSNames) != 0 {
|
||||
s.dnsNames = cert.DNSNames
|
||||
s.logger.DebugContext(
|
||||
ctx,
|
||||
"using certificate's SAN as DNS names",
|
||||
"dns_names", cert.DNSNames,
|
||||
)
|
||||
slices.Sort(s.dnsNames)
|
||||
} else {
|
||||
s.dnsNames = []string{cert.Subject.CommonName}
|
||||
s.logger.DebugContext(
|
||||
ctx,
|
||||
"using certificate's CN as DNS name",
|
||||
"common_name",
|
||||
cert.Subject.CommonName,
|
||||
)
|
||||
}
|
||||
if s.conf.TLSConf.StrictSNICheck && proxyConf.TLSConfig.GetCertificate != nil {
|
||||
s.replaceGetCertificate(proxyConf.TLSConfig)
|
||||
}
|
||||
|
||||
proxyConf.TLSConfig = &tls.Config{
|
||||
GetCertificate: s.onGetCertificate,
|
||||
CipherSuites: s.conf.TLSCiphers,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// isWildcard returns true if host is a wildcard hostname.
|
||||
|
|
@ -790,22 +758,44 @@ func anyNameMatches(dnsNames []string, sni string) (ok bool) {
|
|||
return false
|
||||
}
|
||||
|
||||
// onGetCertificate is called by [tls] package when Client Hello is received. If
|
||||
// the server name (from SNI) supplied by client is incorrect - we terminate the
|
||||
// ongoing TLS handshake.
|
||||
func (s *Server) onGetCertificate(ch *tls.ClientHelloInfo) (*tls.Certificate, error) {
|
||||
if s.conf.TLSConf.StrictSNICheck && !anyNameMatches(s.dnsNames, ch.ServerName) {
|
||||
// TODO(s.chzhen): Pass context.
|
||||
s.logger.WarnContext(
|
||||
context.TODO(),
|
||||
"unknown SNI in Client Hello",
|
||||
"server_name", ch.ServerName,
|
||||
)
|
||||
// replaceGetCertificate replaces the TLS.Config.GetCertificate with a wrapped
|
||||
// version of the previous one, adding a SNI check. It must be called only once
|
||||
// for each instance of orig. orig must not be nil and orig.GetCertificate must
|
||||
// be populated.
|
||||
//
|
||||
// TODO(m.kazantsev): Consider moving this method to aghtls.
|
||||
func (s *Server) replaceGetCertificate(orig *tls.Config) {
|
||||
origGetCert := orig.GetCertificate
|
||||
|
||||
return nil, fmt.Errorf("invalid SNI")
|
||||
orig.GetCertificate = func(chi *tls.ClientHelloInfo) (cert *tls.Certificate, err error) {
|
||||
cert, err = origGetCert(chi)
|
||||
if err != nil {
|
||||
// Don't wrap the error, because it is informative enough as is.
|
||||
return nil, err
|
||||
}
|
||||
if cert == nil || cert.Leaf == nil {
|
||||
return nil, errors.Error("tls certificate is not set")
|
||||
}
|
||||
|
||||
var dnsNames []string
|
||||
if len(cert.Leaf.DNSNames) == 0 {
|
||||
dnsNames = []string{cert.Leaf.Subject.CommonName}
|
||||
} else {
|
||||
dnsNames = cert.Leaf.DNSNames
|
||||
}
|
||||
|
||||
if !anyNameMatches(dnsNames, chi.ServerName) {
|
||||
s.logger.WarnContext(
|
||||
chi.Context(),
|
||||
"unknown sni in client hello",
|
||||
"server_name", chi.ServerName,
|
||||
)
|
||||
|
||||
return nil, errors.Error("unknown sni in client hello")
|
||||
}
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
return s.conf.TLSConf.Cert, nil
|
||||
}
|
||||
|
||||
// preparePlain prepares the plain-DNS configuration for the DNS proxy. The
|
||||
|
|
|
|||
|
|
@ -310,6 +310,7 @@ func TestServer_ServeDNS_dns64(t *testing.T) {
|
|||
LocalPTRResolvers: []string{localUpsAddr},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
startDeferStop(t, s)
|
||||
|
|
@ -354,6 +355,7 @@ func TestServer_dns64WithDisabledRDNS(t *testing.T) {
|
|||
LocalPTRResolvers: []string{localUpsAddr},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
startDeferStop(t, s)
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ import (
|
|||
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghslog"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghtls"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/client"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/querylog"
|
||||
|
|
@ -124,6 +125,10 @@ type Server struct {
|
|||
// PTR resolving.
|
||||
sysResolvers SystemResolvers
|
||||
|
||||
// tlsConfigProvider provides TLS configuration for the server. It must not
|
||||
// be nil.
|
||||
tlsConfigProvider aghtls.TLSConfigProvider
|
||||
|
||||
// access drops disallowed clients.
|
||||
access *accessManager
|
||||
|
||||
|
|
@ -169,10 +174,6 @@ type Server struct {
|
|||
// [upstream.Resolver] interface.
|
||||
bootResolvers []*upstream.UpstreamResolver
|
||||
|
||||
// dnsNames are the DNS names from certificate (SAN) or CN value from
|
||||
// Subject.
|
||||
dnsNames []string
|
||||
|
||||
// conf is the current configuration of the server.
|
||||
conf ServerConfig
|
||||
|
||||
|
|
@ -185,10 +186,6 @@ type Server struct {
|
|||
|
||||
// isRunning is true if the DNS server is running.
|
||||
isRunning bool
|
||||
|
||||
// hasIPAddrs is set during the certificate parsing and is true if the
|
||||
// configured certificate contains at least a single IP address.
|
||||
hasIPAddrs bool
|
||||
}
|
||||
|
||||
// defaultLocalDomainSuffix is the default suffix used to detect internal hosts
|
||||
|
|
@ -207,6 +204,10 @@ type DNSCreateParams struct {
|
|||
Anonymizer *aghnet.IPMut
|
||||
EtcHosts *aghnet.HostsContainer
|
||||
|
||||
// TLSConfigProvider provides a TLS configuration for the server. It must
|
||||
// not be nil.
|
||||
TLSConfigProvider aghtls.TLSConfigProvider
|
||||
|
||||
// Logger is used as a base logger. It must not be nil.
|
||||
Logger *slog.Logger
|
||||
|
||||
|
|
@ -255,6 +256,7 @@ func NewServer(p DNSCreateParams) (s *Server, err error) {
|
|||
conf: ServerConfig{
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
tlsConfigProvider: p.TLSConfigProvider,
|
||||
}
|
||||
|
||||
s.sysResolvers, err = sysresolv.NewSystemResolvers(nil, defaultPlainDNSPort)
|
||||
|
|
@ -478,8 +480,9 @@ func (s *Server) startLocked(ctx context.Context) error {
|
|||
return err
|
||||
}
|
||||
|
||||
// Prepare initializes parameters of s using data from conf. conf must not be
|
||||
// nil.
|
||||
// Prepare initializes parameters of s using data from conf. It can be called
|
||||
// from outside of the package and without acquired s.serverLock only while the
|
||||
// initialization. conf must be non-nil and valid.
|
||||
func (s *Server) Prepare(ctx context.Context, conf *ServerConfig) (err error) {
|
||||
s.conf = *conf
|
||||
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ import (
|
|||
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghos"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghtest"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghtls"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/client"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/filtering/hashprefix"
|
||||
|
|
@ -45,6 +46,9 @@ import (
|
|||
// testLogger is a logger used in tests.
|
||||
var testLogger = slogutil.NewDiscardLogger()
|
||||
|
||||
// testTLSConfigProvider is an empty TLS config provider for tests.
|
||||
var testTLSConfigProvider = aghtls.EmptyTLSConfigProvider{}
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
testutil.DiscardLogOutput(m)
|
||||
}
|
||||
|
|
@ -137,6 +141,7 @@ func createTestServer(
|
|||
tb testing.TB,
|
||||
filterConf *filtering.Config,
|
||||
forwardConf ServerConfig,
|
||||
tlsConfigProvider aghtls.TLSConfigProvider,
|
||||
) (s *Server) {
|
||||
tb.Helper()
|
||||
|
||||
|
|
@ -169,10 +174,11 @@ func createTestServer(
|
|||
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
|
||||
}
|
||||
s, err = NewServer(DNSCreateParams{
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: tlsConfigProvider,
|
||||
})
|
||||
require.NoError(tb, err)
|
||||
|
||||
|
|
@ -209,6 +215,7 @@ func createServerTLSConfig(tb testing.TB) (*tls.Config, []byte, []byte) {
|
|||
IsCA: true,
|
||||
}
|
||||
template.DNSNames = append(template.DNSNames, tlsServerName)
|
||||
template.IPAddresses = append(template.IPAddresses, netutil.IPv4Localhost().AsSlice())
|
||||
|
||||
derBytes, err := x509.CreateCertificate(
|
||||
rand.Reader,
|
||||
|
|
@ -227,23 +234,36 @@ func createServerTLSConfig(tb testing.TB) (*tls.Config, []byte, []byte) {
|
|||
cert, err := tls.X509KeyPair(certPem, keyPem)
|
||||
require.NoErrorf(tb, err, "failed to create certificate: %s", err)
|
||||
|
||||
getCert := func(chi *tls.ClientHelloInfo) (*tls.Certificate, error) {
|
||||
return &cert, nil
|
||||
}
|
||||
|
||||
return &tls.Config{
|
||||
Certificates: []tls.Certificate{cert},
|
||||
ServerName: tlsServerName,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
GetCertificate: getCert,
|
||||
ServerName: tlsServerName,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}, certPem, keyPem
|
||||
}
|
||||
|
||||
func createTestTLS(tb testing.TB, tlsConf *TLSConfig) (s *Server, certPem []byte) {
|
||||
func createTestTLS(
|
||||
tb testing.TB,
|
||||
tlsConf *TLSConfig,
|
||||
) (s *Server, confProvider aghtls.TLSConfigProvider) {
|
||||
tb.Helper()
|
||||
|
||||
var keyPem []byte
|
||||
_, certPem, keyPem = createServerTLSConfig(tb)
|
||||
tlsConfig, certPem, _ := createServerTLSConfig(tb)
|
||||
|
||||
cert, err := tls.X509KeyPair(certPem, keyPem)
|
||||
require.NoError(tb, err)
|
||||
// Add our self-signed generated config to roots.
|
||||
roots := x509.NewCertPool()
|
||||
roots.AppendCertsFromPEM(certPem)
|
||||
|
||||
tlsConf.Cert = &cert
|
||||
tlsConfig.RootCAs = roots
|
||||
|
||||
tlsConfProvider := &aghtest.TLSConfigProvider{}
|
||||
|
||||
tlsConfProvider.OnTLSConfig = func() (conf *tls.Config) { return tlsConfig.Clone() }
|
||||
tlsConfProvider.OnRootCAs = func() (pool *x509.CertPool) { return roots }
|
||||
tlsConfProvider.OnHasIPAddrs = func() (ok bool) { return true }
|
||||
|
||||
s = createTestServer(
|
||||
tb,
|
||||
|
|
@ -261,12 +281,13 @@ func createTestTLS(tb testing.TB, tlsConf *TLSConfig) (s *Server, certPem []byte
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
tlsConfProvider,
|
||||
)
|
||||
|
||||
err = s.Prepare(testutil.ContextWithTimeout(tb, testTimeout), &s.conf)
|
||||
err := s.Prepare(testutil.ContextWithTimeout(tb, testTimeout), &s.conf)
|
||||
require.NoErrorf(tb, err, "failed to prepare server: %s", err)
|
||||
|
||||
return s, certPem
|
||||
return s, tlsConfProvider
|
||||
}
|
||||
|
||||
const googleDomainName = "google-public-dns-a.google.com."
|
||||
|
|
@ -401,6 +422,7 @@ func TestServer(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{newGoogleUpstream()}
|
||||
startDeferStop(t, s)
|
||||
|
|
@ -446,8 +468,9 @@ func TestServer_timeout(t *testing.T) {
|
|||
}
|
||||
|
||||
s, err := NewServer(DNSCreateParams{
|
||||
DNSFilter: createTestDNSFilter(t),
|
||||
Logger: testLogger,
|
||||
DNSFilter: createTestDNSFilter(t),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -459,8 +482,9 @@ func TestServer_timeout(t *testing.T) {
|
|||
|
||||
t.Run("default", func(t *testing.T) {
|
||||
s, err := NewServer(DNSCreateParams{
|
||||
DNSFilter: createTestDNSFilter(t),
|
||||
Logger: testLogger,
|
||||
DNSFilter: createTestDNSFilter(t),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -493,7 +517,8 @@ func TestServer_Prepare_fallbacks(t *testing.T) {
|
|||
}
|
||||
|
||||
s, err := NewServer(DNSCreateParams{
|
||||
Logger: testLogger,
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -521,6 +546,7 @@ func TestServerWithProtectionDisabled(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{newGoogleUpstream()}
|
||||
|
|
@ -537,20 +563,13 @@ func TestServerWithProtectionDisabled(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestDoTServer(t *testing.T) {
|
||||
s, certPem := createTestTLS(t, &TLSConfig{
|
||||
s, tlsConfProvider := createTestTLS(t, &TLSConfig{
|
||||
TLSListenAddrs: []*net.TCPAddr{{}},
|
||||
})
|
||||
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{newGoogleUpstream()}
|
||||
startDeferStop(t, s)
|
||||
|
||||
// Add our self-signed generated config to roots.
|
||||
roots := x509.NewCertPool()
|
||||
roots.AppendCertsFromPEM(certPem)
|
||||
tlsConfig := &tls.Config{
|
||||
ServerName: tlsServerName,
|
||||
RootCAs: roots,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
tlsConfig := tlsConfProvider.TLSConfig()
|
||||
|
||||
// Create a DNS-over-TLS client connection.
|
||||
addr := s.dnsProxy.Addr(proxy.ProtoTLS)
|
||||
|
|
@ -605,7 +624,7 @@ func TestServerRace(t *testing.T) {
|
|||
ConfModifier: agh.EmptyConfigModifier{},
|
||||
ServePlainDNS: true,
|
||||
}
|
||||
s := createTestServer(t, filterConf, forwardConf)
|
||||
s := createTestServer(t, filterConf, forwardConf, testTLSConfigProvider)
|
||||
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{newGoogleUpstream()}
|
||||
startDeferStop(t, s)
|
||||
|
||||
|
|
@ -660,7 +679,7 @@ func TestSafeSearch(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
}
|
||||
s := createTestServer(t, filterConf, forwardConf)
|
||||
s := createTestServer(t, filterConf, forwardConf, testTLSConfigProvider)
|
||||
|
||||
pt := testutil.NewPanicT(t)
|
||||
ups := aghtest.NewUpstreamMock(func(req *dns.Msg) (resp *dns.Msg, err error) {
|
||||
|
|
@ -756,6 +775,7 @@ func TestInvalidRequest(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
startDeferStop(t, s)
|
||||
|
||||
|
|
@ -796,6 +816,7 @@ func TestBlockedRequest(t *testing.T) {
|
|||
BlockingMode: filtering.BlockingModeDefault,
|
||||
},
|
||||
forwardConf,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
startDeferStop(t, s)
|
||||
|
||||
|
|
@ -837,6 +858,7 @@ func TestServerCustomClientUpstream(t *testing.T) {
|
|||
t,
|
||||
&filtering.Config{BlockingMode: filtering.BlockingModeDefault},
|
||||
forwardConf,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
ups := aghtest.NewUpstreamMock(func(req *dns.Msg) (resp *dns.Msg, err error) {
|
||||
|
|
@ -916,6 +938,7 @@ func TestBlockCNAMEProtectionEnabled(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
testUpstm := &aghtest.Upstream{
|
||||
CName: testCNAMEs,
|
||||
|
|
@ -957,6 +980,7 @@ func TestBlockCNAME(t *testing.T) {
|
|||
t,
|
||||
&filtering.Config{ProtectionEnabled: true, BlockingMode: filtering.BlockingModeDefault},
|
||||
forwardConf,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{
|
||||
&aghtest.Upstream{
|
||||
|
|
@ -1033,6 +1057,7 @@ func TestClientRulesForCNAMEMatching(t *testing.T) {
|
|||
t,
|
||||
&filtering.Config{BlockingMode: filtering.BlockingModeDefault},
|
||||
forwardConf,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{
|
||||
&aghtest.Upstream{
|
||||
|
|
@ -1083,6 +1108,7 @@ func TestNullBlockedRequest(t *testing.T) {
|
|||
t,
|
||||
&filtering.Config{ProtectionEnabled: true, BlockingMode: filtering.BlockingModeNullIP},
|
||||
forwardConf,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
startDeferStop(t, s)
|
||||
addr := s.dnsProxy.Addr(proxy.ProtoUDP)
|
||||
|
|
@ -1144,10 +1170,11 @@ func TestBlockedCustomIP(t *testing.T) {
|
|||
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
|
||||
}
|
||||
s, err := NewServer(DNSCreateParams{
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -1226,6 +1253,7 @@ func TestBlockedByHosts(t *testing.T) {
|
|||
t,
|
||||
&filtering.Config{ProtectionEnabled: true, BlockingMode: filtering.BlockingModeDefault},
|
||||
forwardConf,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
startDeferStop(t, s)
|
||||
addr := s.dnsProxy.Addr(proxy.ProtoUDP)
|
||||
|
|
@ -1297,7 +1325,7 @@ func TestBlockedBySafeBrowsing(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
}
|
||||
s := createTestServer(t, filterConf, forwardConf)
|
||||
s := createTestServer(t, filterConf, forwardConf, testTLSConfigProvider)
|
||||
startDeferStop(t, s)
|
||||
addr := s.dnsProxy.Addr(proxy.ProtoUDP)
|
||||
|
||||
|
|
@ -1353,10 +1381,11 @@ func TestRewrite(t *testing.T) {
|
|||
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
|
||||
}
|
||||
s, err := NewServer(DNSCreateParams{
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -1490,9 +1519,10 @@ func TestPTRResponseFromDHCPLeases(t *testing.T) {
|
|||
return "myhost"
|
||||
},
|
||||
},
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
LocalDomain: localDomain,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
LocalDomain: localDomain,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -1579,10 +1609,11 @@ func TestPTRResponseFromHosts(t *testing.T) {
|
|||
|
||||
var s *Server
|
||||
s, err = NewServer(DNSCreateParams{
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: flt,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
DHCPServer: dhcp,
|
||||
DNSFilter: flt,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -1872,6 +1903,7 @@ func TestServer_Exchange(t *testing.T) {
|
|||
UsePrivateRDNS: true,
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
ctx := testutil.ContextWithTimeout(t, testTimeout)
|
||||
|
|
@ -1900,6 +1932,7 @@ func TestServer_Exchange(t *testing.T) {
|
|||
LocalPTRResolvers: []string{},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
ctx := testutil.ContextWithTimeout(t, testTimeout)
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
makeQ := func(qtype rules.RRType) (req *dns.Msg) {
|
||||
|
|
|
|||
|
|
@ -41,10 +41,11 @@ func TestServer_filterDNSResponse(t *testing.T) {
|
|||
f.SetEnabled(true)
|
||||
|
||||
s, err := NewServer(DNSCreateParams{
|
||||
DHCPServer: &testDHCP{},
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
DHCPServer: &testDHCP{},
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@ func TestDNSForwardHTTP_handleGetConfig(t *testing.T) {
|
|||
ConfModifier: agh.EmptyConfigModifier{},
|
||||
ServePlainDNS: true,
|
||||
}
|
||||
s := createTestServer(t, filterConf, forwardConf)
|
||||
s := createTestServer(t, filterConf, forwardConf, testTLSConfigProvider)
|
||||
s.sysResolvers = &emptySysResolvers{}
|
||||
|
||||
require.NoError(t, s.Start(testutil.ContextWithTimeout(t, testTimeout)))
|
||||
|
|
@ -178,7 +178,7 @@ func TestDNSForwardHTTP_handleSetConfig(t *testing.T) {
|
|||
ConfModifier: agh.EmptyConfigModifier{},
|
||||
ServePlainDNS: true,
|
||||
}
|
||||
s := createTestServer(t, filterConf, forwardConf)
|
||||
s := createTestServer(t, filterConf, forwardConf, testTLSConfigProvider)
|
||||
s.sysResolvers = &emptySysResolvers{}
|
||||
|
||||
defaultConf := s.conf
|
||||
|
|
@ -405,6 +405,7 @@ func TestServer_HandleTestUpstreamDNS(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
srv.etcHosts = upstream.NewHostsResolver(hc)
|
||||
startDeferStop(t, srv)
|
||||
|
|
|
|||
|
|
@ -236,6 +236,7 @@ func TestServer_middlewareUDP(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
startDeferStop(t, s)
|
||||
|
|
|
|||
|
|
@ -226,7 +226,34 @@ func (s *Server) makeDDRResponse(req *dns.Msg) (resp *dns.Msg) {
|
|||
resp.Answer = append(resp.Answer, ans)
|
||||
}
|
||||
|
||||
if s.hasIPAddrs {
|
||||
s.appendDoTResolvers(req, resp, domainName)
|
||||
|
||||
for _, addr := range s.dnsProxy.QUICListenAddr {
|
||||
values := []dns.SVCBKeyValue{
|
||||
&dns.SVCBAlpn{Alpn: []string{"doq"}},
|
||||
&dns.SVCBPort{Port: uint16(addr.Port)},
|
||||
}
|
||||
|
||||
ans := &dns.SVCB{
|
||||
Hdr: s.hdr(req, dns.TypeSVCB),
|
||||
Priority: 1,
|
||||
Target: domainName,
|
||||
Value: values,
|
||||
}
|
||||
|
||||
resp.Answer = append(resp.Answer, ans)
|
||||
}
|
||||
|
||||
return resp
|
||||
}
|
||||
|
||||
// appendDoTResolvers appends DNS-over-TLS SVCB resolver entries to resp if the
|
||||
// server's TLS certificate contains IP addresses. req and resp must not be
|
||||
// nil.
|
||||
func (s *Server) appendDoTResolvers(req, resp *dns.Msg, domainName string) {
|
||||
hasIPAddrs := s.tlsConfigProvider.HasIPAddrs()
|
||||
|
||||
if hasIPAddrs {
|
||||
// Only add DNS-over-TLS resolvers in case the certificate contains IP
|
||||
// addresses.
|
||||
//
|
||||
|
|
@ -247,24 +274,6 @@ func (s *Server) makeDDRResponse(req *dns.Msg) (resp *dns.Msg) {
|
|||
resp.Answer = append(resp.Answer, ans)
|
||||
}
|
||||
}
|
||||
|
||||
for _, addr := range s.dnsProxy.QUICListenAddr {
|
||||
values := []dns.SVCBKeyValue{
|
||||
&dns.SVCBAlpn{Alpn: []string{"doq"}},
|
||||
&dns.SVCBPort{Port: uint16(addr.Port)},
|
||||
}
|
||||
|
||||
ans := &dns.SVCB{
|
||||
Hdr: s.hdr(req, dns.TypeSVCB),
|
||||
Priority: 1,
|
||||
Target: domainName,
|
||||
Value: values,
|
||||
}
|
||||
|
||||
resp.Answer = append(resp.Answer, ans)
|
||||
}
|
||||
|
||||
return resp
|
||||
}
|
||||
|
||||
// processDHCPHosts respond to A requests if the target hostname is known to
|
||||
|
|
|
|||
|
|
@ -90,6 +90,7 @@ func TestServer_ProcessInitial(t *testing.T) {
|
|||
t,
|
||||
&filtering.Config{BlockingMode: filtering.BlockingModeDefault},
|
||||
c,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
var gotAddr netip.Addr
|
||||
|
|
@ -193,6 +194,7 @@ func TestServer_ProcessFilteringAfterResponse(t *testing.T) {
|
|||
t,
|
||||
&filtering.Config{BlockingMode: filtering.BlockingModeDefault},
|
||||
c,
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
resp := newResp(dns.RcodeSuccess, tc.req, tc.respAns)
|
||||
|
|
@ -324,9 +326,11 @@ func TestServer_ProcessDDRQuery(t *testing.T) {
|
|||
addrsDoH: addrsDoH,
|
||||
}}
|
||||
|
||||
_, certPem, keyPem := createServerTLSConfig(t)
|
||||
cert, err := tls.X509KeyPair(certPem, keyPem)
|
||||
require.NoError(t, err)
|
||||
tlsConf, _, _ := createServerTLSConfig(t)
|
||||
|
||||
tlsConfProvider := &aghtest.TLSConfigProvider{}
|
||||
tlsConfProvider.OnTLSConfig = func() (conf *tls.Config) { return tlsConf }
|
||||
tlsConfProvider.OnHasIPAddrs = func() (ok bool) { return true }
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
|
|
@ -344,17 +348,14 @@ func TestServer_ProcessDDRQuery(t *testing.T) {
|
|||
},
|
||||
TLSConf: &TLSConfig{
|
||||
ServerName: ddrTestDomainName,
|
||||
Cert: &cert,
|
||||
TLSListenAddrs: tc.addrsDoT,
|
||||
HTTPSListenAddrs: tc.addrsDoH,
|
||||
QUICListenAddrs: tc.addrsDoQ,
|
||||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
tlsConfProvider,
|
||||
)
|
||||
// TODO(e.burkov): Generate a certificate actually containing the
|
||||
// IP addresses.
|
||||
s.hasIPAddrs = true
|
||||
|
||||
req := createTestMessageWithType(tc.host, tc.qtype)
|
||||
|
||||
|
|
@ -691,6 +692,7 @@ func TestServer_ProcessUpstream_localPTR(t *testing.T) {
|
|||
LocalPTRResolvers: []string{localUpsAddr},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
ctx := testutil.ContextWithTimeout(t, testTimeout)
|
||||
pctx := newPrxCtx()
|
||||
|
|
@ -721,6 +723,7 @@ func TestServer_ProcessUpstream_localPTR(t *testing.T) {
|
|||
LocalPTRResolvers: []string{localUpsAddr},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
pctx := newPrxCtx()
|
||||
|
||||
|
|
|
|||
|
|
@ -63,9 +63,10 @@ func TestServer_ServeDNS(t *testing.T) {
|
|||
OnHostByIP: func(ip netip.Addr) (_ string) { panic(testutil.UnexpectedCall(ip)) },
|
||||
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
|
||||
},
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
DNSFilter: f,
|
||||
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: testTLSConfigProvider,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
@ -264,6 +265,7 @@ func TestServer_ServeDNS_restrictLocal(t *testing.T) {
|
|||
LocalPTRResolvers: []string{localUpsAddr},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
startDeferStop(t, s)
|
||||
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ func TestGenAnswerHTTPS_andSVCB(t *testing.T) {
|
|||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
testTLSConfigProvider,
|
||||
)
|
||||
|
||||
req := &dns.Msg{
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ package home
|
|||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
|
|
@ -16,6 +15,7 @@ import (
|
|||
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/aghtls"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/client"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/dnsforward"
|
||||
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
|
||||
|
|
@ -146,15 +146,16 @@ func initDNSServer(
|
|||
confModifier agh.ConfigModifier,
|
||||
) (err error) {
|
||||
globalContext.dnsServer, err = dnsforward.NewServer(dnsforward.DNSCreateParams{
|
||||
Logger: l,
|
||||
DNSFilter: filters,
|
||||
Stats: sts,
|
||||
QueryLog: qlog,
|
||||
PrivateNets: parseSubnetSet(config.DNS.PrivateNets),
|
||||
Anonymizer: anonymizer,
|
||||
DHCPServer: dhcpSrv,
|
||||
EtcHosts: globalContext.etcHosts,
|
||||
LocalDomain: config.DHCP.LocalDomainName,
|
||||
Logger: l,
|
||||
DNSFilter: filters,
|
||||
Stats: sts,
|
||||
QueryLog: qlog,
|
||||
PrivateNets: parseSubnetSet(config.DNS.PrivateNets),
|
||||
Anonymizer: anonymizer,
|
||||
DHCPServer: dhcpSrv,
|
||||
EtcHosts: globalContext.etcHosts,
|
||||
LocalDomain: config.DHCP.LocalDomainName,
|
||||
TLSConfigProvider: tlsMgr,
|
||||
})
|
||||
defer func() {
|
||||
if err != nil {
|
||||
|
|
@ -265,7 +266,7 @@ func newServerConfig(
|
|||
clientSrcConf *clientSourcesConfig,
|
||||
extTLSConf *tlsConfigSettings,
|
||||
dohConf *doHConfig,
|
||||
tlsMgr *tlsManager,
|
||||
tlsConfProvider aghtls.TLSConfigProvider,
|
||||
httpReg aghhttp.Registrar,
|
||||
clientsContainer dnsforward.ClientsContainer,
|
||||
confModifier agh.ConfigModifier,
|
||||
|
|
@ -275,7 +276,7 @@ func newServerConfig(
|
|||
fwdConf := dnsConf.Config
|
||||
fwdConf.ClientsContainer = clientsContainer
|
||||
|
||||
intTLSConf, err := newDNSTLSConfig(extTLSConf, hosts, dohConf.InsecureEnabled)
|
||||
intTLSConf, err := newDNSTLSConfig(extTLSConf, hosts)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("constructing tls config: %w", err)
|
||||
}
|
||||
|
|
@ -287,7 +288,7 @@ func newServerConfig(
|
|||
TLSConf: intTLSConf,
|
||||
TLSAllowUnencryptedDoH: dohConf.InsecureEnabled,
|
||||
UpstreamTimeout: time.Duration(dnsConf.UpstreamTimeout),
|
||||
TLSv12Roots: tlsMgr.rootCerts,
|
||||
TLSv12Roots: tlsConfProvider.RootCAs(),
|
||||
ConfModifier: confModifier,
|
||||
HTTPReg: httpReg,
|
||||
LocalPTRResolvers: dnsConf.PrivateRDNSResolvers,
|
||||
|
|
@ -327,7 +328,6 @@ func newServerConfig(
|
|||
func newDNSTLSConfig(
|
||||
extTLSConf *tlsConfigSettings,
|
||||
addrs []netip.Addr,
|
||||
allowUnencryptedDoH bool,
|
||||
) (dnsConf *dnsforward.TLSConfig, err error) {
|
||||
if !extTLSConf.Enabled {
|
||||
return &dnsforward.TLSConfig{}, nil
|
||||
|
|
@ -359,22 +359,6 @@ func newDNSTLSConfig(
|
|||
dnsConf.QUICListenAddrs = ipsToUDPAddrs(addrs, extTLSConf.PortDNSOverQUIC)
|
||||
}
|
||||
|
||||
cert, err := tls.X509KeyPair(extTLSConf.CertificateChainData, extTLSConf.PrivateKeyData)
|
||||
if err != nil {
|
||||
err = fmt.Errorf("parsing tls key pair: %w", err)
|
||||
if allowUnencryptedDoH || dnsCryptConf != nil {
|
||||
// TODO(s.chzhen): Use [slog.Logger].
|
||||
log.Info("warning: %s", err)
|
||||
|
||||
return dnsConf, nil
|
||||
}
|
||||
|
||||
// Don't wrap the error, because it's already annotated.
|
||||
return nil, err
|
||||
}
|
||||
|
||||
dnsConf.Cert = &cert
|
||||
|
||||
return dnsConf, nil
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -61,6 +61,8 @@ func httpClient(tlsMgr *tlsManager) (c *http.Client) {
|
|||
tr := newCustomUserAgentTransport(&http.Transport{
|
||||
DialContext: dialContext,
|
||||
Proxy: httpProxy,
|
||||
// TODO(m.kazantsev): Do not create TLS config manually, but use
|
||||
// [aghtls.TLSConfigProvider].
|
||||
TLSClientConfig: &tls.Config{
|
||||
RootCAs: tlsMgr.rootCerts,
|
||||
CipherSuites: tlsMgr.customCipherIDs,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import (
|
|||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
|
@ -34,12 +35,25 @@ type tlsManager struct {
|
|||
// logger is used for logging the operation of the TLS Manager.
|
||||
logger *slog.Logger
|
||||
|
||||
// mu protects certLastMod, extTLSConf.
|
||||
// mu protects certLastMod, tlsCert, tlsConf, extTLSConf.
|
||||
mu *sync.Mutex
|
||||
|
||||
// certLastMod is the last modification time of the certificate file.
|
||||
certLastMod time.Time
|
||||
|
||||
// tlsCert is the current TLS certificate. tlsCert must not be stored in
|
||||
// [tls.Config.Certificates], as it violates its documentation.
|
||||
//
|
||||
// TODO(m.kazantsev): Consider a better approach to store the certificate.
|
||||
tlsCert *tls.Certificate
|
||||
|
||||
// tlsConf is a current TLS configuration. It may be nil.
|
||||
tlsConf *tls.Config
|
||||
|
||||
// extTLSConf contains extended TLS configuration settings. It must not be
|
||||
// nil.
|
||||
extTLSConf *tlsConfigSettings
|
||||
|
||||
// rootCerts is a pool of root CAs for TLSv1.2.
|
||||
rootCerts *x509.CertPool
|
||||
|
||||
|
|
@ -49,13 +63,6 @@ type tlsManager struct {
|
|||
// Resolve it.
|
||||
web *webAPI
|
||||
|
||||
// extTLSConf contains extended TLS configuration settings. It must not be
|
||||
// nil.
|
||||
// TODO(m.kazantsev): Add a field of a type of [*tls.Config] which will
|
||||
// represent the TLS settings. This is why these settings are called
|
||||
// 'extended'.
|
||||
extTLSConf *tlsConfigSettings
|
||||
|
||||
// confModifier is used to update the global configuration.
|
||||
confModifier agh.ConfigModifier
|
||||
|
||||
|
|
@ -115,14 +122,14 @@ func newTLSManager(ctx context.Context, conf *tlsManagerConfig) (m *tlsManager,
|
|||
m.extTLSConf.Status = tlsConfigStatus{}
|
||||
|
||||
if len(conf.tlsSettings.OverrideTLSCiphers) > 0 {
|
||||
m.customCipherIDs, err = aghtls.ParseCiphers(config.TLS.OverrideTLSCiphers)
|
||||
m.customCipherIDs, err = aghtls.ParseCiphers(conf.tlsSettings.OverrideTLSCiphers)
|
||||
if err != nil {
|
||||
// Should not happen because upstreams are already validated. See
|
||||
// [validateTLSCipherIDs].
|
||||
panic(err)
|
||||
}
|
||||
|
||||
m.logger.InfoContext(ctx, "overriding ciphers", "ciphers", config.TLS.OverrideTLSCiphers)
|
||||
m.logger.InfoContext(ctx, "overriding ciphers", "ciphers", conf.tlsSettings.OverrideTLSCiphers)
|
||||
} else {
|
||||
m.logger.InfoContext(ctx, "using default ciphers")
|
||||
}
|
||||
|
|
@ -146,14 +153,58 @@ func newTLSManager(ctx context.Context, conf *tlsManagerConfig) (m *tlsManager,
|
|||
if err != nil {
|
||||
m.extTLSConf.Enabled = false
|
||||
|
||||
// Don't wrap the error, because it's informative enough as is.
|
||||
return m, err
|
||||
}
|
||||
|
||||
cert, err := tls.X509KeyPair(m.extTLSConf.CertificateChainData, m.extTLSConf.PrivateKeyData)
|
||||
if err != nil {
|
||||
m.extTLSConf.Enabled = false
|
||||
|
||||
return m, fmt.Errorf("parsing tls certificate: %w", err)
|
||||
}
|
||||
|
||||
slices.Sort(cert.Leaf.DNSNames)
|
||||
|
||||
m.tlsConf = &tls.Config{
|
||||
RootCAs: m.rootCerts,
|
||||
CipherSuites: m.customCipherIDs,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
GetCertificate: m.onGetCertificate,
|
||||
}
|
||||
|
||||
m.tlsCert = &cert
|
||||
m.setCertFileTime(ctx)
|
||||
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// checkIfValidStatus checks if status is valid. If it is valid, certErr is set
|
||||
// to nil. Otherwise, certErr is returned as is. status must not be nil.
|
||||
func (m *tlsManager) checkIfValidStatus(
|
||||
ctx context.Context,
|
||||
status *tlsConfigStatus,
|
||||
certErr error,
|
||||
) (err error) {
|
||||
if certErr == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
status.WarningValidation = certErr.Error()
|
||||
if status.ValidCert && status.ValidKey && status.ValidPair {
|
||||
// Do not return warnings since those aren't critical, just log.
|
||||
m.logger.WarnContext(
|
||||
ctx,
|
||||
"error while loading tls configuration",
|
||||
slogutil.KeyError, certErr,
|
||||
)
|
||||
|
||||
certErr = nil
|
||||
}
|
||||
|
||||
return certErr
|
||||
}
|
||||
|
||||
// setWebAPI stores the provided web API. It must be called before
|
||||
// [tlsManager.start], [tlsManager.reload], [webAPI.handleTLSConfigure], or
|
||||
// [webAPI.validateTLSSettings].
|
||||
|
|
@ -163,7 +214,8 @@ func (m *tlsManager) setWebAPI(webAPI *webAPI) {
|
|||
m.web = webAPI
|
||||
}
|
||||
|
||||
// extendedTLSConfig returns a deep copy of the stored TLS configuration.
|
||||
// extendedTLSConfig returns a deep copy of the stored extended TLS
|
||||
// configuration.
|
||||
func (m *tlsManager) extendedTLSConfig() (extTLSConf *tlsConfigSettings) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
|
@ -174,7 +226,7 @@ func (m *tlsManager) extendedTLSConfig() (extTLSConf *tlsConfigSettings) {
|
|||
// setCertFileTime sets [tlsManager.certLastMod] from the certificate. If there
|
||||
// are errors, setCertFileTime logs them. m.mu is expected to be locked.
|
||||
func (m *tlsManager) setCertFileTime(ctx context.Context) {
|
||||
if len(m.extTLSConf.CertificatePath) == 0 {
|
||||
if m.extTLSConf.CertificatePath == "" {
|
||||
return
|
||||
}
|
||||
|
||||
|
|
@ -252,26 +304,28 @@ func (m *tlsManager) reload(ctx context.Context) {
|
|||
|
||||
m.logger.InfoContext(ctx, "certificate file is modified")
|
||||
|
||||
tlsConf := *tlsConfPtr
|
||||
extTLSConf := *tlsConfPtr
|
||||
status := &tlsConfigStatus{}
|
||||
|
||||
err = m.loadTLSConfig(ctx, &tlsConf, status)
|
||||
err = m.loadTLSConfig(ctx, &extTLSConf, status)
|
||||
if err != nil {
|
||||
m.logger.WarnContext(ctx, "reloading interrupted", slogutil.KeyError, err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
tlsConf.Status = *status
|
||||
|
||||
m.extTLSConf = &tlsConf
|
||||
m.certLastMod = fi.ModTime().UTC()
|
||||
|
||||
err = m.web.reconfigureDNSServer(ctx, m.extTLSConf)
|
||||
err = m.updateTLSCert(&extTLSConf)
|
||||
if err != nil {
|
||||
m.logger.ErrorContext(ctx, "reconfiguring dns server", slogutil.KeyError, err)
|
||||
m.logger.WarnContext(ctx, "failed to update tls certificate", slogutil.KeyError, err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
extTLSConf.Status = *status
|
||||
|
||||
m.extTLSConf = &extTLSConf
|
||||
m.certLastMod = fi.ModTime().UTC()
|
||||
|
||||
// The background context is used because the TLSConfigChanged wraps context
|
||||
// with timeout on its own and shuts down the server, which handles current
|
||||
// request.
|
||||
|
|
@ -281,20 +335,15 @@ func (m *tlsManager) reload(ctx context.Context) {
|
|||
// loadTLSConfig loads and validates the TLS configuration. It also sets
|
||||
// [tlsConfigSettings.CertificateChainData] and
|
||||
// [tlsConfigSettings.PrivateKeyData] properties. The returned error is also
|
||||
// set in status.WarningValidation.
|
||||
// set in status.WarningValidation. All arguments must not be nil. m.mu is
|
||||
// expected to be locked.
|
||||
func (m *tlsManager) loadTLSConfig(
|
||||
ctx context.Context,
|
||||
extTLSConf *tlsConfigSettings,
|
||||
status *tlsConfigStatus,
|
||||
) (err error) {
|
||||
defer func() {
|
||||
if err != nil {
|
||||
status.WarningValidation = err.Error()
|
||||
if status.ValidCert && status.ValidKey && status.ValidPair {
|
||||
// Do not return warnings since those aren't critical.
|
||||
err = nil
|
||||
}
|
||||
}
|
||||
err = m.checkIfValidStatus(ctx, status, err)
|
||||
}()
|
||||
|
||||
err = loadCertificateChainData(extTLSConf)
|
||||
|
|
@ -429,10 +478,18 @@ func (m *tlsManager) setConfig(
|
|||
ctx context.Context,
|
||||
newConf *tlsConfigSettings,
|
||||
servePlain aghalg.NullBool,
|
||||
) (restartHTTPS bool) {
|
||||
) (restartHTTPS bool, err error) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
err = m.updateTLSCert(newConf)
|
||||
if err != nil {
|
||||
m.logger.ErrorContext(ctx, "updating tls certificate", slogutil.KeyError, err)
|
||||
|
||||
// Don't wrap the error, because it is informative enough as is.
|
||||
return false, err
|
||||
}
|
||||
|
||||
m.extTLSConf.updatePlainDNS(newConf, servePlain)
|
||||
|
||||
if !m.extTLSConf.setPrivateFieldsAndCompare(newConf) {
|
||||
|
|
@ -449,7 +506,7 @@ func (m *tlsManager) setConfig(
|
|||
certPath, keyPath = newConf.CertificatePath, newConf.PrivateKeyPath
|
||||
}
|
||||
|
||||
err := m.manager.Set(ctx, aghtls.TLSPair{
|
||||
err = m.manager.Set(ctx, aghtls.TLSPair{
|
||||
CertPath: certPath,
|
||||
KeyPath: keyPath,
|
||||
})
|
||||
|
|
@ -459,7 +516,7 @@ func (m *tlsManager) setConfig(
|
|||
|
||||
m.setCertFileTime(ctx)
|
||||
|
||||
return restartHTTPS
|
||||
return restartHTTPS, nil
|
||||
}
|
||||
|
||||
// updatePlainDNS checks the old value of [tlsConfigSettings.ServePlainDNS] in
|
||||
|
|
@ -804,3 +861,79 @@ func (m *tlsManager) marshalTLS(
|
|||
|
||||
aghhttp.WriteJSONResponseOK(ctx, m.logger, w, r, *data)
|
||||
}
|
||||
|
||||
// TLSConfig implements the [aghtls.TLSConfigProvider] interface for
|
||||
// *tlsManager.
|
||||
func (m *tlsManager) TLSConfig() (conf *tls.Config) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
return m.tlsConf.Clone()
|
||||
}
|
||||
|
||||
// RootCAs implements the [aghtls.TLSConfigProvider] interface for *tlsManager.
|
||||
func (m *tlsManager) RootCAs() (root *x509.CertPool) {
|
||||
return m.rootCerts
|
||||
}
|
||||
|
||||
// HasIPAddrs implements the [aghtls.TLSConfigProvider] interface for
|
||||
// *tlsManager. It returns true if the current TLS configuration has at least
|
||||
// one certificate with an IP address in its SAN extension.
|
||||
func (m *tlsManager) HasIPAddrs() (ok bool) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if m.tlsCert == nil || m.tlsCert.Leaf == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
// TODO(m.kazantsev): Consider storing the value instead of parsing each
|
||||
// time.
|
||||
return aghtls.CertificateHasIP(m.tlsCert.Leaf)
|
||||
}
|
||||
|
||||
// onGetCertificate gets [*tls.Certificate] from [*tls.Config]. If
|
||||
// [tlsManager.extTLSConf.Enabled] is false, nil is returned.
|
||||
//
|
||||
// TODO(m.kazantsev): Consider using tls.SupportsCertificate.
|
||||
func (m *tlsManager) onGetCertificate(chi *tls.ClientHelloInfo) (cert *tls.Certificate, err error) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if !m.extTLSConf.Enabled || m.tlsConf == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
tlsCert := *m.tlsCert
|
||||
|
||||
return &tlsCert, nil
|
||||
}
|
||||
|
||||
// updateTLSCert loads and updates a TLS certificate for m.tlsConf. If
|
||||
// m.tlsConf is nil, it will be initialized. extTLSConf must not be nil. m.mu
|
||||
// must be locked.
|
||||
func (m *tlsManager) updateTLSCert(extTLSConf *tlsConfigSettings) (err error) {
|
||||
if len(extTLSConf.CertificateChainData) == 0 || len(extTLSConf.PrivateKeyData) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
cert, err := tls.X509KeyPair(extTLSConf.CertificateChainData, extTLSConf.PrivateKeyData)
|
||||
if err != nil {
|
||||
return fmt.Errorf("loading tls certificate: %w", err)
|
||||
}
|
||||
|
||||
slices.Sort(cert.Leaf.DNSNames)
|
||||
|
||||
if m.tlsConf == nil {
|
||||
m.tlsConf = &tls.Config{
|
||||
RootCAs: m.rootCerts,
|
||||
CipherSuites: m.customCipherIDs,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
GetCertificate: m.onGetCertificate,
|
||||
}
|
||||
}
|
||||
|
||||
m.tlsCert = &cert
|
||||
|
||||
return nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -229,6 +229,7 @@ func assertCertSerialNumber(tb testing.TB, conf *tlsConfigSettings, wantSN int64
|
|||
assert.Equal(tb, wantSN, cert.Leaf.SerialNumber.Int64())
|
||||
}
|
||||
|
||||
// TODO(m.kazantsev): Refactor.
|
||||
func TestTLSManager_Reload(t *testing.T) {
|
||||
storeGlobals(t)
|
||||
|
||||
|
|
@ -240,10 +241,25 @@ func TestTLSManager_Reload(t *testing.T) {
|
|||
)
|
||||
|
||||
globalContext.dnsServer, err = dnsforward.NewServer(dnsforward.DNSCreateParams{
|
||||
Logger: testLogger,
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: aghtls.EmptyTLSConfigProvider{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = globalContext.dnsServer.Prepare(
|
||||
testutil.ContextWithTimeout(t, testTimeout),
|
||||
&dnsforward.ServerConfig{
|
||||
TLSConf: &dnsforward.TLSConfig{},
|
||||
Config: dnsforward.Config{
|
||||
UpstreamMode: dnsforward.UpstreamModeLoadBalance,
|
||||
EDNSClientSubnet: &dnsforward.EDNSClientSubnet{Enabled: false},
|
||||
ClientsContainer: dnsforward.EmptyClientsContainer{},
|
||||
},
|
||||
ServePlainDNS: true,
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
globalContext.clients.storage, err = client.NewStorage(ctx, &client.StorageConfig{
|
||||
BaseLogger: testLogger,
|
||||
Logger: testLogger,
|
||||
|
|
|
|||
|
|
@ -465,6 +465,8 @@ func (web *webAPI) serveTLS(ctx context.Context) (next bool) {
|
|||
web.httpsServer.server = &http.Server{
|
||||
Addr: addr,
|
||||
Handler: hdlr,
|
||||
// TODO(m.kazantsev): Do not create TLS config manually, but use
|
||||
// [aghtls.TLSConfigProvider].
|
||||
TLSConfig: &tls.Config{
|
||||
Certificates: []tls.Certificate{web.httpsServer.certificate()},
|
||||
RootCAs: web.tlsManager.rootCerts,
|
||||
|
|
@ -505,6 +507,8 @@ func (web *webAPI) mustStartHTTP3(ctx context.Context, address string) {
|
|||
// TODO(a.garipov): See if there is a way to use the error log as
|
||||
// well as timeouts here.
|
||||
Addr: address,
|
||||
// TODO(m.kazantsev): Do not create TLS config manually, but use
|
||||
// [aghtls.TLSConfigProvider].
|
||||
TLSConfig: &tls.Config{
|
||||
Certificates: []tls.Certificate{web.httpsServer.certificate()},
|
||||
RootCAs: web.tlsManager.rootCerts,
|
||||
|
|
@ -809,7 +813,12 @@ func (web *webAPI) handleTLSConfigure(w http.ResponseWriter, r *http.Request) {
|
|||
newTLSConf := &req.tlsConfigSettings
|
||||
newTLSConf.Status = *status
|
||||
|
||||
restartHTTPS = web.tlsManager.setConfig(ctx, newTLSConf, req.ServePlainDNS)
|
||||
restartHTTPS, err = web.tlsManager.setConfig(ctx, newTLSConf, req.ServePlainDNS)
|
||||
if err != nil {
|
||||
aghhttp.ErrorAndLog(ctx, web.logger, r, w, http.StatusInternalServerError, "%s", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
err = web.reconfigureDNSServer(ctx, newTLSConf)
|
||||
if err != nil {
|
||||
|
|
|
|||
|
|
@ -36,7 +36,8 @@ func TestWebAPI_HandleTLSConfigure(t *testing.T) {
|
|||
)
|
||||
|
||||
globalContext.dnsServer, err = dnsforward.NewServer(dnsforward.DNSCreateParams{
|
||||
Logger: testLogger,
|
||||
Logger: testLogger,
|
||||
TLSConfigProvider: aghtls.EmptyTLSConfigProvider{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue