diff --git a/github/actions_artifacts.go b/github/actions_artifacts.go index 9cc3c64b3cb..6df864e4095 100644 --- a/github/actions_artifacts.go +++ b/github/actions_artifacts.go @@ -179,7 +179,7 @@ func (s *ActionsService) downloadArtifactWithoutRateLimit(ctx context.Context, u return nil, newResponse(resp), fmt.Errorf("unexpected status code: %v", resp.Status) } - parsedURL, err := url.Parse(resp.Header.Get("Location")) + parsedURL, err := resp.Location() if err != nil { return nil, newResponse(resp), err } diff --git a/github/actions_artifacts_test.go b/github/actions_artifacts_test.go index 1be74b6222e..07f24e999f0 100644 --- a/github/actions_artifacts_test.go +++ b/github/actions_artifacts_test.go @@ -333,6 +333,40 @@ func TestActionsService_DownloadArtifact(t *testing.T) { } } +func TestActionsService_DownloadArtifact_resolvesRelativeLocation(t *testing.T) { + t.Parallel() + tcs := []struct { + name string + respectRateLimits bool + }{ + {name: "withoutRateLimits"}, + {name: "withRateLimits", respectRateLimits: true}, + } + + for _, tc := range tcs { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + client, mux, serverURL := setup(t) + client.rateLimitRedirectionalEndpoints = tc.respectRateLimits + + mux.HandleFunc("/repos/o/r/actions/artifacts/1/zip", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Location", "download") + w.WriteHeader(http.StatusFound) + }) + + got, _, err := client.Actions.DownloadArtifact(t.Context(), "o", "r", 1, 0) + if err != nil { + t.Fatalf("Actions.DownloadArtifact returned error: %v", err) + } + + want := serverURL + baseURLPath + "/repos/o/r/actions/artifacts/1/download" + if got.String() != want { + t.Errorf("Actions.DownloadArtifact returned %q, want %q", got, want) + } + }) + } +} + func TestActionsService_DownloadArtifact_invalidOwner(t *testing.T) { t.Parallel() tcs := []struct { diff --git a/github/actions_workflow_jobs.go b/github/actions_workflow_jobs.go index 9419cf89711..334d46a427f 100644 --- a/github/actions_workflow_jobs.go +++ b/github/actions_workflow_jobs.go @@ -168,7 +168,7 @@ func (s *ActionsService) getWorkflowJobLogsWithoutRateLimit(ctx context.Context, return nil, newResponse(resp), fmt.Errorf("unexpected status code: %v", resp.Status) } - parsedURL, err := url.Parse(resp.Header.Get("Location")) + parsedURL, err := resp.Location() return parsedURL, newResponse(resp), err } diff --git a/github/actions_workflow_runs.go b/github/actions_workflow_runs.go index 2e12ed9fcf3..405070a28a1 100644 --- a/github/actions_workflow_runs.go +++ b/github/actions_workflow_runs.go @@ -281,7 +281,7 @@ func (s *ActionsService) getWorkflowRunAttemptLogsWithoutRateLimit(ctx context.C return nil, newResponse(resp), fmt.Errorf("unexpected status code: %v", resp.Status) } - parsedURL, err := url.Parse(resp.Header.Get("Location")) + parsedURL, err := resp.Location() return parsedURL, newResponse(resp), err } @@ -401,7 +401,7 @@ func (s *ActionsService) getWorkflowRunLogsWithoutRateLimit(ctx context.Context, return nil, newResponse(resp), fmt.Errorf("unexpected status code: %v", resp.Status) } - parsedURL, err := url.Parse(resp.Header.Get("Location")) + parsedURL, err := resp.Location() return parsedURL, newResponse(resp), err } diff --git a/github/github.go b/github/github.go index 5b1fa27049a..9c59f4c3468 100644 --- a/github/github.go +++ b/github/github.go @@ -1575,7 +1575,7 @@ func (c *Client) bareDoUntilFound(req *http.Request, maxRedirects int) (*url.URL if rerr.Location == nil { return nil, nil, errInvalidLocation } - newURL := c.baseURL.ResolveReference(rerr.Location) + newURL := req.URL.ResolveReference(rerr.Location) return newURL, response, nil } // If permanent redirect response is returned, follow it @@ -1583,7 +1583,7 @@ func (c *Client) bareDoUntilFound(req *http.Request, maxRedirects int) (*url.URL if rerr.Location == nil { return nil, nil, errInvalidLocation } - newURL := c.baseURL.ResolveReference(rerr.Location) + newURL := req.URL.ResolveReference(rerr.Location) // Refuse to follow a permanent redirect outside the origins // this client may send credentials to: the auth transport // attaches them on every hop, so a cross-host target would @@ -1921,8 +1921,8 @@ func equalDurationPtr(a, b *time.Duration) bool { // 307 (Temporary Redirect) // 308 (Permanent Redirect) // -// If there was a valid Location header included, it will be parsed to a URL. You should use -// `BaseURL.ResolveReference()` to enrich it with the correct hostname where needed. +// If there was a valid Location header included, it will be parsed to a URL. Relative locations +// should be resolved against the request URL that received the response. type RedirectionError struct { Response *http.Response // HTTP response that caused this error StatusCode int @@ -2397,34 +2397,39 @@ func (c *Client) roundTripWithOptionalFollowRedirect(ctx context.Context, u stri // If redirect response is returned, follow it if maxRedirects > 0 && resp.StatusCode == http.StatusMovedPermanently { _ = resp.Body.Close() - u = resp.Header.Get("Location") - if err := c.checkRedirectHost(u); err != nil { + location := resp.Header.Get("Location") + target, err := c.resolveAndCheckRedirect(location, req.URL) + if err != nil { return nil, err } - resp, err = c.roundTripWithOptionalFollowRedirect(ctx, u, maxRedirects-1, opts...) + return c.roundTripWithOptionalFollowRedirect(ctx, target.String(), maxRedirects-1, opts...) } return resp, err } -// checkRedirectHost returns an error if the redirect target is outside the -// origins this client may send credentials to. The auth transport attaches -// credentials on every hop, so a cross-origin Location header would otherwise -// carry them to a host the caller never configured, when a compromised or -// malicious API response supplies one. An empty Location is also rejected. -func (c *Client) checkRedirectHost(location string) error { +// resolveAndCheckRedirect resolves a Location header against the request URL +// and returns an error if the target is outside the origins this client may +// send credentials to. The auth transport attaches credentials on every hop, +// so a cross-origin Location header would otherwise carry them to a host the +// caller never configured, when a compromised or malicious API response +// supplies one. An empty Location is also rejected. +func (c *Client) resolveAndCheckRedirect(location string, requestURL *url.URL) (*url.URL, error) { if location == "" { - return errInvalidLocation + return nil, errInvalidLocation + } + if err := checkURLPathTraversal(location); err != nil { + return nil, err } target, err := url.Parse(location) if err != nil { - return fmt.Errorf("invalid redirect location %q: %w", location, err) + return nil, fmt.Errorf("invalid redirect location %q: %w", location, err) } - // Resolve relative locations against BaseURL so relative paths are allowed. - target = c.baseURL.ResolveReference(target) + // Resolve relative locations against the URL that returned the redirect. + target = requestURL.ResolveReference(target) if !c.shouldAuthorizeRequest(target) { - return fmt.Errorf("refusing to follow cross-host redirect from %q to %q", c.baseURL.Host, target.Host) + return nil, fmt.Errorf("refusing to follow cross-host redirect from %q to %q", requestURL.Host, target.Host) } - return nil + return target, nil } // Ptr is a helper routine that allocates a new T value diff --git a/github/github_test.go b/github/github_test.go index dc902f9ed1c..d7501ea9e5c 100644 --- a/github/github_test.go +++ b/github/github_test.go @@ -3826,6 +3826,40 @@ func TestBareDoUntilFound_MissingRedirectLocation(t *testing.T) { } } +func TestBareDoUntilFound_ResolvesRelative301AgainstRequestURL(t *testing.T) { + t.Parallel() + client, mux, _ := setup(t) + + const requestPath = "/repos/owner/repo/actions/artifacts/123/zip" + const redirectPath = "/repos/owner/repo/actions/artifacts/123/download" + var followed atomic.Bool + + mux.HandleFunc(requestPath, func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Location", "download") + w.WriteHeader(http.StatusMovedPermanently) + }) + mux.HandleFunc(redirectPath, func(w http.ResponseWriter, _ *http.Request) { + followed.Store(true) + w.WriteHeader(http.StatusOK) + }) + + req, err := client.NewRequest(t.Context(), "GET", strings.TrimPrefix(requestPath, "/"), nil) + if err != nil { + t.Fatalf("NewRequest returned error: %v", err) + } + _, resp, err := client.bareDoUntilFound(req, 1) + if err != nil { + t.Fatalf("bareDoUntilFound returned error: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Errorf("bareDoUntilFound returned status %v, want %v", resp.StatusCode, http.StatusOK) + } + if !followed.Load() { + t.Error("Expected relative redirect to be resolved against the request URL.") + } +} + // TestRoundTripWithOptionalFollowRedirect_RejectsCrossHostRedirect verifies // that roundTripWithOptionalFollowRedirect refuses to follow a 301 redirect to // a different host, preventing Authorization-header leakage to attacker- @@ -3897,6 +3931,53 @@ func TestRoundTripWithOptionalFollowRedirect_AllowsSameHostRedirect(t *testing.T } } +func TestRoundTripWithOptionalFollowRedirect_ResolvesRelativeLocationAgainstRequestURL(t *testing.T) { + t.Parallel() + client, mux, _ := setup(t) + + const requestPath = "/repos/owner/repo/actions/artifacts/123/zip" + const redirectPath = "/repos/owner/repo/actions/artifacts/123/download" + var followed atomic.Bool + + mux.HandleFunc(requestPath, func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Location", "download") + w.WriteHeader(http.StatusMovedPermanently) + }) + mux.HandleFunc(redirectPath, func(w http.ResponseWriter, _ *http.Request) { + followed.Store(true) + w.WriteHeader(http.StatusOK) + }) + + resp, err := client.roundTripWithOptionalFollowRedirect(t.Context(), strings.TrimPrefix(requestPath, "/"), 1) + if err != nil { + t.Fatalf("Unexpected error following relative redirect: %v", err) + } + if resp != nil && resp.Body != nil { + defer resp.Body.Close() + } + if resp == nil || resp.StatusCode != http.StatusOK { + t.Fatalf("Expected redirect target to return %v, got %#v", http.StatusOK, resp) + } + if !followed.Load() { + t.Error("Expected relative redirect to be resolved against the request URL.") + } +} + +func TestRoundTripWithOptionalFollowRedirect_RejectsPathTraversalInLocation(t *testing.T) { + t.Parallel() + client, mux, _ := setup(t) + + mux.HandleFunc("/repos/owner/repo/archive", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Location", "../download") + w.WriteHeader(http.StatusMovedPermanently) + }) + + _, err := client.roundTripWithOptionalFollowRedirect(t.Context(), "repos/owner/repo/archive", 1) + if !errors.Is(err, ErrPathForbidden) { + t.Fatalf("Expected ErrPathForbidden, got %v", err) + } +} + func TestSanitizeURL(t *testing.T) { t.Parallel() tests := []struct { diff --git a/github/repos_contents.go b/github/repos_contents.go index dc6b4a03c1e..74d6e2b42a7 100644 --- a/github/repos_contents.go +++ b/github/repos_contents.go @@ -370,7 +370,7 @@ func (s *RepositoriesService) getArchiveLinkWithoutRateLimit(ctx context.Context return nil, newResponse(resp), fmt.Errorf("unexpected status code: %v", resp.Status) } - parsedURL, err := url.Parse(resp.Header.Get("Location")) + parsedURL, err := resp.Location() if err != nil { return nil, newResponse(resp), err }