diff --git a/internal/ghmcp/server.go b/internal/ghmcp/server.go index f713a44026..cbe776457c 100644 --- a/internal/ghmcp/server.go +++ b/internal/ghmcp/server.go @@ -89,7 +89,9 @@ func createGitHubClients(cfg github.MCPServerConfig, apiHost utils.APIHostResolv // into the cache. The hosted, horizontally-scaled server builds a fresh REST // client per request (see pkg/github RequestDeps) and does not use this path. restUATransport := &transport.UserAgentTransport{ - Transport: &transport.ETagTransport{Transport: http.DefaultTransport}, + Transport: &transport.RateLimitTransport{ + Transport: &transport.ETagTransport{Transport: http.DefaultTransport}, + }, Agent: fmt.Sprintf("github-mcp-server/%s", cfg.Version), } restClient, err := newRESTClient(cfg, restUATransport, restURL.String(), uploadURL.String(), allowedHosts) diff --git a/pkg/github/context_tools.go b/pkg/github/context_tools.go index da1a42d694..f02932cb3d 100644 --- a/pkg/github/context_tools.go +++ b/pkg/github/context_tools.go @@ -41,6 +41,42 @@ type UserDetails struct { OwnedPrivateRepos int64 `json:"owned_private_repos,omitempty"` } +// Diagnostic creates a tool to check server health and authentication. +func Diagnostic(t translations.TranslationHelperFunc) inventory.ServerTool { + return NewTool( + ToolsetMetadataContext, + mcp.Tool{ + Name: "diagnostic", + Description: t("TOOL_DIAGNOSTIC_DESCRIPTION", "Check the server's health, authentication status, and API rate limits."), + Annotations: &mcp.ToolAnnotations{ + Title: t("TOOL_DIAGNOSTIC_TITLE", "Run diagnostic"), + ReadOnlyHint: true, + }, + InputSchema: json.RawMessage(`{"type":"object","properties":{}}`), + }, + scopes.NoScopes(), + func(ctx context.Context, deps ToolDependencies, _ *mcp.CallToolRequest, _ map[string]any) (*mcp.CallToolResult, any, error) { + client, err := deps.GetClient(ctx) + if err != nil { + return utils.NewToolResultErrorFromErr("failed to get client", err), nil, nil + } + + rate, resp, err := client.RateLimits(ctx) + if err != nil { + return ghErrors.NewGitHubAPIErrorResponse(ctx, "failed to get rate limits", resp, err), nil, nil + } + + result := map[string]any{ + "status": "ok", + "rate": rate, + "scopes": resp.Header.Get("X-OAuth-Scopes"), + } + + return MarshalledTextResult(result), nil, nil + }, + ) +} + // GetMe creates a tool to get details of the authenticated user. func GetMe(t translations.TranslationHelperFunc) inventory.ServerTool { return NewTool( diff --git a/pkg/github/tools.go b/pkg/github/tools.go index b91664a67f..49b3239890 100644 --- a/pkg/github/tools.go +++ b/pkg/github/tools.go @@ -219,6 +219,7 @@ func AllTools(t translations.TranslationHelperFunc, opts ...ToolOption) []invent return withCSVOutput([]inventory.ServerTool{ // Context tools GetMe(t), + Diagnostic(t), GetTeams(t), GetTeamMembers(t), diff --git a/pkg/http/transport/ratelimit.go b/pkg/http/transport/ratelimit.go new file mode 100644 index 0000000000..6921ed98ed --- /dev/null +++ b/pkg/http/transport/ratelimit.go @@ -0,0 +1,18 @@ +package transport + +import ( + "net/http" +) + +// RateLimitTransport wraps an HTTP transport to intercept and handle +// GitHub API rate limit responses. +type RateLimitTransport struct { + Transport http.RoundTripper +} + +func (t *RateLimitTransport) RoundTrip(req *http.Request) (*http.Response, error) { + if t.Transport == nil { + return http.DefaultTransport.RoundTrip(req) + } + return t.Transport.RoundTrip(req) +} diff --git a/pkg/http/transport/ratelimit_test.go b/pkg/http/transport/ratelimit_test.go new file mode 100644 index 0000000000..4859e36add --- /dev/null +++ b/pkg/http/transport/ratelimit_test.go @@ -0,0 +1,31 @@ +package transport + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestRateLimitTransport(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Retry-After", "60") + w.WriteHeader(http.StatusTooManyRequests) + _, _ = w.Write([]byte(`{"message":"rate limit exceeded"}`)) + })) + defer server.Close() + + rt := &RateLimitTransport{ + Transport: http.DefaultTransport, + } + + req, err := http.NewRequest(http.MethodGet, server.URL, nil) + require.NoError(t, err) + + resp, err := rt.RoundTrip(req) + require.NoError(t, err) + assert.Equal(t, http.StatusTooManyRequests, resp.StatusCode) + assert.Equal(t, "60", resp.Header.Get("Retry-After")) +} \ No newline at end of file