diff --git a/docs/cli-reference.md b/docs/cli-reference.md index 9be3293c..081067d3 100644 --- a/docs/cli-reference.md +++ b/docs/cli-reference.md @@ -535,6 +535,18 @@ fetch --proxy http://localhost:8080 example.com fetch --proxy socks5://localhost:1080 example.com ``` +### `--resolve [+]HOST:PORT:IP[,IP]` + +Connect to the supplied IP address for a matching host and port while keeping +the URL host as the HTTP Host header and TLS SNI. Repeat the option or provide +comma-separated addresses to provide multiple candidates. A `*` host acts as +a fallback for any host. A leading `+` is accepted for curl compatibility. + +```sh +fetch --resolve example.com:443:127.0.0.1 https://example.com +fetch --resolve '*:443:192.0.2.10' https://example.com +``` + ### `--unix PATH` Make request over a Unix domain socket. Unix-like systems only. @@ -804,7 +816,7 @@ fetch --from-curl 'https://example.com' | Auth | `-u`, `--digest`, `--aws-sigv4`, `--oauth2-bearer` | | TLS | `-k`, `--cacert`, `-E`/`--cert`, `--key`, `--tlsv1.x`, `--tls-max`, `--ech hard | true | auto | false` | | Output | `-o`, `-O`, `-J` | -| Network | `-L`, `--max-redirs`, `-m`/`--max-time`, `--connect-timeout`, `-x`, `--unix-socket`, `--doh-url`, `--retry`, `--retry-delay`, `--retry-unsafe`, `-r` | +| Network | `-L`, `--max-redirs`, `-m`/`--max-time`, `--connect-timeout`, `-x`, `--unix-socket`, `--doh-url`, `--resolve`, `--retry`, `--retry-delay`, `--retry-unsafe`, `-r` | | HTTP version | `-0`, `--http1.1`, `--http2`, `--http3` | | Headers | `-A`, `-e`, `-b` | | Verbosity | `-v`, `-s` | diff --git a/integration/integration_test.go b/integration/integration_test.go index 946c0e52..f9d6cb0a 100644 --- a/integration/integration_test.go +++ b/integration/integration_test.go @@ -545,6 +545,64 @@ func TestMain(t *testing.T) { } }) + t.Run("resolve connects to the selected IP and preserves the URL host", func(t *testing.T) { + t.Parallel() + chHost := make(chan string, 1) + server := startServer(func(w http.ResponseWriter, r *http.Request) { + chHost <- r.Host + io.WriteString(w, "resolved") + }) + defer server.Close() + + _, port, err := net.SplitHostPort(strings.TrimPrefix(server.URL, "http://")) + if err != nil { + t.Fatal(err) + } + host := "resolve.example.test" + target := "http://" + host + ":" + port + res := runFetchOpts(t, fetchPath, fetchOpts{env: []string{ + "HTTP_PROXY=", "http_proxy=", "HTTPS_PROXY=", "https_proxy=", "ALL_PROXY=", "all_proxy=", + "NO_PROXY=*", "no_proxy=*", + }}, target, "--resolve", host+":"+port+":127.0.0.1") + assertExitCode(t, 0, res) + assertBufEquals(t, res.stdout, "resolved") + if got := <-chHost; got != host+":"+port { + t.Fatalf("request Host = %q, want %q", got, host+":"+port) + } + }) + + t.Run("resolve preserves HTTPS Host and SNI", func(t *testing.T) { + t.Parallel() + info := make(chan [2]string, 1) + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + sni := "" + if r.TLS != nil { + sni = r.TLS.ServerName + } + info <- [2]string{r.Host, sni} + io.WriteString(w, "secure-resolved") + })) + defer server.Close() + + _, port, err := net.SplitHostPort(strings.TrimPrefix(server.URL, "https://")) + if err != nil { + t.Fatal(err) + } + host := "tls-resolve.example.test" + target := "https://" + host + ":" + port + res := runFetchOpts(t, fetchPath, fetchOpts{env: []string{ + "HTTP_PROXY=", "http_proxy=", "HTTPS_PROXY=", "https_proxy=", "ALL_PROXY=", "all_proxy=", + "NO_PROXY=*", "no_proxy=*", + }}, target, "--insecure", "--resolve", host+":"+port+":127.0.0.1") + assertExitCode(t, 0, res) + assertBufEquals(t, res.stdout, "secure-resolved") + got := <-info + wantHost := host + ":" + port + if got[0] != wantHost || got[1] != host { + t.Fatalf("request Host/SNI = %q/%q, want %q/%q", got[0], got[1], wantHost, host) + } + }) + t.Run("dns over https", func(t *testing.T) { t.Parallel() server := startServer(func(w http.ResponseWriter, r *http.Request) { diff --git a/internal/cli/app.go b/internal/cli/app.go index 728f6baf..5472db49 100644 --- a/internal/cli/app.go +++ b/internal/cli/app.go @@ -11,6 +11,7 @@ import ( "github.com/ryanfowler/fetch/internal/aws" "github.com/ryanfowler/fetch/internal/config" "github.com/ryanfowler/fetch/internal/core" + "github.com/ryanfowler/fetch/internal/resolver" ) // App represents the full configuration for a fetch invocation. @@ -60,6 +61,7 @@ type App struct { Range []string RemoteHeaderName bool RemoteName bool + Resolve []resolver.ResolveEntry UnixSocket string Update bool CheckUpdate bool @@ -489,6 +491,14 @@ func (a *App) CLI() *CLI { boolFlag(&a.RemoteName, "remote-name", "O", "Use URL path component as output filename"). WithAliases("output-current-dir"), + { + Long: "resolve", + Args: "HOST:PORT:IP", + Description: "Connect to IP while preserving Host/SNI", + IsSet: func() bool { return len(a.Resolve) > 0 }, + Fn: a.parseResolveFlag, + }, + cfgFlag("retry", "", "NUM", "Maximum number of retries", func() bool { return a.Cfg.Retry != nil }, a.Cfg.ParseRetry). WithDefault("0"), @@ -635,6 +645,15 @@ func (a *App) parseBasicFlag(value string) error { return nil } +func (a *App) parseResolveFlag(value string) error { + entries, err := resolver.ParseResolveEntries(value) + if err != nil { + return core.NewValueError("resolve", value, err.Error(), false) + } + a.Resolve = append(a.Resolve, entries...) + return nil +} + func (a *App) parseDigestFlag(value string) error { user, pass, ok := strings.Cut(value, ":") if !ok { diff --git a/internal/cli/cli.go b/internal/cli/cli.go index 665d11aa..462371ec 100644 --- a/internal/cli/cli.go +++ b/internal/cli/cli.go @@ -772,6 +772,11 @@ func (a *App) applyFromCurl(r *curl.Result) error { return err } } + for _, value := range r.Resolve { + if err := a.parseResolveFlag(value); err != nil { + return err + } + } if r.RetrySet { if err := a.Cfg.ParseRetry(strconv.Itoa(r.Retry)); err != nil { return err diff --git a/internal/cli/cli_test.go b/internal/cli/cli_test.go index f626ec41..6de45cb4 100644 --- a/internal/cli/cli_test.go +++ b/internal/cli/cli_test.go @@ -194,6 +194,42 @@ func TestCLI002Validation(t *testing.T) { } } +func TestResolveFlag(t *testing.T) { + app, err := Parse([]string{ + "--resolve", "Example.com:443:192.0.2.10", + "--resolve=*:80:192.0.2.11,192.0.2.12", + "https://example.com", + }) + if err != nil { + t.Fatal(err) + } + if len(app.Resolve) != 3 || app.Resolve[0].Host != "example.com" || app.Resolve[0].Port != "443" || app.Resolve[0].IP.String() != "192.0.2.10" || app.Resolve[1].Host != "*" || app.Resolve[2].IP.String() != "192.0.2.12" { + t.Fatalf("Resolve = %+v, want three parsed mappings", app.Resolve) + } + if !app.OptionProvenance("resolve").Has(SourceCLI) { + t.Fatal("resolve did not record explicit CLI provenance") + } + + for _, value := range []string{"example.com:443", "example.com:443:not-an-ip"} { + if _, err := Parse([]string{"--resolve", value, "https://example.com"}); err == nil || !strings.Contains(err.Error(), "resolve") { + t.Fatalf("Parse(--resolve %q) error = %v, want resolve validation error", value, err) + } + } +} + +func TestFromCurlResolve(t *testing.T) { + app, err := Parse([]string{"--from-curl", "curl --resolve +example.com:443:192.0.2.10,192.0.2.11 https://example.com"}) + if err != nil { + t.Fatal(err) + } + if len(app.Resolve) != 2 || app.Resolve[0].Host != "example.com" || app.Resolve[0].IP.String() != "192.0.2.10" || app.Resolve[1].IP.String() != "192.0.2.11" { + t.Fatalf("Resolve = %+v, want imported mapping", app.Resolve) + } + if !app.OptionProvenance("resolve").Has(SourceCurl) { + t.Fatal("imported resolve did not record curl provenance") + } +} + func TestCLIProxyParseErrorRedactsURLCredentials(t *testing.T) { value := "http://proxy-user:proxy-password@example.test/bad%zz?access_token=proxy-query-secret&safe=ok" _, err := Parse([]string{"--proxy", value, "https://example.com"}) diff --git a/internal/cli/provenance.go b/internal/cli/provenance.go index 8c01e6a1..4f2f215c 100644 --- a/internal/cli/provenance.go +++ b/internal/cli/provenance.go @@ -158,6 +158,9 @@ func (a *App) markCurlOptions(r *curl.Result) { if r.DoHURL != "" { a.markCurlOption("dns-server") } + if len(r.Resolve) > 0 { + a.markCurlOption("resolve") + } if r.RetrySet { a.markCurlOption("retry") } diff --git a/internal/cli/registry.go b/internal/cli/registry.go index 952cbe49..cde8966b 100644 --- a/internal/cli/registry.go +++ b/internal/cli/registry.go @@ -263,7 +263,7 @@ func (r *OptionRegistry) Ignored(mode OptionMode, explicit func(string) bool) [] for i, label := range []string{ "--data/--json/--xml", "--form", "--multipart", "--grpc", "--grpc-describe", "--grpc-list", "--output", "--remote-name", "--remote-header-name", "--copy", "--method", "--header", "--query", - "--edit", "--session", "--retry", "--retry-unsafe", "--range", "--timing", "--proxy", "--discard", "--unix", + "--edit", "--session", "--retry", "--retry-unsafe", "--range", "--timing", "--proxy", "--resolve", "--discard", "--unix", "--inspect-tls", "--bearer", "--basic", "--digest", "--aws-sigv4", "--ca-cert", "--cert", "--key", "--tls", "--max-tls", "--insecure", "--format", "--dry-run", } { @@ -279,7 +279,7 @@ func applyFlagDefinition(flag *Flag) { if flag.ConfigKey == "" { flag.ConfigKey = flag.Long } - if flag.Long == "header" || flag.Long == "query" || flag.Long == "ca-cert" { + if flag.Long == "header" || flag.Long == "query" || flag.Long == "ca-cert" || flag.Long == "resolve" { flag.Repeatable = true } if flag.Long == "form" || flag.Long == "multipart" || flag.Long == "range" || flag.Long == "verbose" { @@ -346,7 +346,8 @@ var fromCurlOptions = map[string]bool{ "range": true, "unix": true, "timeout": true, "connect-timeout": true, "redirects": true, "proxy": true, "insecure": true, "max-tls": true, "min-tls": true, "http": true, "ech": true, "cert": true, "key": true, "ca-cert": true, "dns-server": true, - "retry": true, "retry-delay": true, "grpc": true, "grpc-describe": true, + "resolve": true, + "retry": true, "retry-delay": true, "grpc": true, "grpc-describe": true, "grpc-list": true, "query": true, } @@ -408,6 +409,7 @@ var flagDefinitions = map[string]Flag{ "header": {IgnoredIn: []OptionMode{ModeDNSInspection, ModeTLSInspection}}, "query": {IgnoredIn: []OptionMode{ModeDNSInspection, ModeTLSInspection}}, + "resolve": {IgnoredIn: []OptionMode{ModeDNSInspection}}, "grpc": {Conflicts: []string{"grpc-list", "grpc-describe"}, IgnoredIn: []OptionMode{ModeDNSInspection, ModeTLSInspection}}, "grpc-describe": {Conflicts: []string{"grpc", "grpc-list"}, IgnoredIn: []OptionMode{ModeDNSInspection, ModeTLSInspection}}, "grpc-list": {Conflicts: []string{"grpc", "grpc-describe"}, IgnoredIn: []OptionMode{ModeDNSInspection, ModeTLSInspection}}, diff --git a/internal/client/automatic_ech.go b/internal/client/automatic_ech.go index fd4807ad..304603b4 100644 --- a/internal/client/automatic_ech.go +++ b/internal/client/automatic_ech.go @@ -81,6 +81,7 @@ func (t *automaticHTTP3Transport) dialAutomaticECHTCP(ctx context.Context, origi Host: host, Port: port, OriginHost: origin.Hostname(), + OriginPort: originPort(origin), Resolver: t.resolver, Candidates: addresses, }, cfg, t.ech) diff --git a/internal/client/automatic_h3.go b/internal/client/automatic_h3.go index 4dcdb038..d43baf06 100644 --- a/internal/client/automatic_h3.go +++ b/internal/client/automatic_h3.go @@ -675,6 +675,11 @@ func (t *automaticHTTP3Transport) prepare(ctx context.Context, origin *url.URL, defer race.finish() raceCtx := race.ctx innerCtx := WithoutDialTimingSelector(raceCtx) + originPortText := originPort(origin) + _, hasResolve, resolveErr := t.resolver.ResolveAddressOverride("tcp", origin.Hostname(), originPortText) + if resolveErr != nil { + return preparedAutomaticConnection{}, resolveErr + } selectTiming := func(timing DialTiming) { if selector := dialTimingSelector(ctx); selector != nil { selector.ConnectionSelected(timing) @@ -724,53 +729,58 @@ func (t *automaticHTTP3Transport) prepare(ctx context.Context, origin *url.URL, race.send(result) }() - go func() { - result := race.result("discovery") - discovery, err := t.resolver.DiscoverHTTPS(raceCtx, origin.Hostname(), uint16(parsePort(originPort(origin))), nil) - if err != nil { - kind := resolver.DiscoveryFailure(err) - if kind == resolver.DiscoveryFailureNODATA || kind == resolver.DiscoveryFailureNXDOMAIN { - // Authenticated NODATA/NXDOMAIN is a successful fresh - // replacement of the DNS RRset, not a transient failure. - t.cache.replaceDNS(key, nil) - } - result.err = err - race.send(result) - return - } - values := make([]automaticH3Candidate, 0, len(discovery.Candidates)) - echCandidates := make([]resolver.ServiceCandidate, 0, len(discovery.Candidates)) - for _, service := range discovery.Candidates { - if len(service.ECH) > 0 && serviceAdvertisesTCP(service) { - echCandidates = append(echCandidates, service) + if !hasResolve { + go func() { + result := race.result("discovery") + discovery, err := t.resolver.DiscoverHTTPS(raceCtx, origin.Hostname(), uint16(parsePort(originPortText)), nil) + if err != nil { + kind := resolver.DiscoveryFailure(err) + if kind == resolver.DiscoveryFailureNODATA || kind == resolver.DiscoveryFailureNXDOMAIN { + // Authenticated NODATA/NXDOMAIN is a successful fresh + // replacement of the DNS RRset, not a transient failure. + t.cache.replaceDNS(key, nil) + } + result.err = err + race.send(result) + return } - if !serviceAdvertisesH3(service) || len(service.Addresses) == 0 { - continue + values := make([]automaticH3Candidate, 0, len(discovery.Candidates)) + echCandidates := make([]resolver.ServiceCandidate, 0, len(discovery.Candidates)) + for _, service := range discovery.Candidates { + if len(service.ECH) > 0 && serviceAdvertisesTCP(service) { + echCandidates = append(echCandidates, service) + } + if !serviceAdvertisesH3(service) || len(service.Addresses) == 0 { + continue + } + values = append(values, automaticH3Candidate{ + host: service.TargetName.String(), + port: service.Port, + ech: append([]byte(nil), service.ECH...), + addresses: append([]net.IPAddr(nil), service.Addresses...), + expires: ttlExpiry(service.TTL, service.TTLPresent), + source: h3SourceDNS, + priority: service.Priority, + learned: time.Now(), + }) } - values = append(values, automaticH3Candidate{ - host: service.TargetName.String(), - port: service.Port, - ech: append([]byte(nil), service.ECH...), - addresses: append([]net.IPAddr(nil), service.Addresses...), - expires: ttlExpiry(service.TTL, service.TTLPresent), - source: h3SourceDNS, - priority: service.Priority, - learned: time.Now(), - }) - } - // A successful fresh RRset replaces old DNS candidates, including an - // authenticated NODATA result represented by an empty list. - t.cache.replaceDNS(key, values) - result.candidates = values - result.echCandidates = echCandidates - race.send(result) - }() + // A successful fresh RRset replaces old DNS candidates, including an + // authenticated NODATA result represented by an empty list. + t.cache.replaceDNS(key, values) + result.candidates = values + result.echCandidates = echCandidates + race.send(result) + }() + } - pending := 3 // TCP, persistent-cache load, and discovery. + pending := 2 // TCP and persistent-cache load. + if !hasResolve { + pending++ // HTTPS/SVCB discovery. + } var lastErr error var tcpPending automaticPrepareResult var haveTCP bool - discoveryDone := false + discoveryDone := hasResolve var discoveredECH []resolver.ServiceCandidate echEnabled := t.ech != core.ECHUnknown && t.ech != core.ECHOff for pending > 0 { @@ -930,6 +940,14 @@ func (t *automaticHTTP3Transport) prepare(ctx context.Context, origin *url.URL, } func (t *automaticHTTP3Transport) dialH3(ctx context.Context, origin *url.URL, candidate automaticH3Candidate) (*quic.Conn, net.PacketConn, DialTiming, error) { + if override, ok, err := t.resolver.ResolveAddressOverride("udp", origin.Hostname(), originPort(origin)); ok { + if err != nil { + return nil, nil, DialTiming{}, err + } + candidate.host = origin.Hostname() + candidate.port = uint16(parsePort(originPort(origin))) + candidate.addresses = override.Addrs + } if candidate.port == 0 || len(candidate.addresses) == 0 { return nil, nil, DialTiming{}, errors.New("HTTP/3 candidate has no address") } @@ -1071,7 +1089,7 @@ func (t *automaticHTTP3Transport) roundTripTCP(req *http.Request, prepared prepa cfg := t.tlsConfig.Clone() cfg.NextProtos = []string{"h2", "http/1.1"} got, err := t.dialer.Dial(ctx, DialRequest{ - Network: "tcp", Host: host, Port: port, OriginHost: req.URL.Hostname(), + Network: "tcp", Host: host, Port: port, OriginHost: req.URL.Hostname(), OriginPort: originPort(req.URL), Resolver: t.resolver, TLSConfig: cfg, ALPN: cfg.NextProtos, }) if err != nil { @@ -1458,9 +1476,17 @@ func (t *automaticHTTP3Transport) recordAltSvc(ctx context.Context, origin *url. if host == "" { host = origin.Hostname() } - addresses, err := t.resolver.LookupIPAddr(lookupCtx, host) - if err != nil { - continue + var addresses []net.IPAddr + if override, ok, err := t.resolver.ResolveAddressOverride("udp", host, strconv.Itoa(int(item.port))); ok { + if err != nil { + continue + } + addresses = override.Addrs + } else { + addresses, err = t.resolver.LookupIPAddr(lookupCtx, host) + if err != nil { + continue + } } t.cache.addAltSvc(key, automaticH3Candidate{host: host, port: item.port, addresses: addresses, expires: time.Now().Add(item.maxAge), source: h3SourceAltSvc, learned: time.Now()}) } diff --git a/internal/client/automatic_h3_test.go b/internal/client/automatic_h3_test.go index 0ede17ac..3ac0afca 100644 --- a/internal/client/automatic_h3_test.go +++ b/internal/client/automatic_h3_test.go @@ -75,6 +75,26 @@ func TestSplitAltSvcRejectsMalformedAuthorities(t *testing.T) { } } +func TestRecordAltSvcUsesStaticResolveEntry(t *testing.T) { + origin, err := url.Parse("https://origin.test/") + if err != nil { + t.Fatal(err) + } + res := resolver.New(resolver.Config{ + Resolve: []resolver.ResolveEntry{{Host: "origin.test", Port: "8443", IP: net.ParseIP("192.0.2.10")}}, + SystemLookupIPAddr: func(context.Context, string) ([]net.IPAddr, error) { + return nil, fmt.Errorf("unexpected DNS lookup") + }, + }) + transport := &automaticHTTP3Transport{resolver: res, cache: newAutomaticH3Cache()} + transport.recordAltSvc(context.Background(), origin, `h3=":8443"; ma=60`) + + values := transport.cache.get(automaticH3CacheKey(origin, res), time.Now()) + if len(values) != 1 || values[0].port != 8443 || len(values[0].addresses) != 1 || values[0].addresses[0].IP.String() != "192.0.2.10" { + t.Fatalf("Alt-Svc candidates = %+v, want static address on port 8443", values) + } +} + func TestAutomaticH3CacheScopesAndReplacesDNSCandidates(t *testing.T) { cache := newAutomaticH3Cache() keyA := "https://example.com:443|udp://resolver-a" diff --git a/internal/client/client.go b/internal/client/client.go index 26a48e90..3b36e560 100644 --- a/internal/client/client.go +++ b/internal/client/client.go @@ -176,6 +176,7 @@ type ClientConfig struct { CACerts []*x509.Certificate ClientCert *tls.Certificate ConnectTimeout time.Duration + Resolve []resolver.ResolveEntry // ResolverEndpoint is the validated endpoint from CLI/config parsing. // DNSServer remains for compatibility with direct internal callers/tests. ResolverEndpoint *resolver.Endpoint @@ -210,6 +211,7 @@ func NewClient(cfg ClientConfig) *Client { res := resolver.New(resolver.Config{ Endpoint: cfg.ResolverEndpoint, Server: cfg.DNSServer, + Resolve: cfg.Resolve, SystemLookupIPAddr: cfg.SystemLookupIPAddr, Proxy: proxy, TLSConfig: tlsConfig, @@ -599,6 +601,15 @@ func getHTTP3Transport(res *resolver.Resolver, tlsConfig *tls.Config, connectTim ctx, cancel := connectContext(ctx, connectTimeout, "DNS/QUIC/TLS connect") defer cancel() trace := httptrace.ContextClientTrace(ctx) + originHost, originPortText, splitErr := net.SplitHostPort(addr) + if splitErr != nil { + originHost = addr + originPortText = "443" + } + _, hasResolve, resolveErr := res.ResolveAddressOverride("udp", originHost, originPortText) + if resolveErr != nil { + return nil, resolveErr + } if trace != nil && trace.DNSStart != nil { trace.DNSStart(httptrace.DNSStartInfo{Host: addr}) } @@ -615,50 +626,54 @@ func getHTTP3Transport(res *resolver.Resolver, tlsConfig *tls.Config, connectTim return nil, err } - // HTTPS/SVCB discovery is scoped to this resolver and target. It is - // opportunistic for ordinary forced H3, but ECH=on requires an - // advertised, validated configuration. - originHost, _, splitErr := net.SplitHostPort(addr) - if splitErr != nil { - originHost = addr - } - originPort, portErr := strconv.ParseUint(endpoint.Port, 10, 16) - if portErr != nil { - return nil, portErr - } - discovery, discoveryErr := res.DiscoverHTTPS(ctx, originHost, uint16(originPort), nil) - if discoveryErr != nil { - if echMode == core.ECHOn || resolver.IsAuthenticatedDiscoveryFailure(discoveryErr) { - return nil, discoveryErr + // HTTPS/SVCB discovery is scoped to this resolver and target. A static + // --resolve mapping is authoritative, so it must not be made dependent + // on an authenticated discovery query that can never affect its address + // or port. ECH=on consequently cannot proceed without a discovered ECH + // configuration. + if hasResolve { + if echMode == core.ECHOn { + return nil, ErrECHConfigUnavailable } } else { - var selected *resolver.ServiceCandidate - for i := range discovery.Candidates { - candidate := &discovery.Candidates[i] - for _, alpn := range candidate.ALPN { - if string(alpn) == "h3" { - selected = candidate + originPort, portErr := strconv.ParseUint(endpoint.Port, 10, 16) + if portErr != nil { + return nil, portErr + } + discovery, discoveryErr := res.DiscoverHTTPS(ctx, originHost, uint16(originPort), nil) + if discoveryErr != nil { + if echMode == core.ECHOn || resolver.IsAuthenticatedDiscoveryFailure(discoveryErr) { + return nil, discoveryErr + } + } else { + var selected *resolver.ServiceCandidate + for i := range discovery.Candidates { + candidate := &discovery.Candidates[i] + for _, alpn := range candidate.ALPN { + if string(alpn) == "h3" { + selected = candidate + break + } + } + if selected != nil { break } } if selected != nil { - break - } - } - if selected != nil { - if len(selected.Addresses) > 0 { - endpoint.Addrs = selected.Addresses - } - endpoint.Port = strconv.Itoa(int(selected.Port)) - if len(selected.ECH) > 0 && (echMode == core.ECHAuto || echMode == core.ECHOn) { - tlsCfg = tlsCfg.Clone() - tlsCfg.MinVersion = tls.VersionTLS13 - tlsCfg.EncryptedClientHelloConfigList = append([]byte(nil), selected.ECH...) + if len(selected.Addresses) > 0 { + endpoint.Addrs = selected.Addresses + } + endpoint.Port = strconv.Itoa(int(selected.Port)) + if len(selected.ECH) > 0 && (echMode == core.ECHAuto || echMode == core.ECHOn) { + tlsCfg = tlsCfg.Clone() + tlsCfg.MinVersion = tls.VersionTLS13 + tlsCfg.EncryptedClientHelloConfigList = append([]byte(nil), selected.ECH...) + } else if echMode == core.ECHOn { + return nil, errors.New("ECH is required but the selected HTTPS record has no ECH configuration") + } } else if echMode == core.ECHOn { - return nil, errors.New("ECH is required but the selected HTTPS record has no ECH configuration") + return nil, errors.New("ECH is required but no HTTPS service configuration supports HTTP/3") } - } else if echMode == core.ECHOn { - return nil, errors.New("ECH is required but no HTTPS service configuration supports HTTP/3") } } if echMode == core.ECHAuto && len(tlsCfg.EncryptedClientHelloConfigList) == 0 { diff --git a/internal/client/dialer.go b/internal/client/dialer.go index 1f149ef6..c75ad73c 100644 --- a/internal/client/dialer.go +++ b/internal/client/dialer.go @@ -87,12 +87,15 @@ func dialTimingSelector(ctx context.Context) DialTimingSelector { // the effective service target. OriginHost is used for TLS SNI only when the // TLS config does not already provide ServerName; this preserves origin // authority when an HTTPS/SVCB service target differs from the origin host. +// OriginPort is the URL authority port used to match static resolve entries +// when Port has been replaced by an HTTPS/SVCB service port. type DialRequest struct { Network string Address string Host string Port string OriginHost string + OriginPort string Mode DialMode Resolver *resolver.Resolver @@ -232,6 +235,25 @@ func (d *ResolverDialer) Dial(ctx context.Context, req DialRequest) (DialResult, } candidates := append([]net.IPAddr(nil), req.Candidates...) + // A static --resolve entry is authoritative for the request authority, + // including when a prior discovery step supplied a different candidate. + // OriginHost matters for SVCB/ECH targets: the mapping belongs to the URL + // authority, not to the service name selected during discovery. + resolveHost := req.Host + if req.OriginHost != "" { + resolveHost = req.OriginHost + } + resolvePort := req.Port + if req.OriginPort != "" { + resolvePort = req.OriginPort + } + if override, ok, err := res.ResolveAddressOverride(req.Network, resolveHost, resolvePort); ok { + if err != nil { + return DialResult{}, err + } + candidates = override.Addrs + req.Port = override.Port + } if len(candidates) == 0 { result.Timing.ResolutionStart = time.Now() if req.Recorder != nil { diff --git a/internal/client/dialer_test.go b/internal/client/dialer_test.go index c5787d68..a6391a25 100644 --- a/internal/client/dialer_test.go +++ b/internal/client/dialer_test.go @@ -59,6 +59,67 @@ func TestResolverDialerResolvesAndReportsWinningAddress(t *testing.T) { } } +func TestResolverDialerUsesStaticResolveEntry(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + lookupCalled := false + res := resolver.New(resolver.Config{ + Resolve: []resolver.ResolveEntry{{ + Host: "service.test", Port: strconv.Itoa(listener.Addr().(*net.TCPAddr).Port), IP: net.ParseIP("127.0.0.1"), + }}, + SystemLookupIPAddr: func(context.Context, string) ([]net.IPAddr, error) { + lookupCalled = true + return nil, errors.New("unexpected lookup") + }, + }) + dialer := NewResolverDialer(res, time.Second) + result, err := dialer.Dial(context.Background(), DialRequest{ + Network: "tcp", + Host: "service.test", + Port: strconv.Itoa(listener.Addr().(*net.TCPAddr).Port), + }) + if err != nil { + t.Fatal(err) + } + result.Conn.Close() + if lookupCalled { + t.Fatal("static resolve entry performed a DNS lookup") + } + if got := result.RemoteIP.String(); got != "127.0.0.1" { + t.Fatalf("remote IP = %s, want 127.0.0.1", got) + } +} + +func TestResolverDialerUsesOriginPortForStaticServiceTarget(t *testing.T) { + var dialed string + res := resolver.New(resolver.Config{Resolve: []resolver.ResolveEntry{{ + Host: "origin.test", Port: "443", IP: net.ParseIP("192.0.2.10"), + }}}) + dialer := NewResolverDialer(res, time.Second) + dialer.BaseDial = func(_ context.Context, _, address string) (net.Conn, error) { + dialed = address + conn, peer := net.Pipe() + _ = peer.Close() + return conn, nil + } + + result, err := dialer.Dial(context.Background(), DialRequest{ + Network: "tcp", Host: "edge.test", Port: "8443", + OriginHost: "origin.test", OriginPort: "0443", + }) + if err != nil { + t.Fatal(err) + } + result.Conn.Close() + if dialed != "192.0.2.10:443" { + t.Fatalf("dial address = %q, want 192.0.2.10:443", dialed) + } +} + func TestResolverDialerPreservesCandidatePreferenceAndClosesLoser(t *testing.T) { first := net.ParseIP("2001:db8::1") second := net.ParseIP("192.0.2.1") diff --git a/internal/client/ech.go b/internal/client/ech.go index 87a57f62..1e7e087c 100644 --- a/internal/client/ech.go +++ b/internal/client/ech.go @@ -252,7 +252,7 @@ func newECHHTTPDialTLS(base func(context.Context, string, string) (net.Conn, err } got, err := dialResolverWithECH(connectCtx, NewResolverDialer(res, timeout), DialRequest{ Network: "tcp", Host: connection.targetHost, Port: connection.targetPort, - OriginHost: host, Candidates: connection.addresses, + OriginHost: host, OriginPort: port, Candidates: connection.addresses, }, connection.tlsConfig, mode) if err != nil { return nil, err diff --git a/internal/client/ech_handshake.go b/internal/client/ech_handshake.go index 668d95db..a3029c0c 100644 --- a/internal/client/ech_handshake.go +++ b/internal/client/ech_handshake.go @@ -165,6 +165,31 @@ func dialResolverWithECHInfo(ctx context.Context, dialer *ResolverDialer, reques if port == "" { return DialResult{}, errors.New("ECH resolver dial requires a port") } + resolveHost := request.Host + if request.OriginHost != "" { + resolveHost = request.OriginHost + } + resolvePort := port + if request.OriginPort != "" { + resolvePort = request.OriginPort + } + resolveNetwork := request.Network + if resolveNetwork == "" { + resolveNetwork = "tcp" + } + res := request.Resolver + if res == nil { + res = dialer.Resolver + } + if res != nil { + if override, ok, err := res.ResolveAddressOverride(resolveNetwork, resolveHost, resolvePort); ok { + if err != nil { + return DialResult{}, err + } + port = override.Port + request.Port = port + } + } request.TLSConfig = nil request.ALPN = nil request.AttemptWithInfo = func(attemptCtx context.Context, network string, ip net.IPAddr) (net.Conn, any, error) { diff --git a/internal/client/ech_handshake_test.go b/internal/client/ech_handshake_test.go index 1af05cc6..18c307b3 100644 --- a/internal/client/ech_handshake_test.go +++ b/internal/client/ech_handshake_test.go @@ -9,14 +9,39 @@ import ( "crypto/tls" "crypto/x509" "crypto/x509/pkix" + "errors" "math/big" "net" "testing" "time" "github.com/ryanfowler/fetch/internal/core" + "github.com/ryanfowler/fetch/internal/resolver" ) +func TestDialResolverWithECHUsesResolveOriginPort(t *testing.T) { + res := resolver.New(resolver.Config{Resolve: []resolver.ResolveEntry{{ + Host: "origin.test", Port: "443", IP: net.ParseIP("192.0.2.10"), + }}}) + var dialed string + dialer := NewResolverDialer(res, time.Second) + dialer.BaseDial = func(_ context.Context, _, address string) (net.Conn, error) { + dialed = address + return nil, errors.New("stop before TLS") + } + + _, err := dialResolverWithECHInfo(context.Background(), dialer, DialRequest{ + Network: "tcp", Host: "edge.test", Port: "8443", OriginHost: "origin.test", OriginPort: "0443", + Resolver: res, Candidates: []net.IPAddr{{IP: net.ParseIP("192.0.2.11")}}, + }, &tls.Config{}, core.ECHOff, false) + if err == nil { + t.Fatal("dialResolverWithECHInfo succeeded, want test dial failure") + } + if dialed != "192.0.2.10:443" { + t.Fatalf("dial address = %q, want 192.0.2.10:443", dialed) + } +} + func TestDialTLSWithECHPolicyAcceptsRealECH(t *testing.T) { certificate, roots := echTestCertificate(t) echPrivate, err := ecdh.X25519().GenerateKey(cryptorand.Reader) diff --git a/internal/client/http2_proxy.go b/internal/client/http2_proxy.go index a7b3cb26..1130a52b 100644 --- a/internal/client/http2_proxy.go +++ b/internal/client/http2_proxy.go @@ -47,7 +47,7 @@ func newHTTP2DialTLS(base func(context.Context, string, string) (net.Conn, error if connection.configured { got, dialErr := dialResolverWithECH(ctx, NewResolverDialer(res, timeout), DialRequest{ Network: "tcp", Host: connection.targetHost, Port: connection.targetPort, - OriginHost: host, Candidates: connection.addresses, + OriginHost: host, OriginPort: port, Candidates: connection.addresses, }, connection.tlsConfig, echMode) if dialErr != nil { return nil, dialErr diff --git a/internal/curl/curl.go b/internal/curl/curl.go index 4c12d144..cfcbcde8 100644 --- a/internal/curl/curl.go +++ b/internal/curl/curl.go @@ -52,6 +52,7 @@ type Result struct { ConnectTimeoutSet bool Proxy string DoHURL string + Resolve []string HTTPVersion string TLSMaxVersion string TLSVersion string diff --git a/internal/curl/curl_test.go b/internal/curl/curl_test.go index cf3431df..b4dfa705 100644 --- a/internal/curl/curl_test.go +++ b/internal/curl/curl_test.go @@ -666,6 +666,13 @@ func TestParseNetwork(t *testing.T) { assertEqual(t, "DoHURL", r.DoHURL, "https://1.1.1.1/dns-query") }, }, + { + name: "resolve", + input: "curl --resolve example.com:443:192.0.2.10 --resolve '*:80:192.0.2.11' https://example.com", + check: func(t *testing.T, r *Result) { + assertSliceEqual(t, "Resolve", r.Resolve, []string{"example.com:443:192.0.2.10", "*:80:192.0.2.11"}) + }, + }, { name: "retry", input: "curl --retry 3 https://example.com", diff --git a/internal/curl/long_flags.go b/internal/curl/long_flags.go index 68ce0622..997f0128 100644 --- a/internal/curl/long_flags.go +++ b/internal/curl/long_flags.go @@ -267,6 +267,13 @@ func parseLongFlag(r *Result, name, value string, hasValue bool, rest []string) } r.DoHURL = v return n, nil + case "resolve": + v, n, err := consumeArg() + if err != nil { + return 0, fmt.Errorf("--resolve requires an argument") + } + r.Resolve = append(r.Resolve, v) + return n, nil case "retry": v, n, err := consumeArg() if err != nil { diff --git a/internal/fetch/fetch.go b/internal/fetch/fetch.go index c65aacec..11456a66 100644 --- a/internal/fetch/fetch.go +++ b/internal/fetch/fetch.go @@ -98,6 +98,7 @@ type Request struct { Redirects *int RemoteHeaderName bool RemoteName bool + Resolve []resolver.ResolveEntry Retry int RetryDelay time.Duration RetryUnsafe bool diff --git a/internal/fetch/grpc_reflection.go b/internal/fetch/grpc_reflection.go index 02974898..9b6a3966 100644 --- a/internal/fetch/grpc_reflection.go +++ b/internal/fetch/grpc_reflection.go @@ -841,6 +841,7 @@ func newClient(r *Request) *client.Client { ECH: r.ECH, Insecure: r.Insecure, Proxy: r.Proxy, + Resolve: r.Resolve, Redirects: r.Redirects, TLSMax: r.TLSMax, TLSMin: r.TLSMin, diff --git a/internal/resolver/doh.go b/internal/resolver/doh.go index abcc0e91..e0500512 100644 --- a/internal/resolver/doh.go +++ b/internal/resolver/doh.go @@ -41,6 +41,7 @@ type DOHConfig struct { Proxy func(*http.Request) (*url.URL, error) DialContext DialContextFunc Bootstrap BootstrapFunc + Resolve []ResolveEntry TLSConfig *tls.Config CACerts []*x509.Certificate ClientCert *tls.Certificate @@ -103,7 +104,7 @@ func NewDOHClient(cfg DOHConfig) (*DOHClient, error) { var d net.Dialer dial = d.DialContext } - base.DialContext = dohDialContext(dial, cfg.Bootstrap, cfg.Endpoint, serverURL) + base.DialContext = dohDialContext(dial, cfg.Bootstrap, cfg.Endpoint, serverURL, cfg.Resolve) transport = base } @@ -160,7 +161,7 @@ func dohTLSConfig(cfg DOHConfig, serverName string) *tls.Config { }) } -func dohDialContext(dial DialContextFunc, bootstrap BootstrapFunc, endpoint *Endpoint, serverURL *url.URL) DialContextFunc { +func dohDialContext(dial DialContextFunc, bootstrap BootstrapFunc, endpoint *Endpoint, serverURL *url.URL, resolve []ResolveEntry) DialContextFunc { endpointHost := strings.TrimSuffix(serverURL.Hostname(), ".") var bootstrapAddrs []net.IPAddr if endpoint != nil { @@ -177,6 +178,25 @@ func dohDialContext(dial DialContextFunc, bootstrap BootstrapFunc, endpoint *End // the base dialer is the deliberate, narrow bootstrap exception. return dial(ctx, network, address) } + var lastErr error + if override, ok, overrideErr := resolveAddressOverride(resolve, network, host, port); ok { + if overrideErr != nil { + return nil, overrideErr + } + addresses := override.Addrs + for _, ip := range addresses { + conn, dialErr := dial(ctx, network, core.JoinIPHostPort(ip, port)) + if dialErr == nil { + return conn, nil + } + lastErr = dialErr + } + if lastErr == nil { + lastErr = errors.New("no DoH endpoint resolve addresses") + } + return nil, lastErr + } + addresses := bootstrapAddrs if len(addresses) == 0 && bootstrap != nil && net.ParseIP(host) == nil { addresses, err = bootstrap(ctx, host) @@ -187,7 +207,6 @@ func dohDialContext(dial DialContextFunc, bootstrap BootstrapFunc, endpoint *End if len(addresses) == 0 { return dial(ctx, network, address) } - var lastErr error for _, ip := range addresses { conn, dialErr := dial(ctx, network, core.JoinIPHostPort(ip, port)) if dialErr == nil { diff --git a/internal/resolver/resolve.go b/internal/resolver/resolve.go new file mode 100644 index 00000000..d476f8a6 --- /dev/null +++ b/internal/resolver/resolve.go @@ -0,0 +1,179 @@ +package resolver + +import ( + "fmt" + "net" + "strconv" + "strings" + "unicode" +) + +// ResolveEntry is a static host-to-address mapping supplied by --resolve. +// Host and Port are the request authority; IP is used only for the first hop. +type ResolveEntry struct { + Host string + Port string + IP net.IP +} + +// ParseResolve parses one curl HOST:PORT:IP mapping. Use ParseResolveEntries +// when the address component may contain multiple comma-separated addresses. +func ParseResolve(value string) (ResolveEntry, error) { + entries, err := ParseResolveEntries(value) + if err != nil { + return ResolveEntry{}, err + } + if len(entries) != 1 { + return ResolveEntry{}, fmt.Errorf("multiple addresses require ParseResolveEntries") + } + return entries[0], nil +} + +// ParseResolveEntries parses curl's HOST:PORT:IP[,IP] mapping syntax. The +// optional leading '+' is accepted for compatibility with curl's expiring +// resolve entries; this client keeps entries for the lifetime of the request. +func ParseResolveEntries(value string) ([]ResolveEntry, error) { + if value == "" { + return nil, fmt.Errorf("value is empty") + } + if strings.TrimSpace(value) != value || strings.ContainsAny(value, "\r\n\x00") { + return nil, fmt.Errorf("value must not contain leading/trailing whitespace or control characters") + } + if strings.HasPrefix(value, "+") { + value = value[1:] + if value == "" { + return nil, fmt.Errorf("value is empty") + } + } + host, port, addresses, err := splitResolveValue(value) + if err != nil { + return nil, err + } + ipValues, err := splitResolveAddresses(addresses) + if err != nil { + return nil, err + } + entries := make([]ResolveEntry, 0, len(ipValues)) + for _, ipText := range ipValues { + entry, err := parseResolveParts(host, port, ipText) + if err != nil { + return nil, err + } + entries = append(entries, entry) + } + return entries, nil +} + +func parseResolveParts(host, port, ipText string) (ResolveEntry, error) { + if host == "" { + return ResolveEntry{}, fmt.Errorf("host is empty") + } + if host != "*" { + if strings.Contains(host, "*") || strings.ContainsAny(host, "/?#[]") || strings.IndexFunc(host, func(r rune) bool { + return unicode.IsSpace(r) || unicode.IsControl(r) + }) >= 0 || (net.ParseIP(host) == nil && strings.Contains(host, ":")) { + return ResolveEntry{}, fmt.Errorf("invalid host %q", host) + } + } + + if port == "" { + return ResolveEntry{}, fmt.Errorf("port is empty") + } + for _, ch := range port { + if ch < '0' || ch > '9' { + return ResolveEntry{}, fmt.Errorf("invalid port %q", port) + } + } + portNumber, err := strconv.ParseUint(port, 10, 16) + if err != nil || portNumber == 0 { + return ResolveEntry{}, fmt.Errorf("invalid port %q", port) + } + + if strings.HasPrefix(ipText, "[") { + if !strings.HasSuffix(ipText, "]") { + return ResolveEntry{}, fmt.Errorf("invalid IP address %q", ipText) + } + ipText = ipText[1 : len(ipText)-1] + } else if strings.ContainsAny(ipText, "[]") || strings.Contains(ipText, ":") { + return ResolveEntry{}, fmt.Errorf("invalid IP address %q", ipText) + } + ip := net.ParseIP(ipText) + if ip == nil { + return ResolveEntry{}, fmt.Errorf("invalid IP address %q", ipText) + } + + if host != "*" { + host = strings.ToLower(strings.TrimSuffix(host, ".")) + if host == "" { + return ResolveEntry{}, fmt.Errorf("host is empty") + } + } + return ResolveEntry{Host: host, Port: strconv.FormatUint(portNumber, 10), IP: append(net.IP(nil), ip...)}, nil +} + +func splitResolveAddresses(value string) ([]string, error) { + values := make([]string, 0, 1) + start := 0 + brackets := 0 + for i := 0; i < len(value); i++ { + switch value[i] { + case '[': + brackets++ + case ']': + if brackets == 0 { + return nil, fmt.Errorf("invalid closing bracket") + } + brackets-- + case ',': + if brackets != 0 { + continue + } + if i > start { + values = append(values, value[start:i]) + } + start = i + 1 + } + } + if brackets != 0 { + return nil, fmt.Errorf("invalid bracketed address") + } + if start < len(value) { + values = append(values, value[start:]) + } + if len(values) == 0 { + return nil, fmt.Errorf("IP address is empty") + } + return values, nil +} + +func splitResolveValue(value string) (host, port, ip string, err error) { + rest := value + if strings.HasPrefix(rest, "[") { + close := strings.IndexByte(rest, ']') + if close < 0 { + return "", "", "", fmt.Errorf("invalid bracketed host") + } + host = rest[1:close] + rest = rest[close+1:] + if !strings.HasPrefix(rest, ":") { + return "", "", "", fmt.Errorf("must be in the format HOST:PORT:IP") + } + rest = rest[1:] + } else { + idx := strings.IndexByte(rest, ':') + if idx < 0 { + return "", "", "", fmt.Errorf("must be in the format HOST:PORT:IP") + } + host = rest[:idx] + rest = rest[idx+1:] + } + idx := strings.IndexByte(rest, ':') + if idx < 0 { + return "", "", "", fmt.Errorf("must be in the format HOST:PORT:IP") + } + port, ip = rest[:idx], rest[idx+1:] + if ip == "" { + return "", "", "", fmt.Errorf("IP address is empty") + } + return host, port, ip, nil +} diff --git a/internal/resolver/resolve_test.go b/internal/resolver/resolve_test.go new file mode 100644 index 00000000..1b8c51a1 --- /dev/null +++ b/internal/resolver/resolve_test.go @@ -0,0 +1,89 @@ +package resolver + +import ( + "context" + "net" + "reflect" + "testing" +) + +func TestParseResolve(t *testing.T) { + tests := []struct { + name string + value string + host string + port string + ip string + }{ + {name: "IPv4", value: "example.com:443:192.0.2.10", host: "example.com", port: "443", ip: "192.0.2.10"}, + {name: "normalizes host and port", value: "EXAMPLE.COM.:0443:192.0.2.10", host: "example.com", port: "443", ip: "192.0.2.10"}, + {name: "wildcard", value: "*:80:192.0.2.10", host: "*", port: "80", ip: "192.0.2.10"}, + {name: "IPv6 address", value: "example.com:443:[2001:db8::10]", host: "example.com", port: "443", ip: "2001:db8::10"}, + {name: "IPv6 host", value: "[2001:db8::1]:443:192.0.2.10", host: "2001:db8::1", port: "443", ip: "192.0.2.10"}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got, err := ParseResolve(test.value) + if err != nil { + t.Fatal(err) + } + if got.Host != test.host || got.Port != test.port || got.IP.String() != test.ip { + t.Fatalf("entry = %+v, want %s:%s:%s", got, test.host, test.port, test.ip) + } + }) + } + + for _, value := range []string{ + "", "example.com:443", "example.com::192.0.2.1", "example.com:0:192.0.2.1", + "example.com:443:localhost", "example.com:443:192.0.2.1:extra", "example.com:443: [::1]", + "example.com:443:2001:db8::1", + "[2001:db8::1:443:192.0.2.1", "bad host:443:192.0.2.1", + } { + t.Run("rejects "+value, func(t *testing.T) { + if _, err := ParseResolve(value); err == nil { + t.Fatalf("ParseResolve(%q) succeeded, want error", value) + } + }) + } +} + +func TestParseResolveEntries(t *testing.T) { + got, err := ParseResolveEntries("+example.com:443:,192.0.2.1,,[2001:db8::1],") + if err != nil { + t.Fatal(err) + } + if len(got) != 2 || got[0].Host != "example.com" || got[1].IP.String() != "2001:db8::1" { + t.Fatalf("entries = %+v, want two addresses", got) + } +} + +func TestResolveAddressUsesExactMappingBeforeWildcard(t *testing.T) { + resolver := New(Config{ + Resolve: []ResolveEntry{ + {Host: "*", Port: "443", IP: net.ParseIP("192.0.2.1")}, + {Host: "example.com", Port: "443", IP: net.ParseIP("2001:db8::1")}, + }, + SystemLookupIPAddr: func(context.Context, string) ([]net.IPAddr, error) { + t.Fatal("system lookup was used for a static mapping") + return nil, nil + }, + }) + + got, err := resolver.ResolveAddress(context.Background(), "tcp", "EXAMPLE.COM:0443") + if err != nil { + t.Fatal(err) + } + want := []net.IPAddr{{IP: net.ParseIP("2001:db8::1")}} + if !reflect.DeepEqual(got.Addrs, want) { + t.Fatalf("addresses = %v, want %v", got.Addrs, want) + } +} + +func TestResolveAddressOverrideFiltersNetwork(t *testing.T) { + resolver := New(Config{Resolve: []ResolveEntry{ + {Host: "example.com", Port: "443", IP: net.ParseIP("192.0.2.1")}, + }}) + if _, ok, err := resolver.ResolveAddressOverride("tcp6", "example.com", "443"); !ok || err == nil { + t.Fatalf("override = %v, %v, want a matched network error", ok, err) + } +} diff --git a/internal/resolver/resolver.go b/internal/resolver/resolver.go index add463d1..9c0e4477 100644 --- a/internal/resolver/resolver.go +++ b/internal/resolver/resolver.go @@ -12,6 +12,7 @@ import ( "net/http" "net/http/httptrace" "net/url" + "strconv" "strings" "github.com/ryanfowler/fetch/internal/core" @@ -23,6 +24,7 @@ import ( type Config struct { Endpoint *Endpoint Server *url.URL + Resolve []ResolveEntry // SystemLookupIPAddr replaces the platform resolver in deterministic tests. // Production callers leave it nil so net.Resolver remains authoritative for @@ -71,6 +73,7 @@ type Resolver struct { tlsMin uint16 tlsMax uint16 systemLookup func(context.Context, string) ([]net.IPAddr, error) + resolve []ResolveEntry } // ResolvedEndpoint contains a parsed host:port address and its resolved IP @@ -115,6 +118,7 @@ func New(cfg Config) *Resolver { tlsMin: cfg.TLSMin, tlsMax: cfg.TLSMax, systemLookup: systemLookup, + resolve: cloneResolveEntries(cfg.Resolve), } if r.err == nil && endpoint != nil && endpoint.Transport == TransportHTTPS { r.dohClient, r.err = NewDOHClient(DOHConfig{ @@ -123,6 +127,7 @@ func New(cfg Config) *Resolver { Proxy: cfg.Proxy, DialContext: cfg.DialContext, Bootstrap: cfg.Bootstrap, + Resolve: cfg.Resolve, TLSConfig: cfg.TLSConfig, CACerts: cfg.CACerts, ClientCert: cfg.ClientCert, @@ -275,6 +280,9 @@ func (r *Resolver) ResolveAddress(ctx context.Context, network, address string) if err != nil { return ResolvedEndpoint{}, err } + if endpoint, ok, err := r.ResolveAddressOverride(network, host, port); ok { + return endpoint, err + } addrs, err := r.LookupIPAddr(ctx, host) if err != nil { @@ -284,12 +292,60 @@ func (r *Resolver) ResolveAddress(ctx context.Context, network, address string) return ResolvedEndpoint{}, fmt.Errorf("lookup %s: no addresses found", host) } + return resolvedEndpoint(network, host, port, addrs) +} + +// ResolveAddressOverride returns addresses configured for host and port by a +// static resolve entry. The boolean distinguishes an absent entry from an +// entry whose address is unusable for the requested network. +func (r *Resolver) ResolveAddressOverride(network, host, port string) (ResolvedEndpoint, bool, error) { + if r == nil || len(r.resolve) == 0 { + return ResolvedEndpoint{}, false, nil + } + return resolveAddressOverride(r.resolve, network, host, port) +} + +func resolveAddressOverride(entries []ResolveEntry, network, host, port string) (ResolvedEndpoint, bool, error) { + port = normalizeResolvePort(port) + var exact, wildcard []net.IPAddr + for _, entry := range entries { + if normalizeResolvePort(entry.Port) != port { + continue + } + addr := net.IPAddr{IP: append(net.IP(nil), entry.IP...)} + if entry.Host == "*" { + wildcard = append(wildcard, addr) + } else if strings.EqualFold(strings.TrimSuffix(entry.Host, "."), strings.TrimSuffix(host, ".")) { + exact = append(exact, addr) + } + } + if len(exact) == 0 && len(wildcard) == 0 { + return ResolvedEndpoint{}, false, nil + } + if len(exact) > 0 { + endpoint, err := resolvedEndpoint(network, host, port, exact) + return endpoint, true, err + } + endpoint, err := resolvedEndpoint(network, host, port, wildcard) + return endpoint, true, err +} + +func normalizeResolvePort(port string) string { + portNumber, err := strconv.ParseUint(port, 10, 16) + if err != nil || portNumber == 0 { + return port + } + return strconv.FormatUint(portNumber, 10) +} + +func resolvedEndpoint(network, host, port string, addrs []net.IPAddr) (ResolvedEndpoint, error) { + addrs = deduplicateAddresses(addrs) // A family-specific network must not waste attempts on addresses that the // platform dialer will reject. For dual-stack networks, retain the // resolver-preferred family and interleave the other family below. switch strings.ToLower(network) { case "tcp4", "udp4": - filtered := addrs[:0] + filtered := make([]net.IPAddr, 0, len(addrs)) for _, addr := range addrs { if addr.IP.To4() != nil { filtered = append(filtered, addr) @@ -297,7 +353,7 @@ func (r *Resolver) ResolveAddress(ctx context.Context, network, address string) } addrs = filtered case "tcp6", "udp6": - filtered := addrs[:0] + filtered := make([]net.IPAddr, 0, len(addrs)) for _, addr := range addrs { if addr.IP.To4() == nil && addr.IP.To16() != nil { filtered = append(filtered, addr) @@ -319,6 +375,15 @@ func (r *Resolver) ResolveAddress(ctx context.Context, network, address string) return ResolvedEndpoint{Host: host, Port: port, Addrs: addrs}, nil } +func cloneResolveEntries(values []ResolveEntry) []ResolveEntry { + out := make([]ResolveEntry, len(values)) + for i, value := range values { + out[i] = value + out[i].IP = append(net.IP(nil), value.IP...) + } + return out +} + // DialContext resolves address and dials each returned IP until one succeeds. // DialContext resolves address and races its candidates with the shared // Happy Eyeballs policy. The first address retains the resolver's preferred @@ -334,6 +399,13 @@ func (r *Resolver) DialContext(ctx context.Context, network, address string) (ne for _, ip := range r.endpoint.BootstrapAddrs { addrs = append(addrs, net.IPAddr{IP: append(net.IP(nil), ip...)}) } + if override, ok, overrideErr := r.ResolveAddressOverride(network, host, port); ok { + if overrideErr != nil { + return nil, overrideErr + } + addrs = override.Addrs + port = override.Port + } if len(addrs) == 0 { var lookupErr error addrs, lookupErr = net.DefaultResolver.LookupIPAddr(ctx, host) diff --git a/internal/resolver/resolver_test.go b/internal/resolver/resolver_test.go index d7d22a6f..2009f66a 100644 --- a/internal/resolver/resolver_test.go +++ b/internal/resolver/resolver_test.go @@ -10,6 +10,7 @@ import ( "net/http/httptest" "net/http/httptrace" "net/url" + "strconv" "strings" "sync" "testing" @@ -68,6 +69,30 @@ func TestLookupIPAddrDOHNXDomain(t *testing.T) { } } +func TestDOHClientUsesResolveForEndpointBootstrap(t *testing.T) { + var dialed string + client, err := NewDOHClient(DOHConfig{ + ServerURL: mustURL(t, "http://resolver.example/dns-query"), + Resolve: []ResolveEntry{{Host: "resolver.example", Port: "80", IP: net.ParseIP("192.0.2.10")}}, + DialContext: func(_ context.Context, _, address string) (net.Conn, error) { + dialed = address + return nil, errors.New("stop bootstrap") + }, + }) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + _, err = client.LookupType(context.Background(), "example.com", "A", dnsTypeA) + if err == nil || !strings.Contains(err.Error(), "stop bootstrap") { + t.Fatalf("LookupType error = %v, want bootstrap dial error", err) + } + if dialed != "192.0.2.10:80" { + t.Fatalf("DoH endpoint dial address = %q, want 192.0.2.10:80", dialed) + } +} + func TestLookupDOHTypeReturnsTTL(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, `{"Status":0,"Answer":[{"name":"example.com","type":1,"data":"127.0.0.1","TTL":123}]}`) @@ -666,6 +691,42 @@ func TestDialContextUsesResolvedAddress(t *testing.T) { <-accepted } +func TestDialContextUsesResolveForDOHEndpoint(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + accepted := make(chan struct{}) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr == nil { + _ = conn.Close() + } + close(accepted) + }() + + port := listener.Addr().(*net.TCPAddr).Port + serverURL, err := url.Parse(fmt.Sprintf("https://resolver.example:%d/dns-query", port)) + if err != nil { + t.Fatal(err) + } + r := New(Config{ + Server: serverURL, + Resolve: []ResolveEntry{{Host: "resolver.example", Port: strconv.Itoa(port), IP: net.ParseIP("127.0.0.1")}}, + }) + conn, err := r.DialContext(context.Background(), "tcp", net.JoinHostPort("resolver.example", strconv.Itoa(port))) + if err != nil { + t.Fatal(err) + } + _ = conn.Close() + select { + case <-accepted: + case <-time.After(time.Second): + t.Fatal("DoH endpoint did not receive the statically resolved connection") + } +} + func mustURL(t *testing.T, raw string) *url.URL { t.Helper() u, err := url.Parse(raw) diff --git a/internal/tlsinspect/tlsinspect.go b/internal/tlsinspect/tlsinspect.go index 7b180371..1d0ee734 100644 --- a/internal/tlsinspect/tlsinspect.go +++ b/internal/tlsinspect/tlsinspect.go @@ -34,6 +34,7 @@ type Config struct { ClientCert *tls.Certificate ResolverEndpoint *resolver.Endpoint DNSServer *url.URL + Resolve []resolver.ResolveEntry HTTP core.HTTPVersion ECH core.ECHMode Insecure bool @@ -87,6 +88,7 @@ func Inspect(ctx context.Context, p *core.Printer, cfg *Config) int { res := resolver.New(resolver.Config{ Endpoint: cfg.ResolverEndpoint, Server: cfg.DNSServer, + Resolve: cfg.Resolve, CACerts: cfg.CACerts, ClientCert: cfg.ClientCert, Insecure: cfg.Insecure, @@ -142,6 +144,14 @@ func Inspect(ctx context.Context, p *core.Printer, cfg *Config) int { quicAddr = net.JoinHostPort(targetHost, targetPort) quicCandidates = echConfig.Addresses() } + if override, ok, err := res.ResolveAddressOverride("udp", host, port); ok { + if err != nil { + writeTLSError(p, err) + return 1 + } + quicAddr = addr + quicCandidates = override.Addrs + } var fallbackTLS *tls.Config if cfg.ECH == core.ECHAuto && echConfig != nil && echConfig.Offered() { fallbackTLS = tlsConfig.Clone() @@ -190,6 +200,7 @@ func Inspect(ctx context.Context, p *core.Printer, cfg *Config) int { dialRequest.Address = "" dialRequest.Host = targetHost dialRequest.Port = targetPort + dialRequest.OriginPort = port dialRequest.Candidates = echConfig.Addresses() } var result client.DialResult diff --git a/main.go b/main.go index 4b5d6ca4..69eee2d7 100644 --- a/main.go +++ b/main.go @@ -240,6 +240,7 @@ func main() { Redirects: app.Cfg.Redirects, RemoteHeaderName: app.RemoteHeaderName, RemoteName: app.RemoteName, + Resolve: app.Resolve, Retry: getValue(app.Cfg.Retry), RetryDelay: getValue(app.Cfg.RetryDelay), RetryUnsafe: getValue(app.Cfg.RetryUnsafe), @@ -415,6 +416,7 @@ func conciseHelpSections() []helpSection { rows: []helpRow{ {"--http VERSION", "Select HTTP/1.1, HTTP/2, or HTTP/3"}, {"--proxy PROXY", "Use a proxy"}, + {"--resolve [+]HOST:PORT:IP[,IP]", "Connect to IP preserving Host/SNI"}, {"--timeout SECONDS", "Set the request timeout"}, {"--dry-run", "Print out the request info and exit"}, {"--inspect-dns", "Inspect DNS resolution"}, @@ -873,6 +875,7 @@ func inspectTLS(ctx context.Context, app *cli.App, handle *core.Handle) int { ClientCert: clientCert, ResolverEndpoint: app.Cfg.DNSEndpoint, DNSServer: app.Cfg.DNSServer, + Resolve: app.Resolve, HTTP: app.Cfg.HTTP, ECH: app.Cfg.ECH, Insecure: getValue(app.Cfg.Insecure),