diff --git a/colly.go b/colly.go index b4b96e61e..5447624e2 100644 --- a/colly.go +++ b/colly.go @@ -673,7 +673,13 @@ func (c *Collector) scrape(u, method string, depth int, requestData io.Reader, c } // note: once 1.13 is minimum supported Go version, // replace this with http.NewRequestWithContext - req = req.WithContext(context.WithValue(c.Context, CheckRevisitKey, checkRevisit)) + // Place a *string holder in the context so the proxy URL chosen by the + // ProxyFunc survives net/http's send() forkReq (triggered by Client.Timeout). + // The fork shallow-copies the request and shares the context pointer, + // so pointer writes through the holder remain visible on the original req. + req = req.WithContext(context.WithValue( + context.WithValue(c.Context, CheckRevisitKey, checkRevisit), + ProxyURLKey, new(string))) if err := c.requestCheck(parsedURL, method, req.GetBody, depth, checkRevisit); err != nil { return err @@ -737,8 +743,8 @@ func (c *Collector) fetch(u, method string, depth int, requestData io.Reader, ct return !request.abort } response, err := c.backend.Cache(req, c.MaxBodySize, checkRequestHeadersFunc, checkResponseHeadersFunc, c.CacheDir, c.CacheExpiration) - if proxyURL, ok := req.Context().Value(ProxyURLKey).(string); ok { - request.ProxyURL = proxyURL + if proxyURL, ok := req.Context().Value(ProxyURLKey).(*string); ok { + request.ProxyURL = *proxyURL } if err := c.handleOnError(response, err, request, ctx); err != nil { return err @@ -747,6 +753,7 @@ func (c *Collector) fetch(u, method string, depth int, requestData io.Reader, ct response.Ctx = ctx response.Request = request response.Trace = hTrace + response.ProxyURL = request.ProxyURL err = response.fixCharset(c.DetectCharset, request.ResponseCharacterEncoding) if err != nil { @@ -1108,7 +1115,23 @@ func (c *Collector) SetProxy(proxyURL string) error { // The proxy type is determined by the URL scheme. "http" // and "socks5" are supported. If the scheme is empty, // "http" is assumed. -func (c *Collector) SetProxyFunc(p ProxyFunc) { +func (c *Collector) SetProxyFunc(f ProxyFunc) { + + var p ProxyFunc = func(pr *http.Request) (*url.URL, error) { + // Capture the context before invoking the user's f. Legacy custom + // ProxyFuncs may do *pr = *pr.WithContext(WithValue(..., ProxyURLKey, "...")), + // which shadows the holder on pr but leaves the original chain (and + // our *string holder) reachable through origCtx. + origCtx := pr.Context() + proxyURL, err := f(pr) + if proxyURL != nil { + if h, _ := origCtx.Value(ProxyURLKey).(*string); h != nil { + *h = proxyURL.String() + } + } + return proxyURL, err + } + t, ok := c.backend.Client.Transport.(*http.Transport) if c.backend.Client.Transport != nil && ok { t.Proxy = p @@ -1340,6 +1363,9 @@ func (c *Collector) handleOnError(response *Response, err error, request *Reques if response.Ctx == nil { response.Ctx = request.Ctx } + if response.ProxyURL == "" { + response.ProxyURL = request.ProxyURL + } for _, f := range c.errorCallbacks { f(response, err) } diff --git a/proxy/proxy.go b/proxy/proxy.go index a4bd84852..a2de86355 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -15,7 +15,6 @@ package proxy import ( - "context" "net/http" "net/url" "sync/atomic" @@ -31,9 +30,9 @@ type roundRobinSwitcher struct { func (r *roundRobinSwitcher) GetProxy(pr *http.Request) (*url.URL, error) { index := atomic.AddUint32(&r.index, 1) - 1 u := r.proxyURLs[index%uint32(len(r.proxyURLs))] - - ctx := context.WithValue(pr.Context(), colly.ProxyURLKey, u.String()) - *pr = *pr.WithContext(ctx) + // SetProxyFunc wraps this and writes the chosen proxy URL through the + // *string holder in the request context, so GetProxy itself only needs + // to return the *url.URL. return u, nil } diff --git a/proxy/proxy_test.go b/proxy/proxy_test.go new file mode 100644 index 000000000..23e682614 --- /dev/null +++ b/proxy/proxy_test.go @@ -0,0 +1,195 @@ +// Copyright 2018 Adam Tauber +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package proxy + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "net/url" + "testing" + + "github.com/gocolly/colly/v2" +) + +// TestRoundRobinProxySwitcher_PropagatesProxyURL is the minimal smoke test: +// after a Visit through the switcher, the response must carry a non-empty +// ProxyURL on both Request and Response. +func TestRoundRobinProxySwitcher_PropagatesProxyURL(t *testing.T) { + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprintln(w, "ok") + })) + defer ts.Close() + + rp, err := RoundRobinProxySwitcher(ts.URL) + if err != nil { + t.Fatalf("RoundRobinProxySwitcher: %v", err) + } + + c := colly.NewCollector(colly.IgnoreRobotsTxt()) + c.SetProxyFunc(rp) + + var called bool + c.OnResponse(func(r *colly.Response) { + called = true + if r.Request.ProxyURL == "" { + t.Errorf("Request.ProxyURL is empty — ProxyURLKey not propagated") + } + if r.ProxyURL == "" { + t.Errorf("Response.ProxyURL is empty") + } + }) + + if err := c.Visit("http://example.com/"); err != nil { + t.Fatalf("Visit: %v", err) + } + if !called { + t.Fatal("OnResponse never fired") + } +} + +// TestRoundRobinProxySwitcher_ProxyURLOnError ensures the chosen proxy URL +// is still recorded when the request fails before any response headers +// arrive (e.g. dial refused) — so OnError can report which proxy was tried. +func TestRoundRobinProxySwitcher_ProxyURLOnError(t *testing.T) { + ln := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + dead := ln.URL + ln.Close() // guarantees dial refused on dead + + rp, err := RoundRobinProxySwitcher(dead) + if err != nil { + t.Fatalf("RoundRobinProxySwitcher: %v", err) + } + c := colly.NewCollector(colly.IgnoreRobotsTxt()) + c.SetProxyFunc(rp) + + var called bool + c.OnError(func(r *colly.Response, _ error) { + called = true + if r.Request.ProxyURL != dead { + t.Errorf("Request.ProxyURL = %q, want %q", r.Request.ProxyURL, dead) + } + if r.ProxyURL != dead { + t.Errorf("Response.ProxyURL = %q, want %q", r.ProxyURL, dead) + } + }) + + if err := c.Visit("http://example.com/"); err == nil { + t.Fatal("expected Visit to fail") + } + if !called { + t.Fatal("OnError never fired") + } +} + +// TestSetProxyFunc_LegacyContextStringPropagates documents the interaction +// between a custom ProxyFunc that follows the legacy "WithContext+string" +// pattern and the current SetProxyFunc wrapper. +// +// The user's *pr = *pr.WithContext(...) mutation only affects the fork that +// net/http.send() created (Client.Timeout triggers forkReq), so the string +// the user writes into ProxyURLKey is discarded along with the fork. What +// actually surfaces on Request.ProxyURL is the *url.URL the ProxyFunc +// returns, written by the wrapper through the *string holder colly placed +// in the (shared) context. To make this concrete the test has the user +// write a marker string that intentionally differs from the returned URL, +// then asserts the URL — not the marker — is what propagates. +func TestSetProxyFunc_LegacyContextStringPropagates(t *testing.T) { + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprintln(w, "ok") + })) + defer ts.Close() + + proxyURL, err := url.Parse(ts.URL) + if err != nil { + t.Fatalf("url.Parse: %v", err) + } + const userMarker = "user-wrote-this-but-it-should-be-ignored" + + c := colly.NewCollector(colly.IgnoreRobotsTxt()) + c.SetProxyFunc(func(pr *http.Request) (*url.URL, error) { + ctx := context.WithValue(pr.Context(), colly.ProxyURLKey, userMarker) + *pr = *pr.WithContext(ctx) + return proxyURL, nil + }) + + var called bool + c.OnResponse(func(r *colly.Response) { + called = true + if r.Request.ProxyURL != proxyURL.String() { + t.Errorf("Request.ProxyURL = %q, want %q (from returned *url.URL)", r.Request.ProxyURL, proxyURL.String()) + } + if r.ProxyURL != proxyURL.String() { + t.Errorf("Response.ProxyURL = %q, want %q", r.ProxyURL, proxyURL.String()) + } + if r.Request.ProxyURL == userMarker { + t.Errorf("Request.ProxyURL leaked the user marker %q — the WithContext+string write must be isolated by forkReq", userMarker) + } + }) + + if err := c.Visit("http://example.com/"); err != nil { + t.Fatalf("Visit: %v", err) + } + if !called { + t.Fatal("OnResponse never fired") + } +} + +// TestSetProxyFunc_LegacyContextStringOnError is the error-path counterpart: +// the same legacy WithContext+string ProxyFunc, but the proxy is a dead port +// so the request fails before any response headers. The returned *url.URL +// (not the user's discarded ctx string) must still be reflected on +// Request.ProxyURL / Response.ProxyURL — proving the *string holder write +// from SetProxyFunc's wrapper survives both forkReq and the error path. +func TestSetProxyFunc_LegacyContextStringOnError(t *testing.T) { + ln := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + dead := ln.URL + ln.Close() // guarantees dial refused on dead + + proxyURL, err := url.Parse(dead) + if err != nil { + t.Fatalf("url.Parse: %v", err) + } + const userMarker = "user-wrote-this-but-it-should-be-ignored" + + c := colly.NewCollector(colly.IgnoreRobotsTxt()) + c.SetProxyFunc(func(pr *http.Request) (*url.URL, error) { + ctx := context.WithValue(pr.Context(), colly.ProxyURLKey, userMarker) + *pr = *pr.WithContext(ctx) + return proxyURL, nil + }) + + var called bool + c.OnError(func(r *colly.Response, _ error) { + called = true + if r.Request.ProxyURL != dead { + t.Errorf("Request.ProxyURL = %q, want %q (from returned *url.URL)", r.Request.ProxyURL, dead) + } + if r.ProxyURL != dead { + t.Errorf("Response.ProxyURL = %q, want %q", r.ProxyURL, dead) + } + if r.Request.ProxyURL == userMarker { + t.Errorf("Request.ProxyURL leaked the user marker %q", userMarker) + } + }) + + if err := c.Visit("http://example.com/"); err == nil { + t.Fatal("expected Visit to fail") + } + if !called { + t.Fatal("OnError never fired") + } +} diff --git a/response.go b/response.go index 30cdeae66..eb2f121c8 100644 --- a/response.go +++ b/response.go @@ -42,6 +42,9 @@ type Response struct { // Trace contains the HTTPTrace for the request. Will only be set by the // collector if Collector.TraceHTTP is set to true. Trace *HTTPTrace + // ProxyURL is the proxy address that handled the request, mirrored from + // Request.ProxyURL for convenience. + ProxyURL string } // Save writes response body to disk