diff --git a/docs/WEB_COVER.md b/docs/WEB_COVER.md index f714cb2..099d290 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -111,6 +111,12 @@ the upstream cannot be reached. Consequently, an upstream that requires an `Authorization` request header is not suitable without a separate authorized front end. +Response credential filtering also covers trailers, including fields that an +upstream adds only when its body ends. Ordinary end-to-end response trailers +remain available, and the body is still streamed rather than buffered in full. +This is defensive handling of upstream metadata, not an additional tunnel +authentication mechanism. + In web mode: - `--listen` is the UDP/H3 bind address and `--tcp-listen` is the TCP/H1/H2 diff --git a/internal/cover/handler.go b/internal/cover/handler.go index eb886d7..f069db1 100644 --- a/internal/cover/handler.go +++ b/internal/cover/handler.go @@ -13,6 +13,7 @@ import ( "os" "path/filepath" "strings" + "sync" ) var hopByHopHeaders = [...]string{ @@ -83,6 +84,11 @@ func NewReverseProxyHandler(origin *url.URL, transport http.RoundTripper) (http. }, ModifyResponse: func(response *http.Response) error { removeUnsafeHeaders(response.Header) + removeUnsafeHeaders(response.Trailer) + // An upgraded body is duplex, not an HTTP message with trailers. + if response.Body != nil && response.StatusCode != http.StatusSwitchingProtocols { + response.Body = &responseTrailerBody{body: response.Body, response: response} + } return nil }, ErrorHandler: func(w http.ResponseWriter, _ *http.Request, _ error) { @@ -93,6 +99,50 @@ func NewReverseProxyHandler(origin *url.URL, transport http.RoundTripper) (http. return proxy, nil } +// responseTrailerBody filters fields that a transport discovers only at EOF +// or Close, including replacement Trailer maps. It does not buffer the body. +// Trailer must not be inspected while Read is in progress. Close may interrupt +// that Read, so neither I/O operation holds mu: defer cleanup until concurrent +// operations have returned rather than blocking Close on the reader. +type responseTrailerBody struct { + body io.ReadCloser + response *http.Response + mu sync.Mutex + active int + pending bool +} + +func (b *responseTrailerBody) Read(p []byte) (int, error) { + b.beginOperation() + n, err := b.body.Read(p) + b.finishOperation(err != nil) + return n, err +} + +func (b *responseTrailerBody) Close() error { + b.beginOperation() + err := b.body.Close() + b.finishOperation(true) + return err +} + +func (b *responseTrailerBody) beginOperation() { + b.mu.Lock() + b.active++ + b.mu.Unlock() +} + +func (b *responseTrailerBody) finishOperation(terminal bool) { + b.mu.Lock() + defer b.mu.Unlock() + b.active-- + b.pending = b.pending || terminal + if b.active == 0 && b.pending { + removeUnsafeHeaders(b.response.Trailer) + b.pending = false + } +} + func normalizeOrigin(origin *url.URL) (*url.URL, error) { if origin == nil { return nil, errors.New("reverse proxy origin is required") diff --git a/internal/cover/response_cancel_test.go b/internal/cover/response_cancel_test.go new file mode 100644 index 0000000..b704df7 --- /dev/null +++ b/internal/cover/response_cancel_test.go @@ -0,0 +1,93 @@ +package cover + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "net/url" + "sync" + "testing" + "time" +) + +func TestReverseProxyResponseBodyCancellationReachesOrigin(t *testing.T) { + canceled := make(chan struct{}) + origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Trailer", "X-End") + _, _ = io.WriteString(w, "x") + w.(http.Flusher).Flush() + <-r.Context().Done() + close(canceled) + })) + t.Cleanup(origin.Close) + originURL, err := url.Parse(origin.URL) + if err != nil { + t.Fatal(err) + } + upstream := &http.Transport{Proxy: nil} + t.Cleanup(upstream.CloseIdleConnections) + bodyClosed := make(chan struct{}) + var closeOnce sync.Once + handler, err := NewReverseProxyHandler(originURL, roundTripFunc(func(r *http.Request) (*http.Response, error) { + response, err := upstream.RoundTrip(r) + if err != nil { + return nil, err + } + body := response.Body + response.Body = &trailerTestBody{read: body.Read, close: func() error { + err := body.Close() + closeOnce.Do(func() { close(bodyClosed) }) + return err + }} + return response, nil + })) + if err != nil { + t.Fatal(err) + } + frontDone := make(chan struct{}) + front := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer close(frontDone) + handler.ServeHTTP(w, r) + })) + t.Cleanup(front.Close) + transport := &http.Transport{Proxy: nil} + t.Cleanup(transport.CloseIdleConnections) + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + t.Cleanup(cancel) + request, err := http.NewRequestWithContext(ctx, http.MethodGet, front.URL, nil) + if err != nil { + t.Fatal(err) + } + response, err := (&http.Client{Transport: transport}).Do(request) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = response.Body.Close() }) + first := make([]byte, 1) + if _, err := io.ReadFull(response.Body, first); err != nil || first[0] != 'x' { + t.Fatalf("initial streamed byte = %q, error %v", first, err) + } + readDone := make(chan error, 1) + go func() { _, err := response.Body.Read(make([]byte, 1)); readDone <- err }() + cancel() + select { + case err := <-readDone: + if err == nil { + t.Fatal("body read succeeded after cancellation") + } + case <-time.After(time.Second): + t.Fatal("response body read did not stop after cancellation") + } + for name, done := range map[string]<-chan struct{}{ + "upstream request": canceled, + "upstream body": bodyClosed, + "front handler": frontDone, + } { + select { + case <-done: + case <-time.After(time.Second): + t.Fatalf("client cancellation did not release %s", name) + } + } +} diff --git a/internal/cover/response_trailer_test.go b/internal/cover/response_trailer_test.go new file mode 100644 index 0000000..84acb22 --- /dev/null +++ b/internal/cover/response_trailer_test.go @@ -0,0 +1,245 @@ +package cover + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "net/http/httputil" + "net/url" + "testing" + "time" +) + +func TestReverseProxyScrubsResponseTrailersOnWire(t *testing.T) { + for _, announced := range []bool{true, false} { + t.Run(fmt.Sprintf("credentials_announced_%t", announced), func(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = listener.Close() }) + done := make(chan error, 1) + go func() { + conn, err := listener.Accept() + if err != nil { + done <- err + return + } + defer conn.Close() + _ = conn.SetDeadline(time.Now().Add(2 * time.Second)) + request, err := http.ReadRequest(bufio.NewReader(conn)) + if err != nil { + done <- err + return + } + _, err = io.Copy(io.Discard, request.Body) + _ = request.Body.Close() + if err != nil { + done <- err + return + } + trailers := "X-End" + if announced { + trailers += ", Authorization, Proxy-Authorization" + } + // The upstream deliberately supplies prohibited credential trailer + // fields. Only fictional values and loopback sockets are used. + _, err = fmt.Fprintf(conn, "HTTP/1.1 200 OK\r\nConnection: close\r\nTransfer-Encoding: chunked\r\nTrailer: %s\r\nAuthorization: fictional-initial-origin\r\nProxy-Authorization: fictional-initial-proxy\r\n\r\n1\r\nx\r\n0\r\nAuthorization: fictional-late-origin\r\nProxy-Authorization: fictional-late-proxy\r\nX-End: retained\r\n\r\n", trailers) + done <- err + }() + origin, err := url.Parse("http://" + listener.Addr().String()) + if err != nil { + t.Fatal(err) + } + transport := &http.Transport{Proxy: nil} + t.Cleanup(transport.CloseIdleConnections) + handler, err := NewReverseProxyHandler(origin, transport) + if err != nil { + t.Fatal(err) + } + front := httptest.NewServer(handler) + t.Cleanup(front.Close) + clientTransport := &http.Transport{Proxy: nil} + t.Cleanup(clientTransport.CloseIdleConnections) + client := &http.Client{Transport: clientTransport, Timeout: 2 * time.Second} + response, err := client.Get(front.URL) + if err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(response.Body) + _ = response.Body.Close() + if err != nil { + t.Fatal(err) + } + if response.StatusCode != http.StatusOK || string(body) != "x" { + t.Fatalf("status/body = %d/%q", response.StatusCode, body) + } + for _, name := range []string{"Authorization", "Proxy-Authorization"} { + if response.Header.Get(name) != "" || response.Trailer.Get(name) != "" { + t.Fatalf("%s escaped response filtering: headers=%v trailers=%v", name, response.Header, response.Trailer) + } + } + if response.Trailer.Get("X-End") != "retained" { + t.Fatalf("ordinary trailer was lost: %v", response.Trailer) + } + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(2 * time.Second): + t.Fatal("origin did not finish") + } + }) + } +} + +func TestResponseTrailerBodyPreservesStreamingAndTerminalRead(t *testing.T) { + readFailure := errors.New("upstream read failure") + for _, terminal := range []error{io.EOF, readFailure, context.Canceled} { + t.Run(terminal.Error(), func(t *testing.T) { + response := &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Trailer: http.Header{ + "Authorization": nil, "Proxy-Authorization": nil, "X-End": nil, + }} + reads := 0 + response.Body = &trailerTestBody{read: func(p []byte) (int, error) { + reads++ + if reads == 1 { + return copy(p, "first"), nil + } + // Replacement, not just mutation, catches a wrapper holding the + // old map instead of following the response's final Trailer map. + response.Trailer = http.Header{ + "Authorization": {"fictional"}, "Proxy-Authorization": {"fictional"}, + "Connection": {"X-Hop"}, "X-Hop": {"remove"}, "X-End": {"retained"}, + } + return copy(p, "last"), terminal + }} + modifyTrailerResponse(t, response) + if reads != 0 { + t.Fatal("response modification read or buffered the body") + } + if _, present := response.Trailer["Authorization"]; present { + t.Fatal("unsafe trailer declaration survived initial filtering") + } + if _, present := response.Trailer["X-End"]; !present { + t.Fatal("ordinary trailer declaration was removed") + } + buffer := make([]byte, 16) + n, err := response.Body.Read(buffer) + if n != 5 || err != nil || string(buffer[:n]) != "first" || reads != 1 { + t.Fatalf("first streaming Read = %d/%v/%q, calls=%d", n, err, buffer[:n], reads) + } + n, err = response.Body.Read(buffer) + if n != 4 || err != terminal || string(buffer[:n]) != "last" { + t.Fatalf("terminal Read = %d/%v/%q, want original n+error", n, err, buffer[:n]) + } + assertFilteredTrailers(t, response.Trailer) + }) + } +} + +func TestResponseTrailerBodyPreservesCloseResultAndLateMap(t *testing.T) { + closeFailure := errors.New("upstream close failure") + response := &http.Response{StatusCode: http.StatusOK} + closes := 0 + response.Body = &trailerTestBody{close: func() error { + closes++ + response.Trailer = http.Header{"Authorization": {"fictional"}, "X-End": {"retained"}} + return closeFailure + }} + modifyTrailerResponse(t, response) + for wantCalls := 1; wantCalls <= 2; wantCalls++ { + if err := response.Body.Close(); err != closeFailure || closes != wantCalls { + t.Fatalf("Close result/calls = %v/%d, want original error and %d", err, closes, wantCalls) + } + assertFilteredTrailers(t, response.Trailer) + } +} + +func TestResponseTrailerBodyCloseDoesNotWaitForRead(t *testing.T) { + started := make(chan struct{}) + releaseRead := make(chan struct{}) + defer close(releaseRead) + response := &http.Response{StatusCode: http.StatusOK} + response.Body = &trailerTestBody{read: func([]byte) (int, error) { + close(started) + <-releaseRead + response.Trailer = http.Header{"Authorization": {"fictional"}, "X-End": {"retained"}} + return 0, context.Canceled + }} + modifyTrailerResponse(t, response) + done := make(chan error, 1) + go func() { _, err := response.Body.Read(make([]byte, 1)); done <- err }() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("body Read did not begin") + } + closed := make(chan error, 1) + go func() { closed <- response.Body.Close() }() + select { + case err := <-closed: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("Close waited for the reader") + } + // Close has returned, but the response Trailer is not yet safe to inspect. + // Release the read without closing the channel, leaving cleanup idempotent. + releaseRead <- struct{}{} + select { + case err := <-done: + if err != context.Canceled { + t.Fatalf("Read error = %v, want original cancellation", err) + } + case <-time.After(time.Second): + t.Fatal("body Read did not finish") + } + assertFilteredTrailers(t, response.Trailer) +} + +func modifyTrailerResponse(t *testing.T, response *http.Response) { + t.Helper() + handler, err := NewReverseProxyHandler(&url.URL{Scheme: "http", Host: "origin.invalid"}, roundTripFunc(func(*http.Request) (*http.Response, error) { + return nil, errors.New("unexpected network call") + })) + if err != nil { + t.Fatal(err) + } + if err := handler.(*httputil.ReverseProxy).ModifyResponse(response); err != nil { + t.Fatal(err) + } +} + +func assertFilteredTrailers(t *testing.T, trailers http.Header) { + t.Helper() + if len(trailers) != 1 || trailers.Get("X-End") != "retained" { + t.Fatalf("filtered trailers = %v, want only X-End retained", trailers) + } +} + +type trailerTestBody struct { + read func([]byte) (int, error) + close func() error +} + +func (b *trailerTestBody) Read(p []byte) (int, error) { + if b.read != nil { + return b.read(p) + } + return 0, io.EOF +} + +func (b *trailerTestBody) Close() error { + if b.close != nil { + return b.close() + } + return nil +}