Skip to content
Merged
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
3 changes: 2 additions & 1 deletion internal/auth/oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"encoding/hex"
"encoding/json"
"fmt"
"html"
"io"
"log/slog"
"net/http"
Expand Down Expand Up @@ -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
}

Expand Down
1 change: 0 additions & 1 deletion internal/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
1 change: 0 additions & 1 deletion internal/server/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
37 changes: 23 additions & 14 deletions internal/strava/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
Comment on lines +102 to 108

Copilot AI Apr 5, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

jsonRequest marshals to jsonBody and then passes a reader to doRequest, which immediately io.ReadAlls it back into a new byte slice. For JSON requests this introduces an extra allocation/copy on every call. Consider changing doRequest to accept []byte (or a func() io.Reader factory) so JSON requests can reuse the already-marshaled bytes without a second read/copy, while still allowing fresh readers per attempt.

Copilot uses AI. Check for mistakes.

// 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.
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand All @@ -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
}
Expand Down
64 changes: 64 additions & 0 deletions internal/strava/client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
25 changes: 4 additions & 21 deletions internal/tools/activities.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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))
Expand Down
6 changes: 0 additions & 6 deletions internal/tools/athlete.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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))
Expand Down
6 changes: 0 additions & 6 deletions internal/tools/clubs.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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) {
Expand All @@ -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))
Expand Down
2 changes: 1 addition & 1 deletion internal/tools/helpers.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
6 changes: 0 additions & 6 deletions internal/tools/register.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
6 changes: 0 additions & 6 deletions internal/tools/streams.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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))
Expand Down
6 changes: 0 additions & 6 deletions internal/tools/uploads.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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))
Expand Down
Loading