From b929243bb827c2f6ccdfa198cbd778545d0cd1c9 Mon Sep 17 00:00:00 2001 From: Stealinglight Date: Sat, 4 Apr 2026 19:33:34 -0700 Subject: [PATCH 1/2] fix: patch 401 retry body-reuse bug, XSS in OAuth callback, and clean up codebase - Fix io.Reader consumed on first request attempt causing 401 retries to send empty bodies for POST/PUT/multipart requests. Buffer body as []byte and create fresh readers for each attempt. - Fix reflected XSS in OAuth error callback by escaping user-controlled Strava error parameter with html.EscapeString(). - Consolidate Post/Put into shared jsonRequest helper. - Remove dead Config.Debug field (debug handled via CLI flag). - Remove stale section-divider comments and outdated references. - Use loop for optional field forwarding in HandleCreateActivity for consistency with HandleUpdateActivity. --- internal/auth/oauth.go | 3 ++- internal/config/config.go | 1 - internal/server/server.go | 1 - internal/strava/client.go | 39 ++++++++++++++++++++++++++---------- internal/tools/activities.go | 25 ++++------------------- internal/tools/athlete.go | 6 ------ internal/tools/clubs.go | 6 ------ internal/tools/helpers.go | 2 +- internal/tools/register.go | 6 ------ internal/tools/streams.go | 6 ------ internal/tools/uploads.go | 6 ------ 11 files changed, 35 insertions(+), 66 deletions(-) diff --git a/internal/auth/oauth.go b/internal/auth/oauth.go index 276b7b2..849b34c 100644 --- a/internal/auth/oauth.go +++ b/internal/auth/oauth.go @@ -6,6 +6,7 @@ import ( "encoding/hex" "encoding/json" "fmt" + "html" "io" "log/slog" "net/http" @@ -109,7 +110,7 @@ func NewCallbackHandler(expectedState string, codeCh chan<- string, errCh chan<- if stravaErr := query.Get("error"); stravaErr != "" { errCh <- fmt.Errorf("Strava authorization error: %s", stravaErr) w.WriteHeader(http.StatusBadRequest) - fmt.Fprintf(w, errorPageHTML, stravaErr) + fmt.Fprintf(w, errorPageHTML, html.EscapeString(stravaErr)) return } diff --git a/internal/config/config.go b/internal/config/config.go index 5020006..37f6974 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -11,7 +11,6 @@ type Config struct { ClientID string ClientSecret string TokenPath string - Debug bool } // Load reads configuration from environment variables and returns a Config. diff --git a/internal/server/server.go b/internal/server/server.go index 7099512..b227803 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -7,7 +7,6 @@ import ( ) // New creates a new MCP server with the given version string and Strava client. -// Tools are registered via tools.RegisterAll (empty in Phase 1, populated in Phase 2). func New(version string, client *strava.Client) *mcpserver.MCPServer { s := mcpserver.NewMCPServer( "strava-mcp", diff --git a/internal/strava/client.go b/internal/strava/client.go index a5d1aaf..2bfe999 100644 --- a/internal/strava/client.go +++ b/internal/strava/client.go @@ -91,22 +91,20 @@ func (c *Client) Get(ctx context.Context, path string, params map[string]string) // Post makes an authenticated POST request to the Strava API. func (c *Client) Post(ctx context.Context, path string, body interface{}) ([]byte, error) { - fullURL := c.baseURL + path - jsonBody, err := json.Marshal(body) - if err != nil { - return nil, fmt.Errorf("marshal request body: %w", err) - } - return c.doRequest(ctx, http.MethodPost, fullURL, bytes.NewReader(jsonBody), "application/json") + return c.jsonRequest(ctx, http.MethodPost, path, body) } // Put makes an authenticated PUT request to the Strava API. func (c *Client) Put(ctx context.Context, path string, body interface{}) ([]byte, error) { - fullURL := c.baseURL + path + return c.jsonRequest(ctx, http.MethodPut, path, body) +} + +func (c *Client) jsonRequest(ctx context.Context, method, path string, body interface{}) ([]byte, error) { jsonBody, err := json.Marshal(body) if err != nil { return nil, fmt.Errorf("marshal request body: %w", err) } - return c.doRequest(ctx, http.MethodPut, fullURL, bytes.NewReader(jsonBody), "application/json") + return c.doRequest(ctx, method, c.baseURL+path, bytes.NewReader(jsonBody), "application/json") } // PostMultipart makes an authenticated POST request with a pre-built multipart body. @@ -146,7 +144,18 @@ func (c *Client) SetTokenURL(u string) { } // doRequest executes an authenticated HTTP request with automatic token refresh. +// The body is buffered so it can be replayed on 401 retry. func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io.Reader, contentType string) ([]byte, error) { + // Buffer the body so we can replay it on 401 retry. + var bodyBytes []byte + if body != nil { + var err error + bodyBytes, err = io.ReadAll(body) + if err != nil { + return nil, fmt.Errorf("buffer request body: %w", err) + } + } + tokens, err := c.tokenStore.Read() if err != nil { return nil, fmt.Errorf("read tokens: %w", err) @@ -160,7 +169,7 @@ func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io. } } - respBody, err := c.executeRequest(ctx, method, fullURL, body, contentType, tokens.AccessToken) + respBody, err := c.executeRequest(ctx, method, fullURL, newReader(bodyBytes), contentType, tokens.AccessToken) if err != nil { // Check for 401 — retry once after refresh var stravaErr *StravaError @@ -169,14 +178,22 @@ func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io. if refreshErr != nil { return nil, fmt.Errorf("token refresh after 401: %w", refreshErr) } - // Retry with new token — if this also fails, return the error directly - return c.executeRequest(ctx, method, fullURL, body, contentType, tokens.AccessToken) + // Retry with new token and a fresh reader + return c.executeRequest(ctx, method, fullURL, newReader(bodyBytes), contentType, tokens.AccessToken) } return nil, err } return respBody, nil } +// newReader returns a reader for the given bytes, or nil if the slice is nil. +func newReader(b []byte) io.Reader { + if b == nil { + return nil + } + return bytes.NewReader(b) +} + // executeRequest builds and executes a single HTTP request. func (c *Client) executeRequest(ctx context.Context, method, fullURL string, body io.Reader, contentType, accessToken string) ([]byte, error) { req, err := http.NewRequestWithContext(ctx, method, fullURL, body) diff --git a/internal/tools/activities.go b/internal/tools/activities.go index b6da55b..7f10fc4 100644 --- a/internal/tools/activities.go +++ b/internal/tools/activities.go @@ -11,8 +11,6 @@ import ( "github.com/Stealinglight/StravaMCP/internal/strava" ) -// --- Tool Definitions --- - var getActivitiesTool = mcp.NewTool("strava_get_activities", mcp.WithDescription(`Retrieves the authenticated athlete's activities. @@ -161,8 +159,6 @@ Useful for: mcp.WithNumber("id", mcp.Description("The ID of the activity"), mcp.Required()), ) -// --- Handler Functions --- - // HandleGetActivities returns a handler for the get_activities tool. func HandleGetActivities(client *strava.Client) server.ToolHandlerFunc { return func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { @@ -241,21 +237,10 @@ func HandleCreateActivity(client *strava.Client) server.ToolHandlerFunc { "elapsed_time": elapsedTime, } - // Optional fields - if v, ok := args["type"]; ok { - body["type"] = v - } - if v, ok := args["description"]; ok { - body["description"] = v - } - if v, ok := args["distance"]; ok { - body["distance"] = v - } - if v, ok := args["trainer"]; ok { - body["trainer"] = v - } - if v, ok := args["commute"]; ok { - body["commute"] = v + for _, field := range []string{"type", "description", "distance", "trainer", "commute"} { + if v, ok := args[field]; ok { + body[field] = v + } } data, err := client.Post(ctx, "/activities", body) @@ -318,8 +303,6 @@ func HandleGetActivityZones(client *strava.Client) server.ToolHandlerFunc { } } -// --- Registration --- - // registerActivities registers all activity tools with the MCP server. func registerActivities(s *server.MCPServer, client *strava.Client) { s.AddTool(getActivitiesTool, HandleGetActivities(client)) diff --git a/internal/tools/athlete.go b/internal/tools/athlete.go index 2b7af32..c7de34b 100644 --- a/internal/tools/athlete.go +++ b/internal/tools/athlete.go @@ -11,8 +11,6 @@ import ( "github.com/Stealinglight/StravaMCP/internal/strava" ) -// --- Tool Definitions --- - var getAthleteTool = mcp.NewTool("strava_get_athlete", mcp.WithDescription(`Retrieves the authenticated athlete's profile information. @@ -75,8 +73,6 @@ Statistics include: mcp.WithNumber("id", mcp.Description("Athlete ID (optional - defaults to authenticated athlete)")), ) -// --- Handler Functions --- - // HandleGetAthlete returns a handler for the get_athlete tool. func HandleGetAthlete(client *strava.Client) server.ToolHandlerFunc { return func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { @@ -118,8 +114,6 @@ func HandleGetAthleteStats(client *strava.Client) server.ToolHandlerFunc { } } -// --- Registration --- - // registerAthlete registers all athlete tools with the MCP server. func registerAthlete(s *server.MCPServer, client *strava.Client) { s.AddTool(getAthleteTool, HandleGetAthlete(client)) diff --git a/internal/tools/clubs.go b/internal/tools/clubs.go index 26291e6..922de6a 100644 --- a/internal/tools/clubs.go +++ b/internal/tools/clubs.go @@ -11,8 +11,6 @@ import ( "github.com/Stealinglight/StravaMCP/internal/strava" ) -// --- Tool Definition --- - var getClubActivitiesTool = mcp.NewTool("strava_get_club_activities", mcp.WithDescription(`Retrieves recent activities from members of a specific club. @@ -44,8 +42,6 @@ Note: Only shows activities from club members who have their privacy settings se mcp.WithNumber("per_page", mcp.Description("Number of items per page (1-200, default 30)")), ) -// --- Handler Function --- - // HandleGetClubActivities returns a handler for the get_club_activities tool. func HandleGetClubActivities(client *strava.Client) server.ToolHandlerFunc { return func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { @@ -71,8 +67,6 @@ func HandleGetClubActivities(client *strava.Client) server.ToolHandlerFunc { } } -// --- Registration --- - // registerClubs registers all club tools with the MCP server. func registerClubs(s *server.MCPServer, client *strava.Client) { s.AddTool(getClubActivitiesTool, HandleGetClubActivities(client)) diff --git a/internal/tools/helpers.go b/internal/tools/helpers.go index bd58773..9f31ba0 100644 --- a/internal/tools/helpers.go +++ b/internal/tools/helpers.go @@ -21,7 +21,7 @@ func FormatResponse(data []byte, client *strava.Client) *mcp.CallToolResult { result := pretty.String() - // Append rate limit warning if usage is high (D-02) + // Append rate limit warning if usage is high if warning := client.RateLimitWarning(); warning != "" { result += "\n\n" + warning } diff --git a/internal/tools/register.go b/internal/tools/register.go index b6478b6..04bace2 100644 --- a/internal/tools/register.go +++ b/internal/tools/register.go @@ -6,12 +6,6 @@ import ( ) // RegisterAll registers all MCP tools with the server. -// Registers 11 tools across 5 resource categories: -// - Activities: strava_get_activities, strava_get_activity_by_id, strava_create_activity, strava_update_activity, strava_get_activity_zones -// - Athlete: strava_get_athlete, strava_get_athlete_stats -// - Streams: strava_get_activity_streams -// - Clubs: strava_get_club_activities -// - Uploads: strava_create_upload, strava_get_upload func RegisterAll(s *server.MCPServer, client *strava.Client) { registerActivities(s, client) registerAthlete(s, client) diff --git a/internal/tools/streams.go b/internal/tools/streams.go index 70586de..9630edd 100644 --- a/internal/tools/streams.go +++ b/internal/tools/streams.go @@ -18,8 +18,6 @@ var streamTypes = []string{ "heartrate", "cadence", "watts", "temp", "moving", "grade_smooth", } -// --- Tool Definition --- - var getActivityStreamsTool = mcp.NewTool("strava_get_activity_streams", mcp.WithDescription(`**[TELEMETRY & DEEP ANALYSIS]** Retrieves time-series sensor data (streams) from an activity. @@ -100,8 +98,6 @@ var getActivityStreamsTool = mcp.NewTool("strava_get_activity_streams", mcp.WithBoolean("key_by_type", mcp.Description("Return streams as an object keyed by type (default: true)")), ) -// --- Handler Function --- - // HandleGetActivityStreams returns a handler for the get_activity_streams tool. func HandleGetActivityStreams(client *strava.Client) server.ToolHandlerFunc { return func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { @@ -141,8 +137,6 @@ func HandleGetActivityStreams(client *strava.Client) server.ToolHandlerFunc { } } -// --- Registration --- - // registerStreams registers all streams tools with the MCP server. func registerStreams(s *server.MCPServer, client *strava.Client) { s.AddTool(getActivityStreamsTool, HandleGetActivityStreams(client)) diff --git a/internal/tools/uploads.go b/internal/tools/uploads.go index afcad1f..fc6bb27 100644 --- a/internal/tools/uploads.go +++ b/internal/tools/uploads.go @@ -17,8 +17,6 @@ import ( "github.com/Stealinglight/StravaMCP/internal/strava" ) -// --- Tool Definitions --- - var createUploadTool = mcp.NewTool("strava_create_upload", mcp.WithDescription(`Uploads a new activity file to Strava. @@ -106,8 +104,6 @@ After an athlete uploads a workout file: mcp.WithNumber("id", mcp.Description("The ID of the upload"), mcp.Required()), ) -// --- Handler Functions --- - // validDataTypes lists all valid upload data types. var validDataTypes = map[string]bool{ "fit": true, "tcx": true, "gpx": true, @@ -220,8 +216,6 @@ func HandleGetUpload(client *strava.Client) server.ToolHandlerFunc { } } -// --- Registration --- - // registerUploads registers all upload tools with the MCP server. func registerUploads(s *server.MCPServer, client *strava.Client) { s.AddTool(createUploadTool, HandleCreateUpload(client)) From 1a09191da53c2d939aa037c38d8fbc660b07e335 Mon Sep 17 00:00:00 2001 From: Stealinglight Date: Sat, 4 Apr 2026 19:57:23 -0700 Subject: [PATCH 2/2] refactor: accept []byte in doRequest to avoid double-copy and add 401 retry body test Address Copilot PR feedback: - Change doRequest to accept []byte instead of io.Reader, eliminating the extra io.ReadAll buffering step. JSON requests pass marshaled bytes directly; PostMultipart reads upfront at the call site. - Use a bodyReader closure to create fresh readers per attempt. - Add TestPostReplaysBodyOn401Retry to verify POST body is fully replayed after token refresh, preventing regression of the consumed-reader bug. --- internal/strava/client.go | 42 +++++++++------------- internal/strava/client_test.go | 64 ++++++++++++++++++++++++++++++++++ 2 files changed, 81 insertions(+), 25 deletions(-) diff --git a/internal/strava/client.go b/internal/strava/client.go index 2bfe999..bbec00b 100644 --- a/internal/strava/client.go +++ b/internal/strava/client.go @@ -104,14 +104,17 @@ func (c *Client) jsonRequest(ctx context.Context, method, path string, body inte if err != nil { return nil, fmt.Errorf("marshal request body: %w", err) } - return c.doRequest(ctx, method, c.baseURL+path, bytes.NewReader(jsonBody), "application/json") + return c.doRequest(ctx, method, c.baseURL+path, jsonBody, "application/json") } // PostMultipart makes an authenticated POST request with a pre-built multipart body. // The contentType must include the multipart boundary (use writer.FormDataContentType()). func (c *Client) PostMultipart(ctx context.Context, path string, body io.Reader, contentType string) ([]byte, error) { - fullURL := c.baseURL + path - return c.doRequest(ctx, http.MethodPost, fullURL, body, contentType) + bodyBytes, err := io.ReadAll(body) + if err != nil { + return nil, fmt.Errorf("buffer multipart body: %w", err) + } + return c.doRequest(ctx, http.MethodPost, c.baseURL+path, bodyBytes, contentType) } // GetRateLimits returns the current rate limit state. @@ -144,18 +147,8 @@ func (c *Client) SetTokenURL(u string) { } // doRequest executes an authenticated HTTP request with automatic token refresh. -// The body is buffered so it can be replayed on 401 retry. -func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io.Reader, contentType string) ([]byte, error) { - // Buffer the body so we can replay it on 401 retry. - var bodyBytes []byte - if body != nil { - var err error - bodyBytes, err = io.ReadAll(body) - if err != nil { - return nil, fmt.Errorf("buffer request body: %w", err) - } - } - +// The body is accepted as []byte so it can be replayed on 401 retry without re-reading. +func (c *Client) doRequest(ctx context.Context, method, fullURL string, body []byte, contentType string) ([]byte, error) { tokens, err := c.tokenStore.Read() if err != nil { return nil, fmt.Errorf("read tokens: %w", err) @@ -169,7 +162,14 @@ func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io. } } - respBody, err := c.executeRequest(ctx, method, fullURL, newReader(bodyBytes), contentType, tokens.AccessToken) + bodyReader := func() io.Reader { + if body == nil { + return nil + } + return bytes.NewReader(body) + } + + respBody, err := c.executeRequest(ctx, method, fullURL, bodyReader(), contentType, tokens.AccessToken) if err != nil { // Check for 401 — retry once after refresh var stravaErr *StravaError @@ -179,21 +179,13 @@ func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io. return nil, fmt.Errorf("token refresh after 401: %w", refreshErr) } // Retry with new token and a fresh reader - return c.executeRequest(ctx, method, fullURL, newReader(bodyBytes), contentType, tokens.AccessToken) + return c.executeRequest(ctx, method, fullURL, bodyReader(), contentType, tokens.AccessToken) } return nil, err } return respBody, nil } -// newReader returns a reader for the given bytes, or nil if the slice is nil. -func newReader(b []byte) io.Reader { - if b == nil { - return nil - } - return bytes.NewReader(b) -} - // executeRequest builds and executes a single HTTP request. func (c *Client) executeRequest(ctx context.Context, method, fullURL string, body io.Reader, contentType, accessToken string) ([]byte, error) { req, err := http.NewRequestWithContext(ctx, method, fullURL, body) diff --git a/internal/strava/client_test.go b/internal/strava/client_test.go index 543ec08..2ae77df 100644 --- a/internal/strava/client_test.go +++ b/internal/strava/client_test.go @@ -616,6 +616,70 @@ func TestPostMultipartReturnsStravaErrorOn4xx(t *testing.T) { } } +// Test 14: Post() replays full body on 401 retry (regression test for body-reuse bug) +func TestPostReplaysBodyOn401Retry(t *testing.T) { + var requestCount atomic.Int32 + var bodies []string + var mu sync.Mutex + + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + mu.Lock() + bodies = append(bodies, string(body)) + mu.Unlock() + + count := requestCount.Add(1) + if count == 1 { + w.WriteHeader(http.StatusUnauthorized) + w.Write([]byte(`{"message":"Authorization Error"}`)) + return + } + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{"id":123}`)) + })) + defer apiSrv.Close() + + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(auth.Tokens{ + AccessToken: "refreshed-token", + RefreshToken: "new-refresh", + ExpiresAt: time.Now().Add(6 * time.Hour).Unix(), + }) + })) + defer tokenSrv.Close() + + store := newMockTokenStore(&auth.Tokens{ + AccessToken: "stale-token", + RefreshToken: "refresh-token", + ExpiresAt: time.Now().Add(1 * time.Hour).Unix(), + }, false) + + cfg := &config.Config{ClientID: "id", ClientSecret: "secret"} + client := strava.NewClient(cfg, store, testLogger()) + client.SetBaseURL(apiSrv.URL) + client.SetTokenURL(tokenSrv.URL) + + payload := map[string]string{"name": "Morning Run", "description": "Easy 5K"} + _, err := client.Post(context.Background(), "/activities", payload) + if err != nil { + t.Fatalf("Post() error: %v", err) + } + + if requestCount.Load() != 2 { + t.Fatalf("API requests = %d, want 2 (initial + retry)", requestCount.Load()) + } + + // Both requests must have received the full JSON body + for i, body := range bodies { + if !strings.Contains(body, "Morning Run") { + t.Errorf("request %d body = %q, want to contain 'Morning Run'", i+1, body) + } + if !strings.Contains(body, "Easy 5K") { + t.Errorf("request %d body = %q, want to contain 'Easy 5K'", i+1, body) + } + } +} + // TestNewClientReturnsNonNil verifies the constructor works func TestNewClientReturnsNonNil(t *testing.T) { dir := t.TempDir()