Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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
29 changes: 25 additions & 4 deletions service/spotify/spotify.go
Original file line number Diff line number Diff line change
Expand Up @@ -285,11 +285,15 @@ func (s *Service) refreshTokenForUser(user *models.User) (string, error) {
s.mu.Lock()
delete(s.userTokens, userID)
s.mu.Unlock()
// Also clear the bad refresh token from the DB
updateErr := s.DB.UpdateUserToken(userID, "", "", time.Now().UTC()) // Clear tokens
if updateErr != nil {
s.logger.Printf("Failed to clear bad refresh token for user %d: %v", userID, updateErr)

// Only discard the refresh token when Spotify says it is genuinely dead.
if isRefreshTokenRejected(resp.StatusCode, body) {
if updateErr := s.DB.UpdateUserToken(userID, "", "", time.Now().UTC()); updateErr != nil {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
s.logger.Printf("Failed to clear bad refresh token for user %d: %v", userID, updateErr)
}
return "", fmt.Errorf("spotify refresh token rejected for user %d (%d): %s", userID, resp.StatusCode, string(body))
}

return "", fmt.Errorf("spotify token refresh failed (%d): %s", resp.StatusCode, string(body))
}

Expand Down Expand Up @@ -326,6 +330,23 @@ func (s *Service) refreshTokenForUser(user *models.User) (string, error) {
return tokenResponse.AccessToken, nil
}

// isRefreshTokenRejected reports whether Spotify permanently rejected the
// refresh token, as opposed to failing for a transient reason.
func isRefreshTokenRejected(statusCode int, body []byte) bool {
if statusCode != http.StatusBadRequest {
return false
}

var errorResponse struct {
Error string `json:"error"`
}
if err := json.Unmarshal(body, &errorResponse); err != nil {
return false
}

return errorResponse.Error == "invalid_grant"
}

// RefreshToken attempts to refresh the token for a given user ID.
// It's less commonly needed now refreshTokenInner handles fetching the user.
func (s *Service) RefreshToken(userID int64) error {
Expand Down
158 changes: 158 additions & 0 deletions service/spotify/spotify_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"testing"
"time"

"github.com/spf13/viper"
"github.com/teal-fm/piper/db"
"github.com/teal-fm/piper/models"
"github.com/teal-fm/piper/session"
Expand Down Expand Up @@ -1413,3 +1414,160 @@ func TestGenerateLocalHash(t *testing.T) {
}
})
}

// ===== Token Refresh Tests =====

// stubRoundTripper answers every request with a canned response, standing in
// for accounts.spotify.com.
type stubRoundTripper struct {
statusCode int
body string
}

func (s stubRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: s.statusCode,
Body: io.NopCloser(strings.NewReader(s.body)),
Header: make(http.Header),
Request: req,
}, nil
}

func newRefreshTestService(t *testing.T, database *db.DB, statusCode int, body string) *Service {
t.Helper()

previousID := viper.Get("spotify.client_id")
previousSecret := viper.Get("spotify.client_secret")
t.Cleanup(func() {
viper.Set("spotify.client_id", previousID)
viper.Set("spotify.client_secret", previousSecret)
})

viper.Set("spotify.client_id", "id")
viper.Set("spotify.client_secret", "secret")

service := newTestService(database, &mockPlayingNowService{})
service.httpClient = &http.Client{
Transport: stubRoundTripper{statusCode: statusCode, body: body},
}
return service
}

// A 502 from Spotify is transient -- we should keep the refresh token & retry later.
func TestRefreshTokenForUser_TransientFailureKeepsRefreshToken(t *testing.T) {
database := setupTestDB(t)
userID := createTestUser(t, database)

user, err := database.AddSpotifySession(userID, "RUSH", "moving@pict.ures", "rush", "access", "refresh", time.Now().UTC().Add(-time.Hour))
if err != nil {
t.Fatalf("Failed to link Spotify session: %v", err)
}

service := newRefreshTestService(t, database, http.StatusBadGateway, "<html><head><title>502 Server Error</title></head></html>")
service.userTokens[userID] = "access"

if _, err := service.refreshTokenForUser(user); err == nil {
t.Fatal("expected the refresh to fail")
}

reloaded, err := database.GetUserByID(userID)
if err != nil {
t.Fatalf("Failed to reload user: %v", err)
}
if reloaded.RefreshToken == nil || *reloaded.RefreshToken != "refresh" {
t.Errorf("RefreshToken = %v, want it kept for the next retry", reloaded.RefreshToken)
}

if _, exists := service.userTokens[userID]; exists {
t.Error("expected the stale cached access token to be dropped")
}
}

// A legitimately bad refresh token should clear the token from the DB.
func TestRefreshTokenForUser_InvalidGrantClearsRefreshToken(t *testing.T) {
database := setupTestDB(t)
userID := createTestUser(t, database)

user, err := database.AddSpotifySession(userID, "YES", "close@to.the.edge", "yes", "access", "refresh", time.Now().UTC().Add(-time.Hour))
if err != nil {
t.Fatalf("Failed to link Spotify session: %v", err)
}

service := newRefreshTestService(t, database, http.StatusBadRequest, `{"error":"invalid_grant","error_description":"Refresh token revoked"}`)
service.userTokens[userID] = "access"

if _, err := service.refreshTokenForUser(user); err == nil {
t.Fatal("expected the refresh to fail")
}

reloaded, err := database.GetUserByID(userID)
if err != nil {
t.Fatalf("Failed to reload user: %v", err)
}
if reloaded.RefreshToken != nil && *reloaded.RefreshToken != "" {
t.Errorf("RefreshToken = %v, want the dead token cleared", *reloaded.RefreshToken)
}
}

func TestIsRefreshTokenRejected(t *testing.T) {
testCases := []struct {
name string
statusCode int
body string
expected bool
}{
{
name: "revoked refresh token",
statusCode: http.StatusBadRequest,
body: `{"error":"invalid_grant","error_description":"Refresh token revoked"}`,
expected: true,
},
{
name: "bad gateway HTML page",
statusCode: http.StatusBadGateway,
body: "<html><head><title>502 Server Error</title></head></html>",
expected: false,
},
{
name: "service unavailable",
statusCode: http.StatusServiceUnavailable,
body: "",
expected: false,
},
{
name: "rate limited",
statusCode: http.StatusTooManyRequests,
body: `{"error":"too_many_requests"}`,
expected: false,
},
{
// Our credentials are wrong, not the user's token.
name: "client misconfigured",
statusCode: http.StatusUnauthorized,
body: `{"error":"invalid_client"}`,
expected: false,
},
{
// Spotify only documents the 400 for a dead token, so an
// invalid_grant under any other status stays retryable.
name: "invalid_grant under an undocumented status",
statusCode: http.StatusUnauthorized,
body: `{"error":"invalid_grant"}`,
expected: false,
},
{
name: "bad request with unparseable body",
statusCode: http.StatusBadRequest,
body: "not json",
expected: false,
},
}

for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
if got := isRefreshTokenRejected(tc.statusCode, []byte(tc.body)); got != tc.expected {
t.Errorf("isRefreshTokenRejected() = %v, want %v", got, tc.expected)
}
})
}
}
Loading