diff --git a/libcore/http.go b/libcore/http.go index 2e4a3bf6f4..52bc611a95 100644 --- a/libcore/http.go +++ b/libcore/http.go @@ -20,7 +20,6 @@ import ( "path/filepath" "strconv" "sync" - "sync/atomic" "time" "github.com/sagernet/quic-go" @@ -248,121 +247,206 @@ func (r *httpRequest) Execute() (HTTPResponse, error) { return httpResp, nil } -type requestFunc func() (response *http.Response, err error) +type requestFunc func(context.Context) (response *http.Response, err error) -func (r *httpRequest) doH3Direct() (HTTPResponse, error) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() +type labeledRequestFunc struct { + label string + request requestFunc +} - successCh := make(chan *http.Response, 1) +var errEmptyHTTPResponse = errors.New("empty response") + +type indexedHTTPResponse struct { + index int + response *http.Response +} + +type cancelOnCloseBody struct { + io.ReadCloser + cancel context.CancelFunc +} + +func (b *cancelOnCloseBody) Close() error { + err := b.ReadCloser.Close() + b.cancel() + return err +} + +func raceHTTPRequests(ctx context.Context, funcs []labeledRequestFunc) (*http.Response, error) { + successCh := make(chan indexedHTTPResponse) + doneCh := make(chan struct{}, len(funcs)) + workerCancels := make([]context.CancelFunc, len(funcs)) var finalErr error - var failedCount atomic.Uint32 - var successCount atomic.Uint32 var mu sync.Mutex - funcs := []requestFunc{ - // Http(s) With Ech - func() (response *http.Response, err error) { - request := r.request.Clone(context.Background()) - echClient := &http.Client{ - Timeout: defaultHTTPRequestTimeout, - Transport: &http.Transport{ - DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) { - var d net.Dialer - c, err := d.DialContext(ctx, network, addr) - if err != nil { - return c, err - } - domain := addr - if host, _, _ := net.SplitHostPort(addr); host != "" { - domain = host - } - echTls := ech.NewECHClientConfig(domain, &r.tls, gLocalDNSTransport) - return echTls.Client(ctx, c) - }, - DisableKeepAlives: true, - }, - } - return echClient.Do(request) - }, - // H3 HTTPS - func() (response *http.Response, err error) { - request := r.request.Clone(context.Background()) - h3Client := &http.Client{ - Timeout: defaultHTTPRequestTimeout, - Transport: &http3.Transport{ - TLSClientConfig: r.tls.Clone(), - QUICConfig: &quic.Config{ - MaxIdleTimeout: time.Second, - }, - }, - } - return h3Client.Do(request) - }, + addError := func(err error) { + mu.Lock() + finalErr = errors.Join(finalErr, err) + mu.Unlock() } - - if r.request.URL.Scheme == "http" { - funcs = funcs[:1] + resultError := func() error { + mu.Lock() + defer mu.Unlock() + if ctxErr := ctx.Err(); ctxErr != nil { + return errors.Join(finalErr, ctxErr) + } + if finalErr == nil { + return errEmptyHTTPResponse + } + return finalErr } - for i, f := range funcs { - go func(f requestFunc) { - defer device.DeferPanicToError("http", func(err error) { log.Println(err) }) + for index, f := range funcs { + workerCtx, workerCancel := context.WithCancel(context.WithoutCancel(ctx)) + workerCancels[index] = workerCancel + go func(index int, f labeledRequestFunc, requestCtx context.Context, cancel context.CancelFunc) { + cancelOnExit := true defer func() { - if successCount.Load() == 0 { - if failedCount.Add(1) >= uint32(len(funcs)) { - // all failed - cancel() - } + if cancelOnExit { + cancel() } + doneCh <- struct{}{} }() + defer device.DeferPanicToError("http", func(err error) { + addError(fmt.Errorf("%s: %w", f.label, err)) + log.Println(err) + }) - var t string - switch i { - case 0: - t = "http(s)" - case 1: - t = "h3" - } - - // execute the HTTP request - rsp, err := f() + rsp, err := f.request(requestCtx) if rsp == nil || err != nil { - mu.Lock() - finalErr = errors.Join(finalErr, fmt.Errorf("%s: %w", t, err)) - mu.Unlock() + if err == nil { + err = errEmptyHTTPResponse + } + addError(fmt.Errorf("%s: %w", f.label, err)) if rsp != nil && rsp.Body != nil { rsp.Body.Close() } return } - // handle the HTTP status code if rsp.StatusCode != http.StatusOK { hr := &httpResponse{Response: rsp} - err = fmt.Errorf("%s: %s", t, hr.errorString()) - mu.Lock() - finalErr = errors.Join(finalErr, err) - mu.Unlock() + addError(fmt.Errorf("%s: %s", f.label, hr.errorString())) return } select { - case successCh <- rsp: - // first successful request, don't close the body - successCount.Add(1) - default: - rsp.Body.Close() + case successCh <- indexedHTTPResponse{index: index, response: rsp}: + // The response body owns the winning request context until Close. + cancelOnExit = false + case <-requestCtx.Done(): + if rsp.Body != nil { + rsp.Body.Close() + } + } + }(index, f, workerCtx, workerCancel) + } + + completed := 0 + cancelAndWait := func(winner int) { + for index, cancel := range workerCancels { + if index != winner { + cancel() + } + } + for completed < len(funcs) { + <-doneCh + completed++ + } + } + + for completed < len(funcs) { + select { + case result := <-successCh: + if ctx.Err() != nil { + workerCancels[result.index]() + if result.response.Body != nil { + result.response.Body.Close() + } + cancelAndWait(-1) + return nil, resultError() } - }(f) + cancelAndWait(result.index) + if result.response.Body == nil { + workerCancels[result.index]() + } else { + result.response.Body = &cancelOnCloseBody{ + ReadCloser: result.response.Body, + cancel: workerCancels[result.index], + } + } + return result.response, nil + case <-doneCh: + completed++ + if completed == len(funcs) { + return nil, resultError() + } + case <-ctx.Done(): + cancelAndWait(-1) + return nil, resultError() + } + } + return nil, resultError() +} + +func (r *httpRequest) doH3Direct() (HTTPResponse, error) { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + funcs := []labeledRequestFunc{ + { + label: "http(s)", + request: func(ctx context.Context) (response *http.Response, err error) { + request := r.request.Clone(ctx) + echClient := &http.Client{ + Timeout: defaultHTTPRequestTimeout, + Transport: &http.Transport{ + DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) { + var d net.Dialer + c, err := d.DialContext(ctx, network, addr) + if err != nil { + return c, err + } + domain := addr + if host, _, _ := net.SplitHostPort(addr); host != "" { + domain = host + } + echTls := ech.NewECHClientConfig(domain, &r.tls, gLocalDNSTransport) + return echTls.Client(ctx, c) + }, + DisableKeepAlives: true, + }, + } + return echClient.Do(request) + }, + }, + { + label: "h3", + request: func(ctx context.Context) (response *http.Response, err error) { + request := r.request.Clone(ctx) + h3Client := &http.Client{ + Timeout: defaultHTTPRequestTimeout, + Transport: &http3.Transport{ + TLSClientConfig: r.tls.Clone(), + QUICConfig: &quic.Config{ + MaxIdleTimeout: time.Second, + }, + }, + } + return h3Client.Do(request) + }, + }, } - select { - case result := <-successCh: - return &httpResponse{Response: result}, nil - case <-ctx.Done(): - return nil, finalErr + if r.request.URL.Scheme == "http" { + funcs = funcs[:1] + } + + result, err := raceHTTPRequests(ctx, funcs) + if err != nil { + return nil, err } + return &httpResponse{Response: result}, nil } type httpResponse struct { diff --git a/libcore/http_test.go b/libcore/http_test.go new file mode 100644 index 0000000000..372cf2f054 --- /dev/null +++ b/libcore/http_test.go @@ -0,0 +1,365 @@ +package libcore + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" +) + +const raceTestTimeout = 2 * time.Second + +type trackingReadCloser struct { + reader io.Reader + closed chan struct{} + once sync.Once +} + +func newTrackingReadCloser(content string) *trackingReadCloser { + return &trackingReadCloser{ + reader: strings.NewReader(content), + closed: make(chan struct{}), + } +} + +func (r *trackingReadCloser) Read(p []byte) (int, error) { + return r.reader.Read(p) +} + +func (r *trackingReadCloser) Close() error { + r.once.Do(func() { close(r.closed) }) + return nil +} + +type httpRaceResult struct { + response *http.Response + err error +} + +func waitForSignal(t *testing.T, signal <-chan struct{}, description string) { + t.Helper() + select { + case <-signal: + case <-time.After(raceTestTimeout): + t.Fatalf("timed out waiting for %s", description) + } +} + +func waitForWorkers(t *testing.T, workers *sync.WaitGroup) { + t.Helper() + done := make(chan struct{}) + go func() { + workers.Wait() + close(done) + }() + waitForSignal(t, done, "request functions") +} + +func waitForRaceResult(t *testing.T, resultCh <-chan httpRaceResult) httpRaceResult { + t.Helper() + select { + case result := <-resultCh: + return result + case <-time.After(raceTestTimeout): + t.Fatal("timed out waiting for request race") + return httpRaceResult{} + } +} + +func assertBodyClosed(t *testing.T, body *trackingReadCloser) { + t.Helper() + waitForSignal(t, body.closed, "response body close") +} + +func TestRaceHTTPRequestsTimeoutCancelsWorkersAndClosesBodies(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond) + defer cancel() + + firstBody := newTrackingReadCloser("first") + secondBody := newTrackingReadCloser("second") + var workers sync.WaitGroup + workers.Add(2) + + blockedRequest := func(body *trackingReadCloser) requestFunc { + return func(requestCtx context.Context) (*http.Response, error) { + defer workers.Done() + <-requestCtx.Done() + return &http.Response{StatusCode: http.StatusOK, Body: body}, nil + } + } + + resultCh := make(chan httpRaceResult, 1) + go func() { + response, err := raceHTTPRequests(ctx, []labeledRequestFunc{ + {label: "first", request: blockedRequest(firstBody)}, + {label: "second", request: blockedRequest(secondBody)}, + }) + resultCh <- httpRaceResult{response: response, err: err} + }() + + result := waitForRaceResult(t, resultCh) + if result.response != nil { + t.Fatal("timeout returned a response") + } + if !errors.Is(result.err, context.DeadlineExceeded) { + t.Fatalf("timeout error = %v, want context deadline exceeded", result.err) + } + waitForWorkers(t, &workers) + assertBodyClosed(t, firstBody) + assertBodyClosed(t, secondBody) +} + +func TestRaceHTTPRequestsEmptyResponse(t *testing.T) { + response, err := raceHTTPRequests(context.Background(), []labeledRequestFunc{ + { + label: "empty", + request: func(context.Context) (*http.Response, error) { + return nil, nil + }, + }, + }) + + if response != nil { + t.Fatal("empty result returned a response") + } + if err == nil || !strings.Contains(err.Error(), "empty response") { + t.Fatalf("empty result error = %v, want empty response", err) + } +} + +func TestRaceHTTPRequestsWorkerPanicIsReturned(t *testing.T) { + response, err := raceHTTPRequests(context.Background(), []labeledRequestFunc{ + { + label: "panicked", + request: func(context.Context) (*http.Response, error) { + panic("worker panic") + }, + }, + }) + + if response != nil { + t.Fatal("panicked request returned a response") + } + if err == nil || !strings.Contains(err.Error(), "panicked: http panic: worker panic") { + t.Fatalf("panic error = %v, want labeled worker panic", err) + } +} + +func TestRaceHTTPRequestsFirstSuccessCancelsLoserAndKeepsWinnerReadable(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + winnerBody := newTrackingReadCloser("winner") + lateBody := newTrackingReadCloser("late") + lateStarted := make(chan struct{}) + lateReturned := make(chan struct{}) + var winnerRequestDone <-chan struct{} + + response, err := raceHTTPRequests(ctx, []labeledRequestFunc{ + { + label: "winner", + request: func(requestCtx context.Context) (*http.Response, error) { + winnerRequestDone = requestCtx.Done() + <-lateStarted + return &http.Response{StatusCode: http.StatusOK, Body: winnerBody}, nil + }, + }, + { + label: "late", + request: func(requestCtx context.Context) (*http.Response, error) { + close(lateStarted) + <-requestCtx.Done() + close(lateReturned) + return &http.Response{StatusCode: http.StatusOK, Body: lateBody}, nil + }, + }, + }) + if err != nil { + cancel() + t.Fatalf("request race failed: %v", err) + } + if response == nil { + cancel() + t.Fatal("request race returned no response") + } + + waitForSignal(t, lateReturned, "late request function") + assertBodyClosed(t, lateBody) + cancel() + + select { + case <-winnerBody.closed: + t.Fatal("winning response body was closed") + case <-winnerRequestDone: + t.Fatal("winning request context was cancelled before body close") + default: + } + + content, err := io.ReadAll(response.Body) + if err != nil { + t.Fatalf("read winning body: %v", err) + } + if string(content) != "winner" { + t.Fatalf("winning body = %q, want winner", content) + } + if err := response.Body.Close(); err != nil { + t.Fatalf("close winning body: %v", err) + } + waitForSignal(t, winnerRequestDone, "winning request cancellation") +} + +func TestRaceHTTPRequestsRealWinnerBodySurvivesParentCancellation(t *testing.T) { + loserStarted := make(chan struct{}) + loserStopped := make(chan struct{}) + releaseWinnerBody := make(chan struct{}) + var releaseOnce sync.Once + releaseWinner := func() { releaseOnce.Do(func() { close(releaseWinnerBody) }) } + + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/winner": + <-loserStarted + writer.WriteHeader(http.StatusOK) + writer.(http.Flusher).Flush() + <-releaseWinnerBody + _, _ = io.WriteString(writer, "winner") + case "/loser": + close(loserStarted) + <-request.Context().Done() + close(loserStopped) + default: + http.NotFound(writer, request) + } + })) + defer server.Close() + defer releaseWinner() + + ctx, cancel := context.WithCancel(context.Background()) + request := func(path string) requestFunc { + return func(requestCtx context.Context) (*http.Response, error) { + req, err := http.NewRequestWithContext(requestCtx, http.MethodGet, server.URL+path, nil) + if err != nil { + return nil, err + } + return server.Client().Do(req) + } + } + + response, err := raceHTTPRequests(ctx, []labeledRequestFunc{ + {label: "winner", request: request("/winner")}, + {label: "loser", request: request("/loser")}, + }) + if err != nil { + cancel() + t.Fatalf("request race failed: %v", err) + } + if response == nil { + cancel() + t.Fatal("request race returned no real response") + } + waitForSignal(t, loserStopped, "real losing request cancellation") + + cancel() + releaseWinner() + content, err := io.ReadAll(response.Body) + if err != nil { + t.Fatalf("read real winning body after parent cancellation: %v", err) + } + if string(content) != "winner" { + t.Fatalf("real winning body = %q, want winner", content) + } + if err := response.Body.Close(); err != nil { + t.Fatalf("close real winning body: %v", err) + } +} + +func TestRaceHTTPRequestsFailureAndDeadlineAreRetained(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond) + defer cancel() + + workerErr := errors.New("worker failed") + failureReturned := make(chan struct{}) + response, err := raceHTTPRequests(ctx, []labeledRequestFunc{ + { + label: "failed", + request: func(context.Context) (*http.Response, error) { + close(failureReturned) + return nil, workerErr + }, + }, + { + label: "blocked", + request: func(requestCtx context.Context) (*http.Response, error) { + <-failureReturned + <-requestCtx.Done() + return nil, requestCtx.Err() + }, + }, + }) + + if response != nil { + t.Fatal("failed request race returned a response") + } + if !errors.Is(err, workerErr) { + t.Fatalf("request race error = %v, want worker failure", err) + } + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("request race error = %v, want deadline exceeded", err) + } +} + +func TestRaceHTTPRequestsAllFailuresAreJoined(t *testing.T) { + firstErr := errors.New("first failure") + secondErr := errors.New("second failure") + response, err := raceHTTPRequests(context.Background(), []labeledRequestFunc{ + { + label: "first", + request: func(context.Context) (*http.Response, error) { + return nil, firstErr + }, + }, + { + label: "second", + request: func(context.Context) (*http.Response, error) { + return nil, secondErr + }, + }, + }) + + if response != nil { + t.Fatal("failed request race returned a response") + } + if err == nil { + t.Fatal("all-failed request race returned a nil error") + } + if !errors.Is(err, firstErr) || !errors.Is(err, secondErr) { + t.Fatalf("joined error = %v, want both worker errors", err) + } +} + +func TestRaceHTTPRequestsRejectsNonOKAndClosesBody(t *testing.T) { + body := newTrackingReadCloser("rejected") + response, err := raceHTTPRequests(context.Background(), []labeledRequestFunc{ + { + label: "http(s)", + request: func(context.Context) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusTeapot, + Status: "418 I'm a teapot", + Body: body, + }, nil + }, + }, + }) + + if response != nil { + t.Fatal("non-200 request returned a response") + } + if err == nil || !strings.Contains(err.Error(), "rejected") { + t.Fatalf("non-200 error = %v, want response body text", err) + } + assertBodyClosed(t, body) +}