diff --git a/README.md b/README.md index d013303d81..221211dc58 100644 --- a/README.md +++ b/README.md @@ -403,7 +403,13 @@ For a complete overview of all installation options, see our **[Installation Gui ### Build from source If you don't have Docker, you can use `go build` to build the binary in the -`cmd/github-mcp-server` directory, and use the `github-mcp-server stdio` command with the `GITHUB_PERSONAL_ACCESS_TOKEN` environment variable set to your token. To specify the output location of the build, use the `-o` flag. You should configure your server to use the built executable as its `command`. For example: +`cmd/github-mcp-server` directory, and use the `github-mcp-server stdio` command with the `GITHUB_PERSONAL_ACCESS_TOKEN` environment variable set to your token. To specify the output location of the build, use the `-o` flag. You should configure your server to use the built executable as its `command`. + +STDIO API requests identify the server as `github-mcp-server/` and retain the upstream MCP client's name/version in parentheses when available. Control characters, quotes, backslashes and parentheses in client metadata are escaped so the HTTP header remains valid; ordinary names, versions, spaces and printable Unicode are preserved. Release builds keep their release version. Source and default Docker builds with revision metadata use `vcs-`. The `-dirty` marker applies only when both the revision and modified state come from embedded VCS metadata, never to a valid, explicitly supplied `main.commit`. This is a VCS build identifier, not a release number. + +Build the complete package with `go build -o github-mcp-server ./cmd/github-mcp-server` from a Git checkout to embed its VCS revision. For builds without VCS metadata, supply the actual release with `-ldflags '-X main.version='` or the full source revision with `-ldflags '-X main.commit='`. Valid metadata is selected in this order: explicit release, explicit source revision, embedded VCS revision, then installed main-module version. Valid explicit revisions remain authoritative even if build-context filtering changes embedded VCS metadata. Missing or malformed candidates fall through to the next usable source. If none is available, STDIO still starts with the development label `dev` and emits a warning on stderr, leaving stdout available for the MCP protocol. Supply real release or revision metadata when version-specific attribution is needed. + +For example: ```JSON { diff --git a/cmd/github-mcp-server/main.go b/cmd/github-mcp-server/main.go index c0cadbbc63..91abef6a08 100644 --- a/cmd/github-mcp-server/main.go +++ b/cmd/github-mcp-server/main.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "os" + "runtime/debug" "strings" "time" @@ -38,7 +39,13 @@ var ( Use: "stdio", Short: "Start stdio server", Long: `Start a server that communicates via standard input/output streams using JSON-RPC messages.`, - RunE: func(_ *cobra.Command, _ []string) error { + RunE: func(cmd *cobra.Command, _ []string) error { + info, _ := debug.ReadBuildInfo() + serverVersion := resolveServerVersion(version, commit, info) + if serverVersion == developmentServerVersion { + cmd.PrintErrln("Warning: no usable server build version metadata; using development version dev") + } + token := viper.GetString("personal_access_token") appID := viper.GetString("app-id") appInstallationID := viper.GetString("app-installation-id") @@ -110,7 +117,7 @@ var ( ttl := viper.GetDuration("repo-access-cache-ttl") stdioServerConfig := ghmcp.StdioServerConfig{ - Version: version, + Version: serverVersion, Host: viper.GetString("host"), Token: token, EnabledToolsets: enabledToolsets, diff --git a/cmd/github-mcp-server/version.go b/cmd/github-mcp-server/version.go new file mode 100644 index 0000000000..51b9e5880c --- /dev/null +++ b/cmd/github-mcp-server/version.go @@ -0,0 +1,75 @@ +package main + +import ( + "encoding/hex" + "runtime/debug" + "strings" + "unicode" +) + +const developmentServerVersion = "dev" + +func resolveServerVersion(release, revision string, info *debug.BuildInfo) string { + isPlaceholder := func(value string) bool { + switch value { + case "", "version", "dev", "unknown", "(devel)": + return true + default: + return false + } + } + validRelease := func(value string) bool { + if isPlaceholder(value) { + return false + } + for _, r := range value { + if r > unicode.MaxASCII || r <= ' ' || r == 127 || strings.ContainsRune("()<>@,;:\\\"/[]?={}", r) { + return false + } + } + return true + } + if validRelease(release) { + return release + } + + var vcsRevision string + var dirty bool + if info != nil { + for _, setting := range info.Settings { + switch setting.Key { + case "vcs.revision": + vcsRevision = setting.Value + case "vcs.modified": + dirty = setting.Value == "true" + } + } + } + revisions := []struct { + value string + dirty bool + }{ + {value: revision}, + {value: vcsRevision, dirty: dirty}, + } + for _, candidate := range revisions { + if len(candidate.value) != 40 && len(candidate.value) != 64 { + continue + } + if _, err := hex.DecodeString(candidate.value); err != nil { + continue + } + if strings.Trim(candidate.value, "0") == "" { + continue + } + resolved := "vcs-" + strings.ToLower(candidate.value) + if candidate.dirty { + resolved += "-dirty" + } + return resolved + } + if info != nil && validRelease(info.Main.Version) { + return info.Main.Version + } + return developmentServerVersion +} diff --git a/cmd/github-mcp-server/version_test.go b/cmd/github-mcp-server/version_test.go new file mode 100644 index 0000000000..4258b817c1 --- /dev/null +++ b/cmd/github-mcp-server/version_test.go @@ -0,0 +1,122 @@ +package main + +import ( + "runtime/debug" + "strings" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestResolveServerVersion(t *testing.T) { + t.Parallel() + sha := strings.Repeat("a1", 20) + otherSHA := strings.Repeat("b2", 20) + source := &debug.BuildInfo{ + Main: debug.Module{Version: "(devel)"}, + Settings: []debug.BuildSetting{ + {Key: "vcs.revision", Value: sha}, {Key: "vcs.modified", Value: "false"}, + }, + } + dirty := &debug.BuildInfo{ + Main: debug.Module{Version: "(devel)"}, + Settings: []debug.BuildSetting{ + {Key: "vcs.revision", Value: sha}, {Key: "vcs.modified", Value: "true"}, + }, + } + tests := []struct { + name string + release string + revision string + info *debug.BuildInfo + want string + }{ + {name: "release unchanged", release: "v1.2.3", revision: sha, info: dirty, want: "v1.2.3"}, + {name: "release suffix unchanged", release: "v1.2.3-rc.1+build.4", want: "v1.2.3-rc.1+build.4"}, + {name: "release before invalid revision", release: "v1.2.3", revision: "short", info: dirty, want: "v1.2.3"}, + {name: "source build", release: "version", revision: "commit", info: source, want: "vcs-" + sha}, + { + name: "source revision before inferred module version", release: "version", revision: "commit", + info: &debug.BuildInfo{ + Main: debug.Module{Version: "v1.2.4-0.20260916095829-a1a1a1a1a1a1+dirty"}, + Settings: dirty.Settings, + }, + want: "vcs-" + sha + "-dirty", + }, + {name: "dirty source", release: "version", revision: "commit", info: dirty, want: "vcs-" + sha + "-dirty"}, + {name: "Docker revision", release: "dev", revision: sha, want: "vcs-" + sha}, + {name: "Docker revision ignores context dirty state", release: "dev", revision: sha, info: dirty, want: "vcs-" + sha}, + {name: "linked revision ignores unrelated dirty state", release: "dev", revision: otherSHA, info: dirty, want: "vcs-" + otherSHA}, + {name: "linked revision precedence", release: "dev", revision: otherSHA, info: source, want: "vcs-" + otherSHA}, + {name: "SHA256", release: "dev", revision: strings.Repeat("ab", 32), want: "vcs-" + strings.Repeat("ab", 32)}, + {name: "canonical hex", release: "dev", revision: strings.ToUpper(sha), want: "vcs-" + sha}, + {name: "installed module", info: &debug.BuildInfo{Main: debug.Module{Version: "v1.2.3"}}, want: "v1.2.3"}, + {name: "absent build info", release: "version", revision: "commit", want: "dev"}, + {name: "missing VCS metadata", release: "dev", info: &debug.BuildInfo{}, want: "dev"}, + {name: "unknown placeholder", release: "unknown", want: "dev"}, + {name: "short revision", release: "dev", revision: "abcdef", want: "dev"}, + {name: "short revision falls through to VCS", release: "dev", revision: "abcdef", info: source, want: "vcs-" + sha}, + {name: "malformed linked revision falls through to VCS", release: "dev", revision: strings.Repeat("x", 40), info: source, want: "vcs-" + sha}, + {name: "malformed linked revision without VCS", release: "dev", revision: strings.Repeat("x", 40), want: "dev"}, + { + name: "malformed VCS falls through to installed module", + info: &debug.BuildInfo{ + Main: debug.Module{Version: "v1.2.3"}, + Settings: []debug.BuildSetting{{Key: "vcs.revision", Value: "short"}}, + }, + want: "v1.2.3", + }, + { + name: "malformed VCS without installed module", + info: &debug.BuildInfo{ + Main: debug.Module{Version: "(devel)"}, + Settings: []debug.BuildSetting{{Key: "vcs.revision", Value: "short"}}, + }, + want: "dev", + }, + {name: "invalid installed module", info: &debug.BuildInfo{Main: debug.Module{Version: "v1.2.3\n"}}, want: "dev"}, + {name: "placeholder installed module", info: &debug.BuildInfo{Main: debug.Module{Version: "(devel)"}}, want: "dev"}, + {name: "release whitespace", release: "v1.2.3 extra", want: "dev"}, + {name: "release newline", release: "v1.2.3\n", want: "dev"}, + {name: "release slash", release: "release/v1.2.3", want: "dev"}, + {name: "release non ASCII", release: "v1.2.3-\u00e9", want: "dev"}, + {name: "invalid release falls through to linked revision", release: "v1.2.3 extra", revision: otherSHA, info: dirty, want: "vcs-" + otherSHA}, + {name: "invalid release falls through to dirty VCS", release: "release/v1.2.3", info: dirty, want: "vcs-" + sha + "-dirty"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, resolveServerVersion(tt.release, tt.revision, tt.info)) + }) + } +} + +func TestResolveServerVersionSkipsNullRevisions(t *testing.T) { + t.Parallel() + for _, revision := range []string{strings.Repeat("0", 40), strings.Repeat("0", 64)} { + for _, modified := range []string{"false", "true"} { + for _, linked := range []bool{false, true} { + name := revision + "/modified=" + modified + if linked { + name += "/linked" + } + t.Run(name, func(t *testing.T) { + info := &debug.BuildInfo{Settings: []debug.BuildSetting{ + {Key: "vcs.revision", Value: revision}, + {Key: "vcs.modified", Value: modified}, + }} + commit := "commit" + want := "dev" + if linked { + commit = revision + info.Settings[0].Value = strings.Repeat("ab", 20) + want = "vcs-" + strings.Repeat("ab", 20) + if modified == "true" { + want += "-dirty" + } + } + assert.Equal(t, want, resolveServerVersion("dev", commit, info)) + }) + } + } + } +} diff --git a/docs/installation-guides/README.md b/docs/installation-guides/README.md index 46581aa77e..888caf64a4 100644 --- a/docs/installation-guides/README.md +++ b/docs/installation-guides/README.md @@ -66,7 +66,7 @@ The GitHub MCP Server can be installed using several methods. **Docker is the mo - **Pros**: Latest features, full customization, no external dependencies - **Cons**: Requires Go development environment, more complex setup - **Prerequisites**: [Go 1.24+](https://go.dev/doc/install) -- **Build command**: `go build -o github-mcp-server cmd/github-mcp-server/main.go` +- **Build command**: `go build -o github-mcp-server ./cmd/github-mcp-server` (build the complete package from a Git checkout to retain its source revision). - **Best for**: Developers who want the latest features or need custom modifications ### Important Notes on the GitHub MCP Server diff --git a/internal/ghmcp/server.go b/internal/ghmcp/server.go index f713a44026..5da57527e8 100644 --- a/internal/ghmcp/server.go +++ b/internal/ghmcp/server.go @@ -8,6 +8,7 @@ import ( "net/http" "os" "os/signal" + "strconv" "strings" "syscall" "time" @@ -33,12 +34,10 @@ import ( // githubClients holds all the GitHub API clients created for a server instance. type githubClients struct { - rest *gogithub.Client - restUATransp *transport.UserAgentTransport - gql *githubv4.Client - gqlHTTP *http.Client // retained for middleware to modify transport - raw *raw.Client - repoAccess *lockdown.RepoAccessCache + rest *gogithub.Client + gql *githubv4.Client + raw *raw.Client + repoAccess *lockdown.RepoAccessCache } // createGitHubClients creates all the GitHub API clients needed by the server. @@ -90,7 +89,7 @@ func createGitHubClients(cfg github.MCPServerConfig, apiHost utils.APIHostResolv // client per request (see pkg/github RequestDeps) and does not use this path. restUATransport := &transport.UserAgentTransport{ Transport: &transport.ETagTransport{Transport: http.DefaultTransport}, - Agent: fmt.Sprintf("github-mcp-server/%s", cfg.Version), + Agent: stdioUserAgent(cfg, nil), } restClient, err := newRESTClient(cfg, restUATransport, restURL.String(), uploadURL.String(), allowedHosts) if err != nil { @@ -102,7 +101,10 @@ func createGitHubClients(cfg github.MCPServerConfig, apiHost utils.APIHostResolv gqlHTTPClient := &http.Client{ Transport: &transport.BearerAuthTransport{ Transport: &transport.GraphQLFeaturesTransport{ - Transport: http.DefaultTransport, + Transport: &transport.UserAgentTransport{ + Transport: http.DefaultTransport, + Agent: stdioUserAgent(cfg, nil), + }, }, Token: cfg.Token, TokenProvider: cfg.TokenProvider, @@ -117,7 +119,7 @@ func createGitHubClients(cfg github.MCPServerConfig, apiHost utils.APIHostResolv // be large and are streamed rather than retained in memory. rawUATransport := &transport.UserAgentTransport{ Transport: http.DefaultTransport, - Agent: fmt.Sprintf("github-mcp-server/%s", cfg.Version), + Agent: stdioUserAgent(cfg, nil), } rawRESTClient, err := newRESTClient(cfg, rawUATransport, restURL.String(), uploadURL.String(), allowedHosts) if err != nil { @@ -141,12 +143,10 @@ func createGitHubClients(cfg github.MCPServerConfig, apiHost utils.APIHostResolv } return &githubClients{ - rest: restClient, - restUATransp: restUATransport, - gql: gqlClient, - gqlHTTP: gqlHTTPClient, - raw: rawClient, - repoAccess: repoAccessCache, + rest: restClient, + gql: gqlClient, + raw: rawClient, + repoAccess: repoAccessCache, }, nil } @@ -232,7 +232,7 @@ func NewStdioMCPServer(ctx context.Context, cfg github.MCPServerConfig) (*mcp.Se return nil, fmt.Errorf("failed to create GitHub MCP server: %w", err) } - ghServer.AddReceivingMiddleware(addUserAgentsMiddleware(cfg, clients.restUATransp, clients.gqlHTTP)) + ghServer.AddReceivingMiddleware(addUserAgentsMiddleware(cfg)) return ghServer, nil } @@ -441,36 +441,28 @@ func createFeatureChecker(enabledFeatures []string, insidersMode bool) inventory } } -func addUserAgentsMiddleware(cfg github.MCPServerConfig, restUATransp *transport.UserAgentTransport, gqlHTTPClient *http.Client) func(next mcp.MethodHandler) mcp.MethodHandler { +func stdioUserAgent(cfg github.MCPServerConfig, client *mcp.Implementation) string { + agent := fmt.Sprintf("github-mcp-server/%s", cfg.Version) + if client != nil { + comment := strconv.Quote(client.Name + "/" + client.Version) + comment = comment[1 : len(comment)-1] + comment = strings.NewReplacer("(", `\(`, ")", `\)`).Replace(comment) + agent += " (" + comment + ")" + } + if cfg.InsidersMode { + agent += " (insiders)" + } + return agent +} + +func addUserAgentsMiddleware(cfg github.MCPServerConfig) func(next mcp.MethodHandler) mcp.MethodHandler { return func(next mcp.MethodHandler) mcp.MethodHandler { return func(ctx context.Context, method string, request mcp.Request) (result mcp.Result, err error) { - if method != "initialize" { - return next(ctx, method, request) - } - - initializeRequest, ok := request.(*mcp.InitializeRequest) - if !ok { - return next(ctx, method, request) + var client *mcp.Implementation + if info, ok := request.(interface{ ClientInfo() *mcp.Implementation }); ok { + client = info.ClientInfo() } - - message := initializeRequest - userAgent := fmt.Sprintf( - "github-mcp-server/%s (%s/%s)", - cfg.Version, - message.Params.ClientInfo.Name, - message.Params.ClientInfo.Version, - ) - if cfg.InsidersMode { - userAgent += " (insiders)" - } - - restUATransp.Agent = userAgent - - gqlHTTPClient.Transport = &transport.UserAgentTransport{ - Transport: gqlHTTPClient.Transport, - Agent: userAgent, - } - + ctx = transport.WithUserAgent(ctx, stdioUserAgent(cfg, client)) return next(ctx, method, request) } } diff --git a/internal/ghmcp/server_test.go b/internal/ghmcp/server_test.go index 6f0e3ac3f3..b0fa1f7544 100644 --- a/internal/ghmcp/server_test.go +++ b/internal/ghmcp/server_test.go @@ -1 +1,251 @@ package ghmcp + +import ( + "context" + "encoding/json" + "net" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + "github.com/github/github-mcp-server/pkg/github" + "github.com/github/github-mcp-server/pkg/observability" + "github.com/github/github-mcp-server/pkg/observability/metrics" + "github.com/github/github-mcp-server/pkg/translations" + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestCreateGitHubClientsGraphQLUserAgent(t *testing.T) { + t.Parallel() + for _, insiders := range []bool{false, true} { + t.Run(map[bool]string{false: "release", true: "insiders"}[insiders], func(t *testing.T) { + headers := make(chan http.Header, 1) + api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + headers <- r.Header.Clone() + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":{"viewer":{"login":"octocat"}}}`)) + })) + defer api.Close() + cfg := github.MCPServerConfig{Version: "v1.2.3", Token: "fixture-token", InsidersMode: insiders} + clients, err := createGitHubClients(cfg, newStaticAPIHostResolver(t, api.URL)) + require.NoError(t, err) + var query struct { + Viewer struct{ Login string } + } + require.NoError(t, clients.gql.Query(t.Context(), &query, nil)) + want := "github-mcp-server/v1.2.3" + if insiders { + want += " (insiders)" + } + header := <-headers + assert.Equal(t, want, header.Get("User-Agent")) + assert.Equal(t, "Bearer fixture-token", header.Get("Authorization")) + }) + } +} + +func TestStdioGraphQLUserAgent(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + handshake string + clientInfo *mcp.Implementation + version string + insiders bool + want string + }{ + { + name: "SDK discovery", handshake: "sdk", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: "test-client", Version: "4.5.6"}, + want: "github-mcp-server/v1.2.3 (test-client/4.5.6)", + }, + { + name: "direct modern call", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: "test-client", Version: "4.5.6"}, + want: "github-mcp-server/v1.2.3 (test-client/4.5.6)", + }, + { + name: "modern call without optional client info", version: "v1.2.3", + want: "github-mcp-server/v1.2.3", + }, + { + name: "legacy initialize", handshake: "legacy", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: "test-client", Version: "4.5.6"}, + want: "github-mcp-server/v1.2.3 (test-client/4.5.6)", + }, + { + name: "insiders", version: "v1.2.3", insiders: true, + clientInfo: &mcp.Implementation{Name: "test-client", Version: "4.5.6"}, + want: "github-mcp-server/v1.2.3 (test-client/4.5.6) (insiders)", + }, + { + name: "different server build", version: "v1.2.4", + clientInfo: &mcp.Implementation{Name: "test-client", Version: "4.5.6"}, + want: "github-mcp-server/v1.2.4 (test-client/4.5.6)", + }, + { + name: "HTTP comment delimiters", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: `Editor (Preview) \client`, Version: `v(1)\build`}, + want: `github-mcp-server/v1.2.3 (Editor \(Preview\) \\client/v\(1\)\\build)`, + }, + { + name: "modern control characters", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: "editor(\x00)\\client\r\nX-Fake:1", Version: "1.0\tbeta\x7f"}, + want: `github-mcp-server/v1.2.3 (editor\(\x00\)\\client\r\nX-Fake:1/1.0\tbeta\x7f)`, + }, + { + name: "SDK control characters", handshake: "sdk", version: "v1.2.3", insiders: true, + clientInfo: &mcp.Implementation{Name: "editor(\x00)\\client\r\nX-Fake:1", Version: "1.0\tbeta\x7f"}, + want: `github-mcp-server/v1.2.3 (editor\(\x00\)\\client\r\nX-Fake:1/1.0\tbeta\x7f) (insiders)`, + }, + { + name: "legacy control characters", handshake: "legacy", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: "editor(\x00)\\client\r\nX-Fake:1", Version: "1.0\tbeta\x7f"}, + want: `github-mcp-server/v1.2.3 (editor\(\x00\)\\client\r\nX-Fake:1/1.0\tbeta\x7f)`, + }, + { + name: "spaces and Unicode preserved", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: "VS Code \u03b2", Version: "\u03b1 1.0"}, + want: "github-mcp-server/v1.2.3 (VS Code \u03b2/\u03b1 1.0)", + }, + { + name: "empty client metadata", version: "v1.2.3", + clientInfo: &mcp.Implementation{}, + want: "github-mcp-server/v1.2.3 (/)", + }, + { + name: "quoted client metadata", version: "v1.2.3", + clientInfo: &mcp.Implementation{Name: `"Editor"`, Version: `v"1`}, + want: `github-mcp-server/v1.2.3 (\"Editor\"/v\"1)`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + + headers := make(chan http.Header, 2) + api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + headers <- r.Header.Clone() + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":{"repository":{"isPrivate":false,"issues":{"nodes":[],"totalCount":0,"pageInfo":{"hasNextPage":false,"endCursor":""}}}}}`)) + })) + defer api.Close() + + cfg := github.MCPServerConfig{ + Version: tt.version, Token: "fixture-token", InsidersMode: tt.insiders, + Translator: translations.NullTranslationHelper, Logger: discardLogger(), + } + clients, err := createGitHubClients(cfg, newStaticAPIHostResolver(t, api.URL)) + require.NoError(t, err) + obs, err := observability.NewExporters(cfg.Logger, metrics.NewNoopMetrics()) + require.NoError(t, err) + deps := github.NewBaseDeps( + clients.rest, clients.gql, clients.raw, nil, cfg.Translator, + github.FeatureFlags{}, 5000, createFeatureChecker(nil, tt.insiders), obs, + ) + inv, err := github.NewInventory(cfg.Translator). + WithToolsets([]string{}).WithTools([]string{"list_issues"}).Build() + require.NoError(t, err) + server, err := github.NewMCPServer(ctx, &cfg, deps, inv) + require.NoError(t, err) + server.AddReceivingMiddleware(addUserAgentsMiddleware(cfg)) + var initialized atomic.Bool + server.AddReceivingMiddleware(func(next mcp.MethodHandler) mcp.MethodHandler { + return func(ctx context.Context, method string, req mcp.Request) (mcp.Result, error) { + if method == "initialize" { + initialized.Store(true) + } + return next(ctx, method, req) + } + }) + + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + require.NoError(t, clientConn.SetDeadline(time.Now().Add(10*time.Second))) + session, err := server.Connect(ctx, &mcp.IOTransport{Reader: serverConn, Writer: serverConn}, nil) + require.NoError(t, err) + defer session.Close() + params := &mcp.CallToolParams{ + Name: "list_issues", Arguments: map[string]any{"owner": "owner", "repo": "repo"}, + } + var changedClient bool + if tt.handshake == "sdk" { + client := mcp.NewClient(tt.clientInfo, nil) + cs, err := client.Connect(ctx, &mcp.IOTransport{Reader: clientConn, Writer: clientConn}, nil) + require.NoError(t, err) + defer cs.Close() + assert.Equal(t, "2026-07-28", cs.InitializeResult().ProtocolVersion) + result, err := cs.CallTool(ctx, params) + require.NoError(t, err) + require.False(t, result.IsError, "%+v", result.Content) + } else { + encoder, decoder := json.NewEncoder(clientConn), json.NewDecoder(clientConn) + id := 0 + call := func(method string, params any) json.RawMessage { + t.Helper() + id++ + require.NoError(t, encoder.Encode(map[string]any{ + "jsonrpc": "2.0", "id": id, "method": method, "params": params, + })) + var response struct { + Result json.RawMessage `json:"result"` + Error json.RawMessage `json:"error"` + } + require.NoError(t, decoder.Decode(&response)) + require.Empty(t, response.Error) + return response.Result + } + if tt.handshake == "legacy" { + call("initialize", &mcp.InitializeParams{ + ProtocolVersion: "2025-11-25", Capabilities: &mcp.ClientCapabilities{}, + ClientInfo: tt.clientInfo, + }) + require.NoError(t, encoder.Encode(map[string]any{ + "jsonrpc": "2.0", "method": "notifications/initialized", + })) + } else { + params.Meta = mcp.Meta{ + mcp.MetaKeyProtocolVersion: "2026-07-28", mcp.MetaKeyClientCapabilities: map[string]any{}, + } + if tt.clientInfo != nil { + params.Meta[mcp.MetaKeyClientInfo] = tt.clientInfo + } + } + var result mcp.CallToolResult + require.NoError(t, json.Unmarshal(call("tools/call", params), &result)) + require.False(t, result.IsError, "%+v", result.Content) + if tt.handshake != "legacy" { + params.Meta[mcp.MetaKeyClientInfo] = &mcp.Implementation{Name: "other-client", Version: "7.8.9"} + require.NoError(t, json.Unmarshal(call("tools/call", params), &result)) + require.False(t, result.IsError, "%+v", result.Content) + changedClient = true + } + } + + assert.Equal(t, tt.handshake == "legacy", initialized.Load()) + select { + case header := <-headers: + assert.Equal(t, tt.want, header.Get("User-Agent")) + assert.Equal(t, "Bearer fixture-token", header.Get("Authorization")) + assert.Equal(t, "issue_fields, repo_issue_fields", header.Get("GraphQL-Features")) + case <-ctx.Done(): + t.Fatal("GraphQL request did not reach the local HTTP server") + } + if changedClient { + want := "github-mcp-server/" + tt.version + " (other-client/7.8.9)" + if tt.insiders { + want += " (insiders)" + } + assert.Equal(t, want, (<-headers).Get("User-Agent")) + } + }) + } +} diff --git a/pkg/http/transport/user_agent.go b/pkg/http/transport/user_agent.go index a489941cce..a40b4cc45b 100644 --- a/pkg/http/transport/user_agent.go +++ b/pkg/http/transport/user_agent.go @@ -1,11 +1,19 @@ package transport import ( + "context" "net/http" "github.com/github/github-mcp-server/pkg/http/headers" ) +type userAgentKey struct{} + +// WithUserAgent supplies a request-scoped identity without mutating a shared transport. +func WithUserAgent(ctx context.Context, agent string) context.Context { + return context.WithValue(ctx, userAgentKey{}, agent) +} + type UserAgentTransport struct { Transport http.RoundTripper Agent string @@ -13,6 +21,10 @@ type UserAgentTransport struct { func (t *UserAgentTransport) RoundTrip(req *http.Request) (*http.Response, error) { req = req.Clone(req.Context()) - req.Header.Set(headers.UserAgentHeader, t.Agent) + agent := t.Agent + if scoped, ok := req.Context().Value(userAgentKey{}).(string); ok { + agent = scoped + } + req.Header.Set(headers.UserAgentHeader, agent) return t.Transport.RoundTrip(req) } diff --git a/pkg/http/transport/user_agent_test.go b/pkg/http/transport/user_agent_test.go new file mode 100644 index 0000000000..58eee9b699 --- /dev/null +++ b/pkg/http/transport/user_agent_test.go @@ -0,0 +1,48 @@ +package transport + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestUserAgentTransportRequestIsolation(t *testing.T) { + t.Parallel() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + _, _ = io.WriteString(w, req.UserAgent()) + })) + t.Cleanup(server.Close) + transport := &UserAgentTransport{ + Transport: http.DefaultTransport, + Agent: "github-mcp-server/remote-abcdef", + } + client := &http.Client{Transport: transport} + + for _, agent := range []string{"", "github-mcp-server/v1.2.3 (first/1)", "github-mcp-server/v1.2.3 (second/2)"} { + t.Run(agent, func(t *testing.T) { + t.Parallel() + ctx := context.Background() + want := transport.Agent + if agent != "" { + ctx = WithUserAgent(ctx, agent) + want = agent + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, server.URL, nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "original") + resp, err := client.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, want, string(body)) + assert.Equal(t, "original", req.UserAgent()) + assert.Equal(t, "github-mcp-server/remote-abcdef", transport.Agent) + }) + } +}