Skip to content
Open
57 changes: 38 additions & 19 deletions github/github.go
Original file line number Diff line number Diff line change
Expand Up @@ -602,25 +602,6 @@ func newClient(opts clientOptions) (*Client, error) {
c.client.Transport = t2
}

if opts.token != nil {
transport := c.client.Transport
if transport == nil {
transport = http.DefaultTransport
}
c.client.Transport = roundTripperFunc(func(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())
req.Header.Set("Authorization", fmt.Sprintf("Bearer %v", *opts.token))
return transport.RoundTrip(req)
})
}

c.clientIgnoreRedirects = &http.Client{
Transport: c.client.Transport,
Timeout: c.client.Timeout,
Jar: c.client.Jar,
CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse },
}

if opts.apiVersionMin != nil {
c.apiVersionMin = *opts.apiVersionMin
}
Expand All @@ -647,6 +628,27 @@ func newClient(opts clientOptions) (*Client, error) {
c.uploadURL, _ = url.Parse(uploadBaseURL)
}

if opts.token != nil {
transport := c.client.Transport
if transport == nil {
transport = http.DefaultTransport
}
c.client.Transport = roundTripperFunc(func(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())
if c.shouldAuthorizeRequest(req) {
req.Header.Set("Authorization", fmt.Sprintf("Bearer %v", *opts.token))
}
return transport.RoundTrip(req)
})
}

c.clientIgnoreRedirects = &http.Client{
Transport: c.client.Transport,
Timeout: c.client.Timeout,
Jar: c.client.Jar,
CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse },
}

c.disableRateLimitCheck = opts.disableRateLimitCheck

if !c.disableRateLimitCheck {
Expand Down Expand Up @@ -705,6 +707,23 @@ func newClient(opts clientOptions) (*Client, error) {
return c, nil
}

func (c *Client) shouldAuthorizeRequest(req *http.Request) bool {
if req == nil || req.URL == nil {
return false
}

return sameOrigin(req.URL, c.baseURL) || sameOrigin(req.URL, c.uploadURL)
}

func sameOrigin(u, base *url.URL) bool {
if u == nil || base == nil {
return false
}

return strings.EqualFold(u.Scheme, base.Scheme) &&
strings.EqualFold(u.Host, base.Host)
}

// UserAgent returns the User-Agent header value for the client.
func (c *Client) UserAgent() string {
return c.userAgent
Expand Down
78 changes: 78 additions & 0 deletions github/github_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -596,6 +596,84 @@ func TestWithAuthToken(t *testing.T) {
})
}

func TestWithAuthTokenAuthorizesConfiguredHostsOnly(t *testing.T) {
t.Parallel()

const token = "secret-token"

trustedAuths := make(chan string, 2)
trusted := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
trustedAuths <- r.Header.Get("Authorization")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{}`))
}))
defer trusted.Close()

attackerAuths := make(chan string, 2)
attacker := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
attackerAuths <- r.Header.Get("Authorization")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{}`))
}))
defer attacker.Close()

baseURL := trusted.URL + "/api/"
uploadURL := trusted.URL + "/upload/"
client := mustNewClient(t, WithURLs(&baseURL, &uploadURL), WithAuthToken(token))

req, err := client.NewRequest(t.Context(), "GET", "repos/o/r", nil)
if err != nil {
t.Fatalf("NewRequest returned unexpected error: %v", err)
}
resp, err := client.BareDo(req)
if err != nil {
t.Fatalf("BareDo returned unexpected error: %v", err)
}
resp.Body.Close()

req, err = client.NewUploadRequest(t.Context(), "assets", strings.NewReader("x"), 1, "")
if err != nil {
t.Fatalf("NewUploadRequest returned unexpected error: %v", err)
}
resp, err = client.BareDo(req)
if err != nil {
t.Fatalf("BareDo returned unexpected error: %v", err)
}
resp.Body.Close()

for range 2 {
if got, want := <-trustedAuths, "Bearer "+token; got != want {
t.Errorf("Authorization on configured host = %q, want %q", got, want)
}
}

req, err = client.NewRequest(t.Context(), "GET", attacker.URL+"/repos/o/r", nil)
if err != nil {
t.Fatalf("NewRequest returned unexpected error: %v", err)
}
resp, err = client.BareDo(req)
if err != nil {
t.Fatalf("BareDo returned unexpected error: %v", err)
}
resp.Body.Close()

req, err = client.NewUploadRequest(t.Context(), attacker.URL+"/assets", strings.NewReader("x"), 1, "")
if err != nil {
t.Fatalf("NewUploadRequest returned unexpected error: %v", err)
}
resp, err = client.BareDo(req)
if err != nil {
t.Fatalf("BareDo returned unexpected error: %v", err)
}
resp.Body.Close()

for range 2 {
if got := <-attackerAuths; got != "" {
t.Errorf("Authorization on cross-host request = %q, want empty", got)
}
}
}

func TestWithURLs(t *testing.T) {
t.Parallel()
for _, tt := range []struct {
Expand Down
Loading