Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion github/actions_artifacts.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
34 changes: 34 additions & 0 deletions github/actions_artifacts_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
2 changes: 1 addition & 1 deletion github/actions_workflow_jobs.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down
4 changes: 2 additions & 2 deletions github/actions_workflow_runs.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down Expand Up @@ -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
}

Expand Down
43 changes: 24 additions & 19 deletions github/github.go
Original file line number Diff line number Diff line change
Expand Up @@ -1575,15 +1575,15 @@ 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
if maxRedirects > 0 && rerr.StatusCode == http.StatusMovedPermanently {
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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
81 changes: 81 additions & 0 deletions github/github_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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-
Expand Down Expand Up @@ -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 {
Expand Down
2 changes: 1 addition & 1 deletion github/repos_contents.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
Loading