Skip to content
Closed
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
4 changes: 3 additions & 1 deletion internal/ghmcp/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
36 changes: 36 additions & 0 deletions pkg/github/context_tools.go
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
1 change: 1 addition & 0 deletions pkg/github/tools.go
Original file line number Diff line number Diff line change
Expand Up @@ -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),

Expand Down
18 changes: 18 additions & 0 deletions pkg/http/transport/ratelimit.go
Original file line number Diff line number Diff line change
@@ -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)
}
31 changes: 31 additions & 0 deletions pkg/http/transport/ratelimit_test.go
Original file line number Diff line number Diff line change
@@ -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"))
}