Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 31 additions & 2 deletions proxy/upstreams.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ import (
"github.com/AdguardTeam/golibs/container"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/netutil"
"golang.org/x/net/idna"
)

// UnqualifiedNames is a key for [UpstreamConfig.DomainReservedUpstreams] map to
Expand All @@ -22,6 +23,14 @@ const UnqualifiedNames = "unqualified_names"
// labelSep is a separator between labels of a domain name.
const labelSep = "."

// domainSpecIDNA converts domains to their canonical lookup form while
// preserving the permissive domain-name validation used by this parser.
var domainSpecIDNA = idna.New(
idna.MapForLookup(),
idna.StrictDomainName(false),
idna.ValidateLabels(false),
)

// UpstreamConfig maps domain names to upstreams.
type UpstreamConfig struct {
// DomainReservedUpstreams maps the domains to the upstreams.
Expand Down Expand Up @@ -221,6 +230,26 @@ func (p *configParser) parseLine(idx int, confLine string) (err error) {
return nil
}

// normalizeDomainSpec validates domain and returns its canonical ASCII form.
func normalizeDomainSpec(domain string) (normalized string, err error) {
wildcard := strings.HasPrefix(domain, "*.")
domain = strings.TrimPrefix(domain, "*.")
err = netutil.ValidateDomainName(domain)
if err != nil {
return "", err
}

domain, err = domainSpecIDNA.ToASCII(domain)
if err != nil {
return "", fmt.Errorf("converting to ASCII: %w", err)
}
if wildcard {
domain = "*." + domain
}

return strings.ToLower(domain + labelSep), nil
}

// splitConfigLine parses upstream configuration line and returns list upstream
// addresses (one or many), list of domains for which this upstream is reserved
// (may be nil). It returns an error if the upstream format is incorrect.
Expand All @@ -243,12 +272,12 @@ func splitConfigLine(confLine string) (upstreams, domains []string, err error) {
continue
}

err = netutil.ValidateDomainName(strings.TrimPrefix(confHost, "*."))
confHost, err = normalizeDomainSpec(confHost)
if err != nil {
return nil, nil, fmt.Errorf("domain at index %d: %w", i, err)
}

domains = append(domains, strings.ToLower(confHost+labelSep))
domains = append(domains, confHost)
}

return strings.Fields(upstreamsLine), domains, nil
Expand Down
20 changes: 20 additions & 0 deletions proxy/upstreams_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,26 @@ func TestUpstreamConfig_GetUpstreamsForDomain(t *testing.T) {
}
}

func TestUpstreamConfig_GetUpstreamsForDomain_IDNA(t *testing.T) {
t.Parallel()

const idnaUpstream = "tcp://idna.upstream:53"

config, err := ParseUpstreamsConfig([]string{
generalUpstream,
"[/恒天.com/]" + idnaUpstream,
"[/GÖPHER.com/]" + idnaUpstream,
}, nil)
require.NoError(t, err)
testutil.CleanupAndRequireSuccess(t, config.Close)

ups := config.getUpstreamsForDomain("xn--rss99n.com.")
assertUpstreamsAddrs(t, ups, []string{idnaUpstream})

ups = config.getUpstreamsForDomain("xn--gpher-jua.com.")
assertUpstreamsAddrs(t, ups, []string{idnaUpstream})
}

func TestUpstreamConfig_GetUpstreamsForDS(t *testing.T) {
t.Parallel()

Expand Down