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..bbec00b 100644 --- a/internal/strava/client.go +++ b/internal/strava/client.go @@ -91,29 +91,30 @@ 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, 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. @@ -146,7 +147,8 @@ func (c *Client) SetTokenURL(u string) { } // doRequest executes an authenticated HTTP request with automatic token refresh. -func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io.Reader, contentType string) ([]byte, error) { +// 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) @@ -160,7 +162,14 @@ func (c *Client) doRequest(ctx context.Context, method, fullURL string, body io. } } - respBody, err := c.executeRequest(ctx, method, fullURL, body, 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 @@ -169,8 +178,8 @@ 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, bodyReader(), contentType, tokens.AccessToken) } return nil, err } 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() 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))