Skip to content
Merged
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
107 changes: 104 additions & 3 deletions internal/ghmcp/oauth_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,12 @@ import (
"log/slog"
"net/http"
"net/http/httptest"
"net/url"
"testing"

"github.com/github/github-mcp-server/internal/oauth"
"github.com/github/github-mcp-server/pkg/github"
"github.com/github/github-mcp-server/pkg/http/headers"
"github.com/github/github-mcp-server/pkg/utils"
"github.com/google/jsonschema-go/jsonschema"
"github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/stretchr/testify/assert"
Expand All @@ -23,6 +23,108 @@ func discardLogger() *slog.Logger {
return slog.New(slog.NewTextHandler(io.Discard, nil))
}

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

var foreignAuth string
foreign := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
foreignAuth = r.Header.Get(headers.AuthorizationHeader)
w.WriteHeader(http.StatusOK)
}))
defer foreign.Close()

var sourceAuth string
source := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sourceAuth = r.Header.Get(headers.AuthorizationHeader)
http.Redirect(w, r, foreign.URL, http.StatusFound)
}))
defer source.Close()

tests := []struct {
name string
cfg github.MCPServerConfig
}{
{
name: "static token",
cfg: github.MCPServerConfig{
Version: "test",
Token: "static-token",
},
},
{
name: "token provider",
cfg: github.MCPServerConfig{
Version: "test",
TokenProvider: func() string { return "provider-token" },
},
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
apiHost := newStaticAPIHostResolver(t, source.URL)
clients, err := createGitHubClients(tt.cfg, apiHost)
require.NoError(t, err)

sourceAuth = ""
foreignAuth = ""
resp, err := clients.rest.Client().Get(source.URL + "/rest")
require.NoError(t, err)
resp.Body.Close()
assert.NotEmpty(t, sourceAuth, "REST request must authenticate to the configured host")
assert.Empty(t, foreignAuth, "REST redirect must not authenticate to a foreign host")

sourceAuth = ""
foreignAuth = ""
resp, err = clients.raw.GetRawContent(context.Background(), "owner", "repo", "file", nil)
require.NoError(t, err)
resp.Body.Close()
assert.NotEmpty(t, sourceAuth, "raw request must authenticate to the configured host")
assert.Empty(t, foreignAuth, "raw redirect must not authenticate to a foreign host")
})
}
}

type staticAPIHostResolver struct {
restURL *url.URL
graphQLURL *url.URL
uploadURL *url.URL
rawURL *url.URL
}

func newStaticAPIHostResolver(t *testing.T, endpoint string) staticAPIHostResolver {
t.Helper()

u, err := url.Parse(endpoint)
require.NoError(t, err)
return staticAPIHostResolver{
restURL: u,
graphQLURL: u,
uploadURL: u,
rawURL: u,
}
}

func (r staticAPIHostResolver) BaseRESTURL(context.Context) (*url.URL, error) {
return r.restURL, nil
}

func (r staticAPIHostResolver) GraphqlURL(context.Context) (*url.URL, error) {
return r.graphQLURL, nil
}

func (r staticAPIHostResolver) UploadURL(context.Context) (*url.URL, error) {
return r.uploadURL, nil
}

func (r staticAPIHostResolver) RawURL(context.Context) (*url.URL, error) {
return r.rawURL, nil
}

func (r staticAPIHostResolver) AuthorizationServerURL(context.Context) (*url.URL, error) {
return r.restURL, nil
}

// probeToolName is the name of the throwaway tool the harness registers; its
// handler runs a probe closure against a sessionPrompter so the adapter can be
// exercised against a real, fully-negotiated server session from the client side.
Expand Down Expand Up @@ -583,8 +685,7 @@ func TestCreateGitHubClientsTokenProvider(t *testing.T) {
defer server.Close()

current := ""
apiHost, err := utils.NewAPIHost(server.URL)
require.NoError(t, err)
apiHost := newStaticAPIHostResolver(t, server.URL)

clients, err := createGitHubClients(github.MCPServerConfig{
Version: "test",
Expand Down
43 changes: 23 additions & 20 deletions internal/ghmcp/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -63,30 +63,32 @@ func createGitHubClients(cfg github.MCPServerConfig, apiHost utils.APIHostResolv
return nil, fmt.Errorf("failed to get Raw URL: %w", err)
}

// Construct REST client. When a TokenProvider is configured, we
// authenticate via BearerAuthTransport and skip go-github's WithAuthToken:
// the latter installs its own round tripper that would pin the static token
// and shadow the dynamic one.
// allowedHosts scopes the bearer token to the configured GitHub hosts, so a
// response that redirects off them does not carry the token to the redirect
// target. See transport.BearerAuthTransport.
allowedHosts := []string{
restURL.Host,
uploadURL.Host,
graphQLURL.Host,
rawURL.Host,
}

// Construct REST client. BearerAuthTransport handles both static and
// provider-backed tokens so every authentication mode uses the same host
// restrictions.
restUATransport := &transport.UserAgentTransport{
Transport: http.DefaultTransport,
Agent: fmt.Sprintf("github-mcp-server/%s", cfg.Version),
}
var restClient *gogithub.Client
if cfg.TokenProvider != nil {
restClient, err = gogithub.NewClient(
gogithub.WithHTTPClient(&http.Client{Transport: &transport.BearerAuthTransport{
Transport: restUATransport,
TokenProvider: cfg.TokenProvider,
}}),
gogithub.WithEnterpriseURLs(restURL.String(), uploadURL.String()),
)
} else {
restClient, err = gogithub.NewClient(
gogithub.WithHTTPClient(&http.Client{Transport: restUATransport}),
gogithub.WithAuthToken(cfg.Token),
gogithub.WithEnterpriseURLs(restURL.String(), uploadURL.String()),
)
}
restClient, err := gogithub.NewClient(
gogithub.WithHTTPClient(&http.Client{Transport: &transport.BearerAuthTransport{
Transport: restUATransport,
Token: cfg.Token,
TokenProvider: cfg.TokenProvider,
AllowedHosts: allowedHosts,
}}),
gogithub.WithEnterpriseURLs(restURL.String(), uploadURL.String()),
)
if err != nil {
return nil, fmt.Errorf("failed to create REST client: %w", err)
}
Expand All @@ -100,6 +102,7 @@ func createGitHubClients(cfg github.MCPServerConfig, apiHost utils.APIHostResolv
},
Token: cfg.Token,
TokenProvider: cfg.TokenProvider,
AllowedHosts: allowedHosts,
},
}

Expand Down
56 changes: 49 additions & 7 deletions pkg/github/dependencies.go
Original file line number Diff line number Diff line change
Expand Up @@ -330,10 +330,29 @@ func (d *RequestDeps) GetClient(ctx context.Context) (*gogithub.Client, error) {
if err != nil {
return nil, fmt.Errorf("failed to get upload URL: %w", err)
}
graphqlURL, err := d.apiHosts.GraphqlURL(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get GraphQL URL: %w", err)
}
rawURL, err := d.apiHosts.RawURL(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get Raw URL: %w", err)
}

allowedHosts := []string{
baseRestURL.Host,
uploadURL.Host,
graphqlURL.Host,
rawURL.Host,
}

// Construct REST client
restClient, err := gogithub.NewClient(
gogithub.WithAuthToken(token),
gogithub.WithHTTPClient(&http.Client{Transport: &transport.BearerAuthTransport{
Transport: http.DefaultTransport,
Token: token,
AllowedHosts: allowedHosts,
}}),
gogithub.WithUserAgent(fmt.Sprintf("github-mcp-server/%s", d.version)),
gogithub.WithEnterpriseURLs(baseRestURL.String(), uploadURL.String()),
)
Expand All @@ -355,6 +374,33 @@ func (d *RequestDeps) GetGQLClient(ctx context.Context) (*githubv4.Client, error
}
token := tokenInfo.Token

baseRestURL, err := d.apiHosts.BaseRESTURL(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get base REST URL: %w", err)
}
uploadURL, err := d.apiHosts.UploadURL(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get upload URL: %w", err)
}
graphqlURL, err := d.apiHosts.GraphqlURL(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get GraphQL URL: %w", err)
}
rawURL, err := d.apiHosts.RawURL(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get Raw URL: %w", err)
}

// allowedHosts scopes the bearer token to the configured GitHub hosts, so a
// response that redirects off them does not carry the token to the redirect
// target. See transport.BearerAuthTransport.
allowedHosts := []string{
baseRestURL.Host,
uploadURL.Host,
graphqlURL.Host,
rawURL.Host,
}

// Construct GraphQL client
// We use NewEnterpriseClient unconditionally since we already parsed the API host
// Wrap transport with GraphQLFeaturesTransport to inject feature flags from context,
Expand All @@ -364,15 +410,11 @@ func (d *RequestDeps) GetGQLClient(ctx context.Context) (*githubv4.Client, error
Transport: &transport.GraphQLFeaturesTransport{
Transport: http.DefaultTransport,
},
Token: token,
Token: token,
AllowedHosts: allowedHosts,
},
}

graphqlURL, err := d.apiHosts.GraphqlURL(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get GraphQL URL: %w", err)
}

gqlClient := githubv4.NewEnterpriseClient(graphqlURL.String(), gqlHTTPClient)
return gqlClient, nil
}
Expand Down
Loading
Loading