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

Squashed commit of the following:

commit e4d7f8a690f7e0731b86f41c05b42f816a4f4c38
Merge: d4c296f79 336e9c9df
Author: 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;

commit ba44e52537
Author: Maksim Kazantsev <m.kazantsev@adguard.com>
Date:   Fri Jul 24 18:50:02 2026 +0300

    home: rm supports certificate;

commit 4ebec73863
Author: Maksim Kazantsev <m.kazantsev@adguard.com>
Date:   Wed Jul 22 13:04:17 2026 +0300

    internal: imp code; imp tests;

commit 715baa5bfa
Author: Maksim Kazantsev <m.kazantsev@adguard.com>
Date:   Tue Jul 21 17:24:22 2026 +0300

    internal: upd interface; do not check certificates any more;

commit 92592a017d
Author: Maksim Kazantsev <m.kazantsev@adguard.com>
Date:   Mon Jul 20 14:25:16 2026 +0300

    dnsforward: dont ignore error from origGetCert;

commit afc6b1ba76
Author: Maksim Kazantsev <m.kazantsev@adguard.com>
Date:   Mon Jul 20 13:46:32 2026 +0300

    dnsforward: add common name fallback; home: imp onGetCertificate;

commit 30b71c50ce
Author: 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;

commit 3b55b20cb2
Author: 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:
Maksim Kazantsev 2026-07-27 14:42:13 +00:00
parent 336e9c9df4
commit 254e17dabf
21 changed files with 435 additions and 224 deletions

View file

@ -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:

View file

@ -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()
}

View file

@ -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
}

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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)

View file

@ -49,6 +49,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
},
ServePlainDNS: true,
},
testTLSConfigProvider,
)
makeQ := func(qtype rules.RRType) (req *dns.Msg) {

View file

@ -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)

View file

@ -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)

View file

@ -236,6 +236,7 @@ func TestServer_middlewareUDP(t *testing.T) {
},
ServePlainDNS: true,
},
testTLSConfigProvider,
)
startDeferStop(t, s)

View file

@ -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

View file

@ -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()

View file

@ -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)

View file

@ -28,6 +28,7 @@ func TestGenAnswerHTTPS_andSVCB(t *testing.T) {
},
ServePlainDNS: true,
},
testTLSConfigProvider,
)
req := &dns.Msg{

View file

@ -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
}

View file

@ -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,

View file

@ -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
}

View file

@ -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,

View file

@ -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 {

View file

@ -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)