From ac261cc34db3177990f69bc4e853c06520cef684 Mon Sep 17 00:00:00 2001 From: abhinavgautam01 Date: Sat, 26 Sep 2026 15:44:28 +0530 Subject: [PATCH] fix: centralize pass-through response relaying --- internal/handler/composer.go | 10 +- internal/handler/conan.go | 11 +- internal/handler/conda.go | 8 +- internal/handler/container.go | 18 +- internal/handler/container_manifest.go | 8 +- internal/handler/container_tags.go | 5 +- internal/handler/gem.go | 28 +-- internal/handler/handler.go | 38 +--- internal/handler/hex.go | 8 +- internal/handler/nuget.go | 10 +- internal/handler/pypi.go | 14 +- internal/handler/relay.go | 110 ++++++++++ internal/handler/relay_test.go | 273 +++++++++++++++++++++++++ internal/handler/swift.go | 28 ++- internal/server/relay_test.go | 102 +++++++++ internal/server/server.go | 6 + 16 files changed, 538 insertions(+), 139 deletions(-) create mode 100644 internal/handler/relay.go create mode 100644 internal/handler/relay_test.go create mode 100644 internal/server/relay_test.go diff --git a/internal/handler/composer.go b/internal/handler/composer.go index fbbd7a4e..47378f79 100644 --- a/internal/handler/composer.go +++ b/internal/handler/composer.go @@ -5,7 +5,6 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "path" "strings" @@ -491,12 +490,5 @@ func (h *ComposerHandler) proxyUpstream(w http.ResponseWriter, r *http.Request) } defer func() { _ = resp.Body.Close() }() - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, nil) } diff --git a/internal/handler/conan.go b/internal/handler/conan.go index 7142f0dd..862f48c1 100644 --- a/internal/handler/conan.go +++ b/internal/handler/conan.go @@ -2,7 +2,6 @@ package handler import ( "fmt" - "io" "net/http" "strings" ) @@ -194,13 +193,5 @@ func (h *ConanHandler) proxyUpstream(w http.ResponseWriter, r *http.Request) { } defer func() { _ = resp.Body.Close() }() - // Copy response headers - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, nil) } diff --git a/internal/handler/conda.go b/internal/handler/conda.go index 91f72663..ad01f167 100644 --- a/internal/handler/conda.go +++ b/internal/handler/conda.go @@ -156,13 +156,7 @@ func (h *CondaHandler) handleRepodata(w http.ResponseWriter, r *http.Request) { defer func() { _ = resp.Body.Close() }() if resp.StatusCode != http.StatusOK { - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, nil) return } diff --git a/internal/handler/container.go b/internal/handler/container.go index 4bd677bb..5a47d795 100644 --- a/internal/handler/container.go +++ b/internal/handler/container.go @@ -274,16 +274,16 @@ func (h *ContainerHandler) proxyBlobHead(w http.ResponseWriter, r *http.Request, } defer func() { _ = resp.Body.Close() }() - for _, header := range []string{headerContentType, headerContentLength, "Docker-Content-Digest", headerETag, headerLastModified} { - if v := resp.Header.Get(header); v != "" { - w.Header().Set(header, v) + h.proxy.relayResponse(w, r, resp, func(dst, src http.Header) { + for _, header := range []string{headerContentType, headerContentLength, "Docker-Content-Digest", headerETag, headerLastModified} { + if v := src.Get(header); v != "" { + dst.Set(header, v) + } } - } - if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices && w.Header().Get("Docker-Content-Digest") == "" { - w.Header().Set("Docker-Content-Digest", digest) - } - - w.WriteHeader(resp.StatusCode) + if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices && dst.Get("Docker-Content-Digest") == "" { + dst.Set("Docker-Content-Digest", digest) + } + }) } // registryForName resolves a client-visible OCI repository name to an upstream diff --git a/internal/handler/container_manifest.go b/internal/handler/container_manifest.go index 7b030868..01df3afb 100644 --- a/internal/handler/container_manifest.go +++ b/internal/handler/container_manifest.go @@ -5,7 +5,6 @@ import ( "crypto/sha256" "encoding/hex" "fmt" - "io" "mime" "net/http" "regexp" @@ -77,15 +76,12 @@ func (h *ContainerHandler) serveManifest(w http.ResponseWriter, r *http.Request, writeContainerManifest(w, r, cached, true) return } - copyContainerManifestHeaders(w.Header(), resp.Header) - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, copyContainerManifestHeaders) return } if r.Method == http.MethodHead { - copyContainerManifestHeaders(w.Header(), resp.Header) - w.WriteHeader(http.StatusOK) + h.proxy.relayResponse(w, r, resp, copyContainerManifestHeaders) return } diff --git a/internal/handler/container_tags.go b/internal/handler/container_tags.go index dcc08f81..e9803ed1 100644 --- a/internal/handler/container_tags.go +++ b/internal/handler/container_tags.go @@ -5,7 +5,6 @@ import ( "crypto/sha256" "encoding/hex" "fmt" - "io" "net/http" "net/url" "regexp" @@ -73,9 +72,7 @@ func (h *ContainerHandler) serveTagsList(w http.ResponseWriter, r *http.Request, writeContainerTags(w, cached, true) return } - copyContainerTagsHeaders(w.Header(), resp.Header) - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, copyContainerTagsHeaders) return } diff --git a/internal/handler/gem.go b/internal/handler/gem.go index 8fc50399..efe964b6 100644 --- a/internal/handler/gem.go +++ b/internal/handler/gem.go @@ -4,7 +4,6 @@ import ( "bufio" "encoding/json" "fmt" - "io" "net/http" "strings" "time" @@ -117,17 +116,13 @@ func (h *GemHandler) handleCompactIndex(w http.ResponseWriter, r *http.Request) defer func() { _ = indexResp.Body.Close() }() if indexResp.StatusCode != http.StatusOK { - copyResponseHeaders(w, indexResp.Header) - w.WriteHeader(indexResp.StatusCode) - _, _ = io.Copy(w, indexResp.Body) + h.proxy.relayResponse(w, r, indexResp, nil) return } if filteredVersions == nil { h.proxy.Logger.Warn("failed to fetch version timestamps, proxying unfiltered", "name", name) - copyResponseHeaders(w, indexResp.Header) - w.WriteHeader(http.StatusOK) - _, _ = io.Copy(w, indexResp.Body) + h.proxy.relayResponse(w, r, indexResp, nil) return } @@ -215,15 +210,6 @@ func (h *GemHandler) writeFilteredIndex(w http.ResponseWriter, resp *http.Respon } } -// copyResponseHeaders copies HTTP headers from a response to a writer. -func copyResponseHeaders(w http.ResponseWriter, headers http.Header) { - for k, vv := range headers { - for _, v := range vv { - w.Header().Add(k, v) - } - } -} - // gemVersion represents a version entry from the RubyGems versions API. type gemVersion struct { Number string `json:"number"` @@ -314,15 +300,7 @@ func (h *GemHandler) proxyUpstream(w http.ResponseWriter, r *http.Request) { } defer func() { _ = resp.Body.Close() }() - // Copy response headers - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, nil) } func init() { diff --git a/internal/handler/handler.go b/internal/handler/handler.go index 36e6f16d..8040b47c 100644 --- a/internal/handler/handler.go +++ b/internal/handler/handler.go @@ -768,7 +768,8 @@ func serveArtifact(w http.ResponseWriter, method string, result *CacheResult) { // ProxyUpstream forwards a request to an upstream URL without caching. // It copies the request, forwards specified headers, and streams the response back. -// If forwardHeaders is nil, all response headers are copied. +// forwardHeaders controls the request headers sent upstream. End-to-end response +// headers and trailers are relayed independently of that list. func (p *Proxy) ProxyUpstream(w http.ResponseWriter, r *http.Request, upstreamURL string, forwardHeaders []string) { p.Logger.Debug("proxying to upstream", "url", upstreamURL) @@ -794,17 +795,10 @@ func (p *Proxy) ProxyUpstream(w http.ResponseWriter, r *http.Request, upstreamUR } defer func() { _ = resp.Body.Close() }() - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + p.relayResponse(w, r, resp, nil) } -// ProxyFile forwards a file request to upstream, copying all response headers. +// ProxyFile forwards a file request, relaying end-to-end headers and trailers. func (p *Proxy) ProxyFile(w http.ResponseWriter, r *http.Request, upstreamURL string) { req, err := http.NewRequestWithContext(r.Context(), r.Method, upstreamURL, nil) if err != nil { @@ -820,14 +814,7 @@ func (p *Proxy) ProxyFile(w http.ResponseWriter, r *http.Request, upstreamURL st } defer func() { _ = resp.Body.Close() }() - for key, values := range resp.Header { - for _, v := range values { - w.Header().Add(key, v) - } - } - - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + p.relayResponse(w, r, resp, nil) } // JSONError writes a JSON error response. @@ -1273,16 +1260,13 @@ func (p *Proxy) proxyMetadataStream(w http.ResponseWriter, r *http.Request, upst } defer func() { _ = resp.Body.Close() }() - for _, header := range []string{headerContentType, headerContentLength, headerContentEncoding, headerLastModified, headerETag} { - if v := resp.Header.Get(header); v != "" { - w.Header().Set(header, v) + p.relayResponse(w, r, resp, func(dst, src http.Header) { + for _, header := range []string{headerContentType, headerContentLength, headerContentEncoding, headerLastModified, headerETag} { + if v := src.Get(header); v != "" { + dst.Set(header, v) + } } - } - - w.WriteHeader(resp.StatusCode) - if r.Method != http.MethodHead { - _, _ = io.Copy(w, resp.Body) - } + }) } func (p *Proxy) applyUpstreamAuth(req *http.Request) { diff --git a/internal/handler/hex.go b/internal/handler/hex.go index 15d28915..c475e5a9 100644 --- a/internal/handler/hex.go +++ b/internal/handler/hex.go @@ -117,13 +117,7 @@ func (h *HexHandler) handlePackages(w http.ResponseWriter, r *http.Request) { defer func() { _ = protoResp.Body.Close() }() if protoResp.StatusCode != http.StatusOK { - for k, vv := range protoResp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - w.WriteHeader(protoResp.StatusCode) - _, _ = io.Copy(w, protoResp.Body) + h.proxy.relayResponse(w, r, protoResp, nil) return } diff --git a/internal/handler/nuget.go b/internal/handler/nuget.go index 3216f23d..d37e3910 100644 --- a/internal/handler/nuget.go +++ b/internal/handler/nuget.go @@ -4,7 +4,6 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "strings" ) @@ -222,14 +221,7 @@ func (h *NuGetHandler) proxyUpstream(w http.ResponseWriter, r *http.Request) { } defer func() { _ = resp.Body.Close() }() - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, nil) } // buildUpstreamURL constructs the upstream URL for a request. diff --git a/internal/handler/pypi.go b/internal/handler/pypi.go index 7b050cf5..21d98051 100644 --- a/internal/handler/pypi.go +++ b/internal/handler/pypi.go @@ -6,7 +6,6 @@ import ( "encoding/json" "errors" "fmt" - "io" "mime" "net/http" "regexp" @@ -817,13 +816,8 @@ func (h *PyPIHandler) proxySimple(w http.ResponseWriter, r *http.Request, path s } defer func() { _ = resp.Body.Close() }() - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } - } - ensureVaryAccept(w.Header()) - - w.WriteHeader(resp.StatusCode) - _, _ = io.Copy(w, resp.Body) + h.proxy.relayResponse(w, r, resp, func(dst, src http.Header) { + copyRelayHeaders(dst, src) + ensureVaryAccept(dst) + }) } diff --git a/internal/handler/relay.go b/internal/handler/relay.go new file mode 100644 index 00000000..75b859a6 --- /dev/null +++ b/internal/handler/relay.go @@ -0,0 +1,110 @@ +package handler + +import ( + "io" + "net/http" + "strings" + + "golang.org/x/net/http/httpguts" +) + +// relayResponse forwards an unmodified upstream body. Callers still own and +// close resp.Body, including when a failed transfer aborts the handler. A custom +// header copier may select or rewrite end-to-end headers; it only sees headers +// after hop-by-hop fields have been removed. +func (p *Proxy) relayResponse(w http.ResponseWriter, r *http.Request, resp *http.Response, copyHeaders func(http.Header, http.Header)) { + blocked := relayHopHeaders(resp.Header) + headers := resp.Header.Clone() + for name := range blocked { + headers.Del(name) + } + bodyAllowed := r.Method != http.MethodHead && resp.StatusCode >= http.StatusOK && + resp.StatusCode != http.StatusNoContent && resp.StatusCode != http.StatusResetContent && resp.StatusCode != http.StatusNotModified + if resp.StatusCode < http.StatusOK || resp.StatusCode == http.StatusNoContent { + headers.Del(headerContentLength) + } else if resp.StatusCode == http.StatusResetContent { + headers.Set(headerContentLength, "0") + } + if copyHeaders == nil { + copyHeaders = copyRelayHeaders + } + copyHeaders(w.Header(), headers) + + announced := make(map[string]bool) + if bodyAllowed { + for name := range resp.Trailer { + name = http.CanonicalHeaderKey(name) + if validRelayTrailer(name, blocked) { + announced[name] = true + w.Header().Add("Trailer", name) + } + } + if len(announced) > 0 { + // HTTP/1 trailers require chunking, not a fixed Content-Length. + w.Header().Del(headerContentLength) + } + } + w.WriteHeader(resp.StatusCode) + if !bodyAllowed || resp.Body == nil { + return + } + + written, err := io.Copy(w, resp.Body) + if err != nil { + upstreamURL := "" + if resp.Request != nil && resp.Request.URL != nil { + upstreamURL = resp.Request.URL.Redacted() + } + p.Logger.Warn("upstream response relay failed", "url", upstreamURL, + "status", resp.StatusCode, "bytes", written, "error", err) + // Headers are already committed: an error page would turn a truncated + // download into a seemingly successful response. Let net/http close the + // HTTP/1 connection or reset the HTTP/2 stream instead. + panic(http.ErrAbortHandler) + } + relayTrailers(w, resp.Trailer, announced, blocked) +} + +func copyRelayHeaders(dst, src http.Header) { + for name, values := range src { + for _, value := range values { + dst.Add(name, value) + } + } +} + +func relayHopHeaders(headers http.Header) map[string]bool { + blocked := map[string]bool{ + "Connection": true, "Proxy-Connection": true, "Keep-Alive": true, + "Proxy-Authenticate": true, "Proxy-Authorization": true, "Te": true, + "Trailer": true, "Transfer-Encoding": true, "Upgrade": true, + } + for _, value := range headers.Values("Connection") { + for _, name := range strings.Split(value, ",") { + if name = strings.TrimSpace(name); name != "" { + blocked[http.CanonicalHeaderKey(name)] = true + } + } + } + return blocked +} + +func validRelayTrailer(name string, blocked map[string]bool) bool { + return !blocked[name] && httpguts.ValidHeaderFieldName(name) && httpguts.ValidTrailerHeader(name) +} + +func relayTrailers(w http.ResponseWriter, trailers http.Header, announced, blocked map[string]bool) { + for name, values := range trailers { + name = http.CanonicalHeaderKey(name) + if !validRelayTrailer(name, blocked) { + continue + } + if !announced[name] { + // A trailer discovered only at EOF must not become a regular header + // or trigger automatic Content-Length on a short buffered response. + _ = http.NewResponseController(w).Flush() + name = http.TrailerPrefix + name + } + w.Header()[name] = append([]string(nil), values...) + } +} diff --git a/internal/handler/relay_test.go b/internal/handler/relay_test.go new file mode 100644 index 00000000..23f2cba0 --- /dev/null +++ b/internal/handler/relay_test.go @@ -0,0 +1,273 @@ +package handler + +import ( + "bytes" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" + "testing/iotest" + + "github.com/git-pkgs/cooldown" + "github.com/go-chi/chi/v5/middleware" +) + +// Exercise the actual routes as well as the two shared entry points. Error +// responses also reach relay branches in handlers that normally parse metadata. +func relayTestRoutes(proxy *Proxy, upstream string) http.Handler { + mux := http.NewServeMux() + mux.HandleFunc("/upstream", func(w http.ResponseWriter, r *http.Request) { + proxy.ProxyUpstream(w, r, upstream, nil) + }) + mux.HandleFunc("/file", func(w http.ResponseWriter, r *http.Request) { + proxy.ProxyFile(w, r, upstream) + }) + mux.HandleFunc("/metadata", func(w http.ResponseWriter, r *http.Request) { + proxy.ProxyCached(w, r, upstream, "test", "index") + }) + const proxyURL = "http://proxy.local" + mount := func(prefix string, h http.Handler) { + mux.Handle(prefix+"/", http.StripPrefix(prefix, h)) + } + mount("/nuget", NewNuGetHandlerWithUpstreams(proxy, proxyURL, upstream, upstream).Routes()) + mount("/conan", NewConanHandlerWithUpstream(proxy, proxyURL, upstream).Routes()) + mount("/composer", NewComposerHandlerWithUpstreams(proxy, proxyURL, upstream, upstream).Routes()) + mount("/pypi", NewPyPIHandlerWithUpstreams(proxy, proxyURL, upstream, upstream).Routes()) + mount("/gem", NewGemHandlerWithUpstream(proxy, proxyURL, upstream).Routes()) + mount("/conda", NewCondaHandlerWithUpstream(proxy, proxyURL, upstream).Routes()) + mount("/hex", NewHexHandlerWithUpstreams(proxy, proxyURL, upstream, upstream).Routes()) + mount("/swift", NewSwiftHandler(proxy, proxyURL, upstream).Routes()) + mount("/v2", NewContainerHandlerWithRegistry(proxy, proxyURL, upstream).Routes()) + // Match the production recovery middleware: it must not swallow the abort. + return middleware.Recoverer(mux) +} + +func TestRelayRoutes(t *testing.T) { + for _, route := range []string{ + "/upstream", "/file", "/metadata", "/nuget/query", + "/conan/v2/files/demo/1.0/user/stable/rev/recipe/other.txt", + "/composer/search.json", "/pypi/simple/", "/gem/api/v1/dependencies", + "/gem/info/demo", "/conda/conda-forge/noarch/repodata.json", + "/hex/packages/demo", "/swift/scope/demo/1.0.0/Package.swift", + "/v2/library/demo/manifests/latest", "/v2/library/demo/tags/list", + } { + t.Run(route, func(t *testing.T) { + for _, truncated := range []bool{false, true} { + t.Run(fmt.Sprintf("truncated=%v", truncated), func(t *testing.T) { + testRelayRoute(t, route, truncated) + }) + } + }) + } +} + +func testRelayRoute(t *testing.T, route string, truncated bool) { + t.Helper() + // Large enough to commit downstream headers before a body-copy failure. + body := strings.Repeat("x", 64*1024) + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + conn, rw, err := w.(http.Hijacker).Hijack() + if err != nil { + t.Error(err) + return + } + defer func() { _ = conn.Close() }() + _, _ = fmt.Fprintf(rw, "HTTP/1.1 502 Bad Gateway\r\nContent-Type: text/plain\r\nConnection: X-Private\r\nX-Private: secret\r\nTransfer-Encoding: chunked\r\nTrailer: X-Checksum, X-Private\r\n\r\n%x\r\n%s\r\n", len(body), body) + if !truncated { + _, _ = fmt.Fprint(rw, "0\r\nX-Checksum: verified\r\nX-Late: discovered-at-eof\r\nX-Private: still-secret\r\n\r\n") + } + if err := rw.Flush(); err != nil { + t.Error(err) + } + })) + defer upstream.Close() + proxy := testProxy() + proxy.Logger = slog.New(slog.NewTextHandler(io.Discard, nil)) + proxy.HTTPClient = upstream.Client() + proxy.Cooldown = &cooldown.Config{Default: "3d"} + downstream := httptest.NewServer(relayTestRoutes(proxy, upstream.URL)) + defer downstream.Close() + + resp, err := downstream.Client().Get(downstream.URL + route) + if err != nil { + t.Fatal(err) + } + defer func() { _ = resp.Body.Close() }() + got, readErr := io.ReadAll(resp.Body) + if resp.StatusCode != http.StatusBadGateway { + t.Fatalf("status = %d, want 502", resp.StatusCode) + } + if resp.Header.Get("Connection") != "" || resp.Header.Get("X-Private") != "" || resp.Trailer.Get("X-Private") != "" { + t.Fatalf("connection-scoped fields leaked: headers=%v trailers=%v", resp.Header, resp.Trailer) + } + if truncated { + if readErr == nil { + t.Fatal("truncated upstream was delivered as a complete response") + } + return + } + if readErr != nil || string(got) != body { + t.Fatalf("body length = %d, want %d; error = %v", len(got), len(body), readErr) + } + if resp.Trailer.Get("X-Checksum") != "verified" || resp.Trailer.Get("X-Late") != "discovered-at-eof" { + t.Fatalf("trailers not relayed: %v", resp.Trailer) + } + if route == "/pypi/simple/" && !strings.Contains(resp.Header.Get("Vary"), "Accept") { + t.Error("PyPI lost Vary: Accept") + } +} + +func TestRelayResponseHeaders(t *testing.T) { + headers := http.Header{ + "Connection": {"keep-alive, x-private", " X-Second, Content-Length "}, + "X-Private": {"secret"}, "X-Second": {"secret"}, + "Keep-Alive": {"timeout=5"}, "Proxy-Connection": {"keep-alive"}, + "Proxy-Authenticate": {"challenge"}, "Proxy-Authorization": {"credentials"}, + "Te": {"trailers"}, "Trailer": {"X-Untrusted"}, + "Transfer-Encoding": {"chunked"}, "Upgrade": {"websocket"}, + "Content-Length": {"999"}, "Content-Type": {"application/octet-stream"}, + "Content-Encoding": {"gzip"}, "Etag": {`"v1"`}, + "Set-Cookie": {"a=1", "b=2"}, "Www-Authenticate": {"Bearer realm=test"}, + } + resp := &http.Response{StatusCode: http.StatusOK, Header: headers, Body: io.NopCloser(strings.NewReader("body"))} + w := httptest.NewRecorder() + w.Header().Set("X-Request-ID", "local") + testProxy().relayResponse(w, httptest.NewRequest(http.MethodGet, "/", nil), resp, nil) + for _, name := range []string{ + "Connection", "Keep-Alive", "Proxy-Connection", "Proxy-Authenticate", + "Proxy-Authorization", "Te", "Trailer", "Transfer-Encoding", "Upgrade", + "X-Private", "X-Second", "Content-Length", + } { + if got := w.Header().Get(name); got != "" { + t.Errorf("hop-by-hop header %s survived: %q", name, got) + } + } + for _, name := range []string{"Content-Type", "Content-Encoding", "Etag", "Set-Cookie", "Www-Authenticate"} { + if got := strings.Join(w.Header().Values(name), ","); got != strings.Join(headers.Values(name), ",") { + t.Errorf("end-to-end header %s changed: %q", name, got) + } + } + if w.Header().Get("X-Request-ID") != "local" || w.Body.String() != "body" { + t.Fatal("local header or response body changed") + } + if headers.Get("X-Private") != "secret" { + t.Fatal("relay mutated the upstream response headers") + } +} + +func TestRelayResponseBodiless(t *testing.T) { + for _, tt := range []struct { + method string + status int + length string + }{ + {http.MethodHead, http.StatusOK, "99"}, + {http.MethodGet, http.StatusNoContent, ""}, + {http.MethodGet, http.StatusResetContent, "0"}, + {http.MethodGet, http.StatusNotModified, "99"}, + {http.MethodGet, http.StatusEarlyHints, ""}, + } { + t.Run(fmt.Sprintf("%s/%d", tt.method, tt.status), func(t *testing.T) { + resp := &http.Response{ + StatusCode: tt.status, Header: http.Header{"Content-Length": {"99"}}, + Body: io.NopCloser(iotest.ErrReader(errors.New("body must not be read"))), + Trailer: http.Header{"X-Checksum": {"ignored"}}, + } + w := httptest.NewRecorder() + testProxy().relayResponse(w, httptest.NewRequest(tt.method, "/", nil), resp, nil) + if w.Code != tt.status || w.Body.Len() != 0 || w.Header().Get("Trailer") != "" { + t.Fatalf("unexpected bodiless response: %d %v %q", w.Code, w.Header(), w.Body.String()) + } + if got := w.Header().Get("Content-Length"); got != tt.length { + t.Errorf("Content-Length = %q, want %q", got, tt.length) + } + }) + } +} + +type relayFailWriter struct{ http.ResponseWriter } + +func (w relayFailWriter) Write([]byte) (int, error) { + return 0, errors.New("downstream disconnected") +} + +func TestRelayResponseCopyFailure(t *testing.T) { + for _, downstreamError := range []bool{false, true} { + t.Run(fmt.Sprintf("downstreamError=%v", downstreamError), func(t *testing.T) { + var logs bytes.Buffer + proxy := testProxy() + proxy.Logger = slog.New(slog.NewTextHandler(&logs, nil)) + resp := &http.Response{ + StatusCode: http.StatusOK, Header: make(http.Header), + Request: httptest.NewRequest(http.MethodGet, "http://upstream.test/file", nil), + Body: io.NopCloser(io.MultiReader(strings.NewReader("part"), iotest.ErrReader(io.ErrUnexpectedEOF))), + } + var w http.ResponseWriter = httptest.NewRecorder() + wantBytes := "bytes=4" + if downstreamError { + w = relayFailWriter{httptest.NewRecorder()} + wantBytes = "bytes=0" + } + defer func() { + if got := recover(); got != http.ErrAbortHandler { + t.Errorf("panic = %v, want http.ErrAbortHandler", got) + } + for _, field := range []string{"url=http://upstream.test/file", "status=200", wantBytes, "error="} { + if !strings.Contains(logs.String(), field) { + t.Errorf("log missing %q: %s", field, logs.String()) + } + } + }() + proxy.relayResponse(w, httptest.NewRequest(http.MethodGet, "/", nil), resp, nil) + }) + } +} + +func TestRelayResponseTrailerValidation(t *testing.T) { + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Connection": {"X-Private"}, "Content-Length": {"4"}}, + Body: io.NopCloser(strings.NewReader("body")), + Trailer: http.Header{ + "X-Checksum": {"verified"}, "X-Private": {"secret"}, + "Content-Length": {"999"}, "Content-Type": {"bad/type"}, + "Authorization": {"secret"}, "Transfer-Encoding": {"chunked"}, + "If-Match": {"secret"}, "Invalid Name": {"invalid"}, + }, + } + w := httptest.NewRecorder() + testProxy().relayResponse(w, httptest.NewRequest(http.MethodGet, "/", nil), resp, nil) + result := w.Result() + defer func() { _ = result.Body.Close() }() + if result.Header.Get("Content-Length") != "" { + t.Error("declared trailers must remove Content-Length") + } + if len(result.Trailer) != 1 || result.Trailer.Get("X-Checksum") != "verified" { + t.Fatalf("invalid trailers leaked: %v", result.Trailer) + } +} + +func TestRelayClosesBodyOnAbort(t *testing.T) { + for _, mode := range []string{"upstream", "file", "metadata"} { + t.Run(mode, func(t *testing.T) { + body := &closeTrackingReader{Reader: iotest.ErrReader(io.ErrUnexpectedEOF)} + proxy := testProxy() + proxy.HTTPClient = &http.Client{Transport: pypiRoundTripFunc(func(r *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: body, Request: r}, nil + })} + defer func() { + if got := recover(); got != http.ErrAbortHandler { + t.Errorf("panic = %v, want http.ErrAbortHandler", got) + } + if !body.closed { + t.Error("upstream body was not closed on abort") + } + }() + relayTestRoutes(proxy, "http://upstream.test").ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/"+mode, nil)) + }) + } +} diff --git a/internal/handler/swift.go b/internal/handler/swift.go index 1c289d0b..93bc9a71 100644 --- a/internal/handler/swift.go +++ b/internal/handler/swift.go @@ -8,7 +8,6 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "net/url" "strconv" @@ -411,21 +410,18 @@ func (h *SwiftHandler) proxySwiftResource(w http.ResponseWriter, r *http.Request } defer func() { _ = resp.Body.Close() }() - copySwiftResponseHeaders(w.Header(), resp.Header) - if location := resp.Header.Get("Location"); location != "" { - w.Header().Set("Location", h.rewriteRegistryURL(location, upstreamURL)) - } - for _, link := range resp.Header.Values("Link") { - w.Header().Add("Link", h.rewriteLinkHeader(link, upstreamURL)) - } - if w.Header().Get("Content-Version") == "" { - w.Header().Set("Content-Version", swiftContentVersion) - } - - w.WriteHeader(resp.StatusCode) - if r.Method != http.MethodHead { - _, _ = io.Copy(w, resp.Body) - } + h.proxy.relayResponse(w, r, resp, func(dst, src http.Header) { + copySwiftResponseHeaders(dst, src) + if location := src.Get("Location"); location != "" { + dst.Set("Location", h.rewriteRegistryURL(location, upstreamURL)) + } + for _, link := range src.Values("Link") { + dst.Add("Link", h.rewriteLinkHeader(link, upstreamURL)) + } + if dst.Get("Content-Version") == "" { + dst.Set("Content-Version", swiftContentVersion) + } + }) } func copySwiftResponseHeaders(dst, src http.Header) { diff --git a/internal/server/relay_test.go b/internal/server/relay_test.go new file mode 100644 index 00000000..d2c6bd60 --- /dev/null +++ b/internal/server/relay_test.go @@ -0,0 +1,102 @@ +package server + +import ( + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/git-pkgs/proxy/internal/handler" + "github.com/go-chi/chi/v5/middleware" +) + +// Small bodies normally acquire Content-Length when the handler returns. Late, +// undeclared trailers need a flush through the production responseWriter first. +func TestRelayLateTrailersThroughMiddleware(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, "small body") + if err := http.NewResponseController(w).Flush(); err != nil { + t.Error(err) + } + w.Header().Set(http.TrailerPrefix+"X-Late", "verified") + })) + defer upstream.Close() + for _, protocol := range []string{"http1", "http2"} { + t.Run(protocol, func(t *testing.T) { + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + s := &Server{logger: logger} + proxy := &handler.Proxy{Logger: logger, HTTPClient: upstream.Client()} + h := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + proxy.ProxyUpstream(w, r, upstream.URL, nil) + }) + downstream := httptest.NewUnstartedServer(s.LoggerMiddleware(middleware.Recoverer(h))) + wantProto := 1 + if protocol == "http2" { + downstream.EnableHTTP2 = true + downstream.StartTLS() + wantProto = 2 + } else { + downstream.Start() + } + defer downstream.Close() + resp, err := downstream.Client().Get(downstream.URL) + if err != nil { + t.Fatal(err) + } + defer func() { _ = resp.Body.Close() }() + body, err := io.ReadAll(resp.Body) + if err != nil || string(body) != "small body" || resp.ProtoMajor != wantProto { + t.Fatalf("unexpected response: protocol=%s body=%q error=%v", resp.Proto, body, err) + } + if resp.Trailer.Get("X-Late") != "verified" || resp.Header.Get("X-Late") != "" { + t.Fatalf("late trailer lost or sent as a header: headers=%v trailers=%v", resp.Header, resp.Trailer) + } + }) + } +} + +func TestRelayAbortThroughMiddleware(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, strings.Repeat("x", 64*1024)) + if err := http.NewResponseController(w).Flush(); err != nil { + t.Error(err) + } + // End the upstream chunked body without its terminating chunk. + panic(http.ErrAbortHandler) + })) + defer upstream.Close() + for _, protocol := range []string{"http1", "http2"} { + t.Run(protocol, func(t *testing.T) { + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + s := &Server{logger: logger} + proxy := &handler.Proxy{Logger: logger, HTTPClient: upstream.Client()} + h := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + proxy.ProxyUpstream(w, r, upstream.URL, nil) + }) + downstream := httptest.NewUnstartedServer(s.LoggerMiddleware(middleware.Recoverer(h))) + wantProto := 1 + if protocol == "http2" { + downstream.EnableHTTP2 = true + downstream.StartTLS() + wantProto = 2 + } else { + downstream.Start() + } + defer downstream.Close() + resp, err := downstream.Client().Get(downstream.URL) + if err != nil { + t.Fatal(err) + } + defer func() { _ = resp.Body.Close() }() + _, err = io.ReadAll(resp.Body) + if err == nil { + t.Error("truncated upstream was delivered as a complete response") + } + if resp.StatusCode != http.StatusOK || resp.ProtoMajor != wantProto { + t.Fatalf("unexpected response: status=%d protocol=%s", resp.StatusCode, resp.Proto) + } + }) + } +} diff --git a/internal/server/server.go b/internal/server/server.go index 6d2a2cb3..236109fb 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -1216,6 +1216,12 @@ type responseWriter struct { status int } +// Unwrap lets ResponseController reach capabilities such as flushing when a +// relayed response discovers trailers only after its body has been copied. +func (rw *responseWriter) Unwrap() http.ResponseWriter { + return rw.ResponseWriter +} + func (rw *responseWriter) WriteHeader(code int) { rw.status = code rw.ResponseWriter.WriteHeader(code)