diff --git a/cmd/handlers.go b/cmd/handlers.go index 4c29f99..9a5d506 100644 --- a/cmd/handlers.go +++ b/cmd/handlers.go @@ -11,7 +11,7 @@ import ( "github.com/teal-fm/piper/db/apikey" "github.com/teal-fm/piper/models" atprotoauth "github.com/teal-fm/piper/oauth/atproto" - pages "github.com/teal-fm/piper/pages" + "github.com/teal-fm/piper/pages" "github.com/teal-fm/piper/service/applemusic" atprotoservice "github.com/teal-fm/piper/service/atproto" "github.com/teal-fm/piper/service/musicbrainz" @@ -139,26 +139,26 @@ func handleLinkLastfmSubmit(database *db.DB) http.HandlerFunc { } func handleAppleMusicLink(pg *pages.Pages, am *applemusic.Service) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "text/html") - devToken, _, errTok := am.GenerateDeveloperToken() - if errTok != nil { - log.Printf("Error generating Apple Music developer token: %v", errTok) - http.Error(w, "Failed to prepare Apple Music", http.StatusInternalServerError) - return - } - data := struct{ - NavBar pages.NavBar - DevToken string - }{DevToken: devToken} - err := pg.Execute("applemusic_link", w, data) - if err != nil { - log.Printf("Error executing template: %v", err) - } - } + return func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/html") + devToken, _, errTok := am.GenerateDeveloperToken() + if errTok != nil { + log.Printf("Error generating Apple Music developer token: %v", errTok) + http.Error(w, "Failed to prepare Apple Music", http.StatusInternalServerError) + return + } + data := struct { + NavBar pages.NavBar + DevToken string + }{DevToken: devToken} + err := pg.Execute("applemusic_link", w, data) + if err != nil { + log.Printf("Error executing template: %v", err) + } + } } -func apiCurrentTrack(spotifyService *spotify.SpotifyService) http.HandlerFunc { +func apiCurrentTrack(spotifyService *spotify.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { userID, ok := session.GetUserID(r.Context()) if !ok { @@ -181,7 +181,7 @@ func apiCurrentTrack(spotifyService *spotify.SpotifyService) http.HandlerFunc { } } -func apiTrackHistory(spotifyService *spotify.SpotifyService) http.HandlerFunc { +func apiTrackHistory(spotifyService *spotify.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { userID, ok := session.GetUserID(r.Context()) if !ok { @@ -210,7 +210,7 @@ func apiTrackHistory(spotifyService *spotify.SpotifyService) http.HandlerFunc { } } -func apiMusicBrainzSearch(mbService *musicbrainz.MusicBrainzService) http.HandlerFunc { +func apiMusicBrainzSearch(mbService *musicbrainz.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { if mbService == nil { jsonResponse(w, http.StatusServiceUnavailable, map[string]string{"error": "MusicBrainz service is not available"}) @@ -350,64 +350,64 @@ func apiUnlinkLastfmHandler(database *db.DB) http.HandlerFunc { // apiAppleMusicAuthorize stores a MusicKit user token for the current user func apiAppleMusicAuthorize(database *db.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, authenticated := session.GetUserID(r.Context()) - if !authenticated { - jsonResponse(w, http.StatusUnauthorized, map[string]string{"error": "Unauthorized"}) - return - } - if r.Method != http.MethodPost { - jsonResponse(w, http.StatusMethodNotAllowed, map[string]string{"error": "Method not allowed"}) - return - } - - var req struct { - UserToken string `json:"userToken"` - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonResponse(w, http.StatusBadRequest, map[string]string{"error": "Invalid request body"}) - return - } - if req.UserToken == "" { - jsonResponse(w, http.StatusBadRequest, map[string]string{"error": "userToken is required"}) - return - } - - if err := database.UpdateAppleMusicUserToken(userID, req.UserToken); err != nil { - log.Printf("apiAppleMusicAuthorize: failed to save token for user %d: %v", userID, err) - jsonResponse(w, http.StatusInternalServerError, map[string]string{"error": "Failed to save token"}) - return - } - - jsonResponse(w, http.StatusOK, map[string]any{"status": "ok"}) - } + return func(w http.ResponseWriter, r *http.Request) { + userID, authenticated := session.GetUserID(r.Context()) + if !authenticated { + jsonResponse(w, http.StatusUnauthorized, map[string]string{"error": "Unauthorized"}) + return + } + if r.Method != http.MethodPost { + jsonResponse(w, http.StatusMethodNotAllowed, map[string]string{"error": "Method not allowed"}) + return + } + + var req struct { + UserToken string `json:"userToken"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + jsonResponse(w, http.StatusBadRequest, map[string]string{"error": "Invalid request body"}) + return + } + if req.UserToken == "" { + jsonResponse(w, http.StatusBadRequest, map[string]string{"error": "userToken is required"}) + return + } + + if err := database.UpdateAppleMusicUserToken(userID, req.UserToken); err != nil { + log.Printf("apiAppleMusicAuthorize: failed to save token for user %d: %v", userID, err) + jsonResponse(w, http.StatusInternalServerError, map[string]string{"error": "Failed to save token"}) + return + } + + jsonResponse(w, http.StatusOK, map[string]any{"status": "ok"}) + } } // apiAppleMusicUnlink clears the MusicKit user token for the current user func apiAppleMusicUnlink(database *db.DB) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - userID, authenticated := session.GetUserID(r.Context()) - if !authenticated { - jsonResponse(w, http.StatusUnauthorized, map[string]string{"error": "Unauthorized"}) - return - } - if r.Method != http.MethodPost { - jsonResponse(w, http.StatusMethodNotAllowed, map[string]string{"error": "Method not allowed"}) - return - } - - if err := database.ClearAppleMusicUserToken(userID); err != nil { - log.Printf("apiAppleMusicUnlink: failed to clear token for user %d: %v", userID, err) - jsonResponse(w, http.StatusInternalServerError, map[string]string{"error": "Failed to unlink Apple Music"}) - return - } - - jsonResponse(w, http.StatusOK, map[string]any{"status": "ok"}) - } + return func(w http.ResponseWriter, r *http.Request) { + userID, authenticated := session.GetUserID(r.Context()) + if !authenticated { + jsonResponse(w, http.StatusUnauthorized, map[string]string{"error": "Unauthorized"}) + return + } + if r.Method != http.MethodPost { + jsonResponse(w, http.StatusMethodNotAllowed, map[string]string{"error": "Method not allowed"}) + return + } + + if err := database.ClearAppleMusicUserToken(userID); err != nil { + log.Printf("apiAppleMusicUnlink: failed to clear token for user %d: %v", userID, err) + jsonResponse(w, http.StatusInternalServerError, map[string]string{"error": "Failed to unlink Apple Music"}) + return + } + + jsonResponse(w, http.StatusOK, map[string]any{"status": "ok"}) + } } // apiSubmitListensHandler handles ListenBrainz-compatible submissions -func apiSubmitListensHandler(database *db.DB, atprotoService *atprotoauth.ATprotoAuthService, playingNowService *playingnow.PlayingNowService, mbService *musicbrainz.MusicBrainzService) http.HandlerFunc { +func apiSubmitListensHandler(database *db.DB, atprotoService *atprotoauth.AuthService, playingNowService *playingnow.Service, mbService *musicbrainz.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { userID, authenticated := session.GetUserID(r.Context()) if !authenticated { @@ -471,7 +471,7 @@ func apiSubmitListensHandler(database *db.DB, atprotoService *atprotoauth.ATprot } // Convert to internal Track format - track := listen.ConvertToTrack(userID) + track := listen.ConvertToTrack() // Attempt to hydrate with MusicBrainz data if service is available and track doesn't have MBIDs if mbService != nil && track.RecordingMBID == nil { @@ -538,7 +538,7 @@ func apiSubmitListensHandler(database *db.DB, atprotoService *atprotoauth.ATprot } // apiMbTokenValidateHandler handles ListenBrainz token validation requests -func apiMbTokenValidateHandler(sm *session.SessionManager) http.HandlerFunc { +func apiMbTokenValidateHandler(sm *session.Manager) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { apiKeyStr, apiKeyErr := apikey.ExtractApiKey(r) diff --git a/cmd/listenbrainz_test.go b/cmd/listenbrainz_test.go index bca37ea..384c75b 100644 --- a/cmd/listenbrainz_test.go +++ b/cmd/listenbrainz_test.go @@ -490,7 +490,7 @@ func TestListenBrainzDataConversion(t *testing.T) { }, } - track := payload.ConvertToTrack(123) + track := payload.ConvertToTrack() // Verify conversion if track.Name != "Test Track" { diff --git a/cmd/main.go b/cmd/main.go index 76de3b0..f3d0fde 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -16,7 +16,7 @@ import ( "github.com/teal-fm/piper/db" "github.com/teal-fm/piper/oauth" "github.com/teal-fm/piper/oauth/atproto" - pages "github.com/teal-fm/piper/pages" + "github.com/teal-fm/piper/pages" apikeyService "github.com/teal-fm/piper/service/apikey" "github.com/teal-fm/piper/service/musicbrainz" "github.com/teal-fm/piper/service/spotify" @@ -25,13 +25,13 @@ import ( type application struct { database *db.DB - sessionManager *session.SessionManager - oauthManager *oauth.OAuthServiceManager - spotifyService *spotify.SpotifyService + sessionManager *session.Manager + oauthManager *oauth.ServiceManager + spotifyService *spotify.Service apiKeyService *apikeyService.Service - mbService *musicbrainz.MusicBrainzService - atprotoService *atproto.ATprotoAuthService - playingNowService *playingnow.PlayingNowService + mbService *musicbrainz.Service + atprotoService *atproto.AuthService + playingNowService *playingnow.Service appleMusicService *applemusic.Service pages *pages.Pages } @@ -89,38 +89,38 @@ func main() { playingNowService := playingnow.NewPlayingNowService(database, atprotoService) spotifyService := spotify.NewSpotifyService(database, atprotoService, mbService, playingNowService) lastfmService := lastfm.NewLastFMService(database, viper.GetString("lastfm.api_key"), mbService, atprotoService, playingNowService) - // Read Apple Music settings with env fallbacks - teamID := viper.GetString("applemusic.team_id") - if teamID == "" { - teamID = viper.GetString("APPLE_MUSIC_TEAM_ID") - } - keyID := viper.GetString("applemusic.key_id") - if keyID == "" { - keyID = viper.GetString("APPLE_MUSIC_KEY_ID") - } - keyPath := viper.GetString("applemusic.private_key_path") - if keyPath == "" { - keyPath = viper.GetString("APPLE_MUSIC_PRIVATE_KEY_PATH") - } - - var appleMusicService *applemusic.Service - // Only initialize Apple Music service if all required credentials are present - if teamID != "" && keyID != "" && keyPath != "" { - appleMusicService = applemusic.NewService( - teamID, - keyID, - keyPath, - ).WithPersistence( - func() (string, time.Time, bool, error) { - return database.GetAppleMusicDeveloperToken() - }, - func(token string, exp time.Time) error { - return database.SaveAppleMusicDeveloperToken(token, exp) - }, - ).WithDeps(database, atprotoService, mbService, playingNowService) - } else { - log.Println("Apple Music credentials not configured (missing team_id, key_id, or private_key_path). Apple Music features will be disabled.") - } + // Read Apple Music settings with env fallbacks + teamID := viper.GetString("applemusic.team_id") + if teamID == "" { + teamID = viper.GetString("APPLE_MUSIC_TEAM_ID") + } + keyID := viper.GetString("applemusic.key_id") + if keyID == "" { + keyID = viper.GetString("APPLE_MUSIC_KEY_ID") + } + keyPath := viper.GetString("applemusic.private_key_path") + if keyPath == "" { + keyPath = viper.GetString("APPLE_MUSIC_PRIVATE_KEY_PATH") + } + + var appleMusicService *applemusic.Service + // Only initialize Apple Music service if all required credentials are present + if teamID != "" && keyID != "" && keyPath != "" { + appleMusicService = applemusic.NewService( + teamID, + keyID, + keyPath, + ).WithPersistence( + func() (string, time.Time, bool, error) { + return database.GetAppleMusicDeveloperToken() + }, + func(token string, exp time.Time) error { + return database.SaveAppleMusicDeveloperToken(token, exp) + }, + ).WithDeps(database, atprotoService, mbService, playingNowService) + } else { + log.Println("Apple Music credentials not configured (missing team_id, key_id, or private_key_path). Apple Music features will be disabled.") + } oauthManager := oauth.NewOAuthServiceManager() @@ -158,12 +158,12 @@ func main() { go spotifyService.StartListeningTracker(trackerInterval) - go lastfmService.StartListeningTracker(lastfmInterval) - // Apple Music tracker uses same tracker.interval as Spotify for now - // Only start if Apple Music service is configured - if appleMusicService != nil { - go appleMusicService.StartListeningTracker(trackerInterval) - } + go lastfmService.StartListeningTracker(lastfmInterval) + // Apple Music tracker uses same tracker.interval as Spotify for now + // Only start if Apple Music service is configured + if appleMusicService != nil { + go appleMusicService.StartListeningTracker(trackerInterval) + } serverAddr := fmt.Sprintf("%s:%s", viper.GetString("server.host"), viper.GetString("server.port")) server := &http.Server{ diff --git a/cmd/routes.go b/cmd/routes.go index 9609aaf..eda0246 100644 --- a/cmd/routes.go +++ b/cmd/routes.go @@ -28,7 +28,7 @@ func (app *application) routes() http.Handler { mux.HandleFunc("/api-keys", session.WithAuth(app.apiKeyService.HandleAPIKeyManagement(app.database, app.pages), app.sessionManager)) mux.HandleFunc("/link-lastfm", session.WithAuth(handleLinkLastfmForm(app.database, app.pages), app.sessionManager)) // GET form mux.HandleFunc("/link-lastfm/submit", session.WithAuth(handleLinkLastfmSubmit(app.database), app.sessionManager)) // POST submit - Changed route slightly - mux.HandleFunc("/link-applemusic", session.WithAuth(handleAppleMusicLink(app.pages, app.appleMusicService), app.sessionManager)) + mux.HandleFunc("/link-applemusic", session.WithAuth(handleAppleMusicLink(app.pages, app.appleMusicService), app.sessionManager)) mux.HandleFunc("/logout", app.oauthManager.HandleLogout("atproto")) mux.HandleFunc("/debug/", session.WithAuth(app.sessionManager.HandleDebug, app.sessionManager)) @@ -40,19 +40,17 @@ func (app *application) routes() http.Handler { mux.HandleFunc("/api/v1/history", session.WithAPIAuth(apiTrackHistory(app.spotifyService), app.sessionManager)) // Spotify History mux.HandleFunc("/api/v1/musicbrainz/search", apiMusicBrainzSearch(app.mbService)) // MusicBrainz (public?) - // Apple Music user authorization (protected with session auth) - mux.HandleFunc("/api/v1/applemusic/authorize", session.WithAuth(apiAppleMusicAuthorize(app.database), app.sessionManager)) - mux.HandleFunc("/api/v1/applemusic/unlink", session.WithAuth(apiAppleMusicUnlink(app.database), app.sessionManager)) + // Apple Music user authorization (protected with session auth) + mux.HandleFunc("/api/v1/applemusic/authorize", session.WithAuth(apiAppleMusicAuthorize(app.database), app.sessionManager)) + mux.HandleFunc("/api/v1/applemusic/unlink", session.WithAuth(apiAppleMusicUnlink(app.database), app.sessionManager)) // ListenBrainz-compatible endpoint mux.HandleFunc("/1/submit-listens", session.WithAPIAuth(apiSubmitListensHandler(app.database, app.atprotoService, app.playingNowService, app.mbService), app.sessionManager)) mux.HandleFunc("/1/validate-token", apiMbTokenValidateHandler(app.sessionManager)) serverUrlRoot := viper.GetString("server.root_url") - atpClientId := viper.GetString("atproto.client_id") - atpCallbackUrl := viper.GetString("atproto.callback_url") mux.HandleFunc("/oauth-client-metadata.json", func(w http.ResponseWriter, r *http.Request) { - app.atprotoService.HandleClientMetadata(w, r, serverUrlRoot, atpClientId, atpCallbackUrl) + app.atprotoService.HandleClientMetadata(w, r, serverUrlRoot) }) mux.HandleFunc("/oauth/jwks.json", app.atprotoService.HandleJwks) diff --git a/config/config.go b/config/config.go index 054440c..ce0b239 100644 --- a/config/config.go +++ b/config/config.go @@ -1,6 +1,7 @@ package config import ( + "errors" "log" "strings" @@ -48,7 +49,8 @@ func Load() { viper.AddConfigPath(".") if err := viper.ReadInConfig(); err != nil { - if _, ok := err.(viper.ConfigFileNotFoundError); !ok { + var configFileNotFoundError viper.ConfigFileNotFoundError + if !errors.As(err, &configFileNotFoundError) { log.Fatalf("Error reading config file: %v", err) } log.Println("Config file not found, using default values and environment variables") @@ -58,7 +60,7 @@ func Load() { // check for required settings requiredVars := []string{"spotify.client_id", "spotify.client_secret"} - missingVars := []string{} + var missingVars []string for _, v := range requiredVars { if !viper.IsSet(v) { diff --git a/db/apikey/apikey.go b/db/apikey/apikey.go index 7388aa5..a02e2aa 100644 --- a/db/apikey/apikey.go +++ b/db/apikey/apikey.go @@ -22,15 +22,15 @@ type ApiKey struct { ExpiresAt time.Time } -// ApiKeyManager manages API keys -type ApiKeyManager struct { +// Manager ApiKeyManager manages API keys +type Manager struct { db *db.DB apiKeys map[string]*ApiKey mu sync.RWMutex } // NewApiKeyManager creates a new API key manager -func NewApiKeyManager(database *db.DB) *ApiKeyManager { +func NewApiKeyManager(database *db.DB) *Manager { // Initialize API keys table if it doesn't exist _, err := database.Exec(` CREATE TABLE IF NOT EXISTS api_keys ( @@ -46,14 +46,14 @@ func NewApiKeyManager(database *db.DB) *ApiKeyManager { log.Printf("Error creating api_keys table: %v", err) } - return &ApiKeyManager{ + return &Manager{ db: database, apiKeys: make(map[string]*ApiKey), } } // CreateApiKey creates a new API key for a user -func (am *ApiKeyManager) CreateApiKey(userID int64, name string, validityDays int) (*ApiKey, error) { +func (am *Manager) CreateApiKey(userID int64, name string, validityDays int) (*ApiKey, error) { am.mu.Lock() defer am.mu.Unlock() @@ -92,7 +92,7 @@ func (am *ApiKeyManager) CreateApiKey(userID int64, name string, validityDays in } // GetApiKey retrieves an API key by ID -func (am *ApiKeyManager) GetApiKey(apiKeyID string) (*ApiKey, bool) { +func (am *Manager) GetApiKey(apiKeyID string) (*ApiKey, bool) { // First check in-memory cache am.mu.RLock() apiKey, exists := am.apiKeys[apiKeyID] @@ -132,7 +132,7 @@ func (am *ApiKeyManager) GetApiKey(apiKeyID string) (*ApiKey, bool) { } // DeleteApiKey removes an API key -func (am *ApiKeyManager) DeleteApiKey(apiKeyID string) error { +func (am *Manager) DeleteApiKey(apiKeyID string) error { am.mu.Lock() delete(am.apiKeys, apiKeyID) am.mu.Unlock() @@ -142,7 +142,7 @@ func (am *ApiKeyManager) DeleteApiKey(apiKeyID string) error { } // GetUserApiKeys retrieves all API keys for a user -func (am *ApiKeyManager) GetUserApiKeys(userID int64) ([]*ApiKey, error) { +func (am *Manager) GetUserApiKeys(userID int64) ([]*ApiKey, error) { rows, err := am.db.Query(` SELECT id, user_id, name, created_at, expires_at FROM api_keys diff --git a/db/atproto.go b/db/atproto.go index 5516994..07ed843 100644 --- a/db/atproto.go +++ b/db/atproto.go @@ -3,6 +3,7 @@ package db import ( "context" "database/sql" + "errors" "fmt" "strings" "time" @@ -20,7 +21,7 @@ func (db *DB) FindOrCreateUserByDID(did string) (*models.User, error) { WHERE atproto_did = ?`, did).Scan(&user.ID, &user.ATProtoDID, &user.CreatedAt, &user.UpdatedAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { now := time.Now().UTC() // create user! result, insertErr := db.Exec(` @@ -162,7 +163,7 @@ func (s *SqliteATProtoStore) GetSession(ctx context.Context, did syntax.DID, ses &dpopPrivateKeyMultibase, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, fmt.Errorf("session not found: %s", lookUpKey) } if err != nil { @@ -271,7 +272,7 @@ func (s *SqliteATProtoStore) GetAuthRequestInfo(ctx context.Context, state strin &dpopAuthServerNonce, &dpopPrivateKeyMultibase, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, fmt.Errorf("request info not found: %s", state) } if err != nil { @@ -307,7 +308,7 @@ func (s *SqliteATProtoStore) SaveAuthRequestInfo(ctx context.Context, info oauth if err == nil { return fmt.Errorf("auth request already saved for state %s", info.State) } - if err != nil && err != sql.ErrNoRows { + if !errors.Is(err, sql.ErrNoRows) { return err } var accountDIDStr interface{} diff --git a/db/db.go b/db/db.go index 8bfd9d3..5635804 100644 --- a/db/db.go +++ b/db/db.go @@ -3,6 +3,7 @@ package db import ( "database/sql" "encoding/json" + "errors" "fmt" "log" "os" @@ -20,9 +21,13 @@ type DB struct { func New(dbPath string) (*DB, error) { dir := filepath.Dir(dbPath) - if dir != "." && dir != "/" { - os.MkdirAll(dir, 0755) - } + if dir != "." && dir != "/" { + err := os.MkdirAll(dir, 0755) + if err != nil { + fmt.Println("Failed to create directory for database:", err) + return nil, err + } + } db, err := sql.Open("sqlite3", dbPath) if err != nil { @@ -150,44 +155,44 @@ func (db *DB) Initialize() error { // Apple Music developer token persistence func (db *DB) ensureAppleMusicTokenTable() error { - _, err := db.Exec(` + _, err := db.Exec(` CREATE TABLE IF NOT EXISTS applemusic_token ( token TEXT, expires_at TIMESTAMP )`) - return err + return err } func (db *DB) GetAppleMusicDeveloperToken() (string, time.Time, bool, error) { - if err := db.ensureAppleMusicTokenTable(); err != nil { - return "", time.Time{}, false, err - } - var token string - var exp time.Time - err := db.QueryRow(`SELECT token, expires_at FROM applemusic_token LIMIT 1`).Scan(&token, &exp) - if err == sql.ErrNoRows { - return "", time.Time{}, false, nil - } - if err != nil { - return "", time.Time{}, false, err - } - return token, exp, true, nil + if err := db.ensureAppleMusicTokenTable(); err != nil { + return "", time.Time{}, false, err + } + var token string + var exp time.Time + err := db.QueryRow(`SELECT token, expires_at FROM applemusic_token LIMIT 1`).Scan(&token, &exp) + if errors.Is(err, sql.ErrNoRows) { + return "", time.Time{}, false, nil + } + if err != nil { + return "", time.Time{}, false, err + } + return token, exp, true, nil } func (db *DB) SaveAppleMusicDeveloperToken(token string, exp time.Time) error { - if err := db.ensureAppleMusicTokenTable(); err != nil { - return err - } - // Replace existing single row - _, err := db.Exec(`DELETE FROM applemusic_token`) - if err != nil { - return err - } - _, err = db.Exec(`INSERT INTO applemusic_token (token, expires_at) VALUES (?, ?)`, token, exp) - return err + if err := db.ensureAppleMusicTokenTable(); err != nil { + return err + } + // Replace existing single row + _, err := db.Exec(`DELETE FROM applemusic_token`) + if err != nil { + return err + } + _, err = db.Exec(`INSERT INTO applemusic_token (token, expires_at) VALUES (?, ?)`, token, exp) + return err } -// create user without spotify id +// CreateUser create user without spotify id func (db *DB) CreateUser(user *models.User) (int64, error) { now := time.Now().UTC() @@ -203,7 +208,7 @@ func (db *DB) CreateUser(user *models.User) (int64, error) { return result.LastInsertId() } -// add spotify session to user, returning the updated user +// AddSpotifySession add spotify session to user, returning the updated user func (db *DB) AddSpotifySession(userID int64, username, email, spotifyId, accessToken, refreshToken string, tokenExpiry time.Time) (*models.User, error) { now := time.Now().UTC() @@ -227,7 +232,7 @@ func (db *DB) AddSpotifySession(userID int64, username, email, spotifyId, access func (db *DB) GetUserByID(ID int64) (*models.User, error) { user := &models.User{} - err := db.QueryRow(` + err := db.QueryRow(` SELECT id, username, email, @@ -242,12 +247,12 @@ func (db *DB) GetUserByID(ID int64) (*models.User, error) { created_at, updated_at FROM users WHERE id = ?`, ID).Scan( - &user.ID, &user.Username, &user.Email, &user.ATProtoDID, &user.MostRecentAtProtoSessionID, &user.SpotifyID, - &user.AccessToken, &user.RefreshToken, &user.TokenExpiry, - &user.LastFMUsername, &user.AppleMusicUserToken, - &user.CreatedAt, &user.UpdatedAt) + &user.ID, &user.Username, &user.Email, &user.ATProtoDID, &user.MostRecentAtProtoSessionID, &user.SpotifyID, + &user.AccessToken, &user.RefreshToken, &user.TokenExpiry, + &user.LastFMUsername, &user.AppleMusicUserToken, + &user.CreatedAt, &user.UpdatedAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, nil } @@ -261,15 +266,15 @@ func (db *DB) GetUserByID(ID int64) (*models.User, error) { func (db *DB) GetUserBySpotifyID(spotifyID string) (*models.User, error) { user := &models.User{} - err := db.QueryRow(` + err := db.QueryRow(` SELECT id, username, email, spotify_id, access_token, refresh_token, token_expiry, lastfm_username, applemusic_user_token, created_at, updated_at FROM users WHERE spotify_id = ?`, spotifyID).Scan( - &user.ID, &user.Username, &user.Email, &user.SpotifyID, - &user.AccessToken, &user.RefreshToken, &user.TokenExpiry, - &user.LastFMUsername, &user.AppleMusicUserToken, - &user.CreatedAt, &user.UpdatedAt) + &user.ID, &user.Username, &user.Email, &user.SpotifyID, + &user.AccessToken, &user.RefreshToken, &user.TokenExpiry, + &user.LastFMUsername, &user.AppleMusicUserToken, + &user.CreatedAt, &user.UpdatedAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, nil } @@ -293,56 +298,56 @@ func (db *DB) UpdateUserToken(userID int64, accessToken, refreshToken string, ex } func (db *DB) UpdateAppleMusicUserToken(userID int64, userToken string) error { - now := time.Now().UTC() - _, err := db.Exec(` + now := time.Now().UTC() + _, err := db.Exec(` UPDATE users SET applemusic_user_token = ?, updated_at = ? WHERE id = ?`, - userToken, now, userID) - return err + userToken, now, userID) + return err } // ClearAppleMusicUserToken removes the stored Apple Music user token for a user func (db *DB) ClearAppleMusicUserToken(userID int64) error { - now := time.Now().UTC() - _, err := db.Exec(` + now := time.Now().UTC() + _, err := db.Exec(` UPDATE users SET applemusic_user_token = NULL, updated_at = ? WHERE id = ?`, - now, userID) - return err + now, userID) + return err } // GetAllAppleMusicLinkedUsers returns users who have an Apple Music user token set func (db *DB) GetAllAppleMusicLinkedUsers() ([]*models.User, error) { - rows, err := db.Query(` + rows, err := db.Query(` SELECT id, username, email, atproto_did, most_recent_at_session_id, spotify_id, access_token, refresh_token, token_expiry, lastfm_username, applemusic_user_token, created_at, updated_at FROM users WHERE applemusic_user_token IS NOT NULL AND applemusic_user_token != '' ORDER BY id`) - if err != nil { - return nil, err - } - defer rows.Close() - - var users []*models.User - for rows.Next() { - u := &models.User{} - if err := rows.Scan( - &u.ID, &u.Username, &u.Email, &u.ATProtoDID, &u.MostRecentAtProtoSessionID, - &u.SpotifyID, &u.AccessToken, &u.RefreshToken, &u.TokenExpiry, - &u.LastFMUsername, &u.AppleMusicUserToken, &u.CreatedAt, &u.UpdatedAt, - ); err != nil { - return nil, err - } - users = append(users, u) - } - if err := rows.Err(); err != nil { - return nil, err - } - return users, nil + if err != nil { + return nil, err + } + defer rows.Close() + + var users []*models.User + for rows.Next() { + u := &models.User{} + if err := rows.Scan( + &u.ID, &u.Username, &u.Email, &u.ATProtoDID, &u.MostRecentAtProtoSessionID, + &u.SpotifyID, &u.AccessToken, &u.RefreshToken, &u.TokenExpiry, + &u.LastFMUsername, &u.AppleMusicUserToken, &u.CreatedAt, &u.UpdatedAt, + ); err != nil { + return nil, err + } + users = append(users, u) + } + if err := rows.Err(); err != nil { + return nil, err + } + return users, nil } func (db *DB) SaveTrack(userID int64, track *models.Track) (int64, error) { @@ -352,9 +357,9 @@ func (db *DB) SaveTrack(userID int64, track *models.Track) (int64, error) { bytes, err := json.Marshal(track.Artist) if err != nil { return 0, err - } else { - artistString = string(bytes) } + + artistString = string(bytes) } var trackID int64 @@ -376,9 +381,9 @@ func (db *DB) UpdateTrack(trackID int64, track *models.Track) error { bytes, err := json.Marshal(track.Artist) if err != nil { return err - } else { - artistString = string(bytes) } + + artistString = string(bytes) } _, err := db.Exec(` @@ -521,7 +526,7 @@ func (db *DB) GetAllActiveUsersWithUnExpiredTokens() ([]*models.User, error) { return SpotifyQueryMapping(rows) } -// debug to view current user's information +// DebugViewUserInformation debug to view current user's information // put everything in an 'any' type func (db *DB) DebugViewUserInformation(userID int64) (map[string]any, error) { // Use Query instead of QueryRow to get access to column names and ensure only one row is processed @@ -595,7 +600,7 @@ func (db *DB) GetLastKnownTimestamp(userID int64) (*time.Time, error) { LIMIT 1`, userID).Scan(&lastTimestamp) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, nil } return nil, fmt.Errorf("failed to query last scrobble timestamp for user %d: %w", userID, err) diff --git a/models/listenbrainz.go b/models/listenbrainz.go index a39b140..887dc5e 100644 --- a/models/listenbrainz.go +++ b/models/listenbrainz.go @@ -47,7 +47,7 @@ type ListenBrainzAdditionalInfo struct { } // ConvertToTrack converts ListenBrainz format to internal Track format -func (lbp *ListenBrainzPayload) ConvertToTrack(userID int64) Track { +func (lbp *ListenBrainzPayload) ConvertToTrack() Track { track := Track{ Name: lbp.TrackMetadata.TrackName, Artist: []Artist{{Name: lbp.TrackMetadata.ArtistName}}, diff --git a/models/user.go b/models/user.go index 580f68d..f648e65 100644 --- a/models/user.go +++ b/models/user.go @@ -2,7 +2,7 @@ package models import "time" -// an end user of piper +// User an end user of piper type User struct { ID int64 Username *string diff --git a/oauth/atproto/atproto.go b/oauth/atproto/atproto.go index 95808cc..811dae1 100644 --- a/oauth/atproto/atproto.go +++ b/oauth/atproto/atproto.go @@ -20,16 +20,16 @@ import ( "slices" ) -type ATprotoAuthService struct { +type AuthService struct { clientApp *oauth.ClientApp DB *db.DB - sessionManager *session.SessionManager + sessionManager *session.Manager clientId string callbackUrl string logger *log.Logger } -func NewATprotoAuthService(database *db.DB, sessionManager *session.SessionManager, clientSecretKey string, clientId string, callbackUrl string, clientSecretId string) (*ATprotoAuthService, error) { +func NewATprotoAuthService(database *db.DB, sessionManager *session.Manager, clientSecretKey string, clientId string, callbackUrl string, clientSecretId string) (*AuthService, error) { fmt.Println(clientId, callbackUrl) scopes := []string{"atproto", "repo:fm.teal.alpha.feed.play", "repo:fm.teal.alpha.actor.status"} @@ -49,7 +49,7 @@ func NewATprotoAuthService(database *db.DB, sessionManager *session.SessionManag logger := log.New(os.Stdout, "ATProto oauth: ", log.LstdFlags|log.Lmsgprefix) - svc := &ATprotoAuthService{ + svc := &AuthService{ clientApp: oauthClient, callbackUrl: callbackUrl, DB: database, @@ -60,7 +60,7 @@ func NewATprotoAuthService(database *db.DB, sessionManager *session.SessionManag return svc, nil } -func (a *ATprotoAuthService) GetATProtoClient(accountDID string, sessionID string, ctx context.Context) (*client.APIClient, error) { +func (a *AuthService) GetATProtoClient(accountDID string, sessionID string, ctx context.Context) (*client.APIClient, error) { did, err := syntax.ParseDID(accountDID) if err != nil { return nil, err @@ -75,7 +75,7 @@ func (a *ATprotoAuthService) GetATProtoClient(accountDID string, sessionID strin } -func (a *ATprotoAuthService) HandleLogin(w http.ResponseWriter, r *http.Request) { +func (a *AuthService) HandleLogin(w http.ResponseWriter, r *http.Request) { handle := r.URL.Query().Get("handle") if handle == "" { a.logger.Printf("ATProto Login Error: handle is required") @@ -96,21 +96,26 @@ func (a *ATprotoAuthService) HandleLogin(w http.ResponseWriter, r *http.Request) http.Redirect(w, r, authUrl.String(), http.StatusFound) } -func (a *ATprotoAuthService) HandleLogout(w http.ResponseWriter, r *http.Request) { +func (a *AuthService) HandleLogout(w http.ResponseWriter, r *http.Request) { cookie, err := r.Cookie("session") if err == nil { - session, exists := a.sessionManager.GetSession(cookie.Value) + webSession, exists := a.sessionManager.GetSession(cookie.Value) if !exists { http.Redirect(w, r, "/", http.StatusSeeOther) return } - dbUser, err := a.DB.GetUserByID(session.UserID) + dbUser, err := a.DB.GetUserByID(webSession.UserID) if err != nil { http.Redirect(w, r, "/", http.StatusSeeOther) return } + if dbUser == nil { + http.Redirect(w, r, "/", http.StatusSeeOther) + return + } + did, err := syntax.ParseDID(*dbUser.ATProtoDID) if err != nil { @@ -120,7 +125,7 @@ func (a *ATprotoAuthService) HandleLogout(w http.ResponseWriter, r *http.Request } ctx := r.Context() - err = a.clientApp.Logout(ctx, did, session.ATProtoSessionID) + err = a.clientApp.Logout(ctx, did, webSession.ATProtoSessionID) if err != nil { a.logger.Printf("Error logging the user: %s out: %s", did, err) } @@ -132,7 +137,7 @@ func (a *ATprotoAuthService) HandleLogout(w http.ResponseWriter, r *http.Request http.Redirect(w, r, "/", http.StatusSeeOther) } -func (a *ATprotoAuthService) HandleCallback(w http.ResponseWriter, r *http.Request) (int64, error) { +func (a *AuthService) HandleCallback(w http.ResponseWriter, r *http.Request) (int64, error) { ctx := r.Context() sessData, err := a.clientApp.ProcessCallback(ctx, r.URL.Query()) @@ -165,6 +170,6 @@ func (a *ATprotoAuthService) HandleCallback(w http.ResponseWriter, r *http.Reque a.logger.Printf("Failed to set latest atproto session id for user %d: %v", user.ID, err) } - a.logger.Printf("ATProto Callback Success: User %d (DID: %s) authenticated.", user.ID, user.ATProtoDID) + a.logger.Printf("ATProto Callback Success: User %d (DID: %v) authenticated.", user.ID, user.ATProtoDID) return user.ID, nil } diff --git a/oauth/atproto/http.go b/oauth/atproto/http.go index d5f5be2..c194b05 100644 --- a/oauth/atproto/http.go +++ b/oauth/atproto/http.go @@ -1,4 +1,4 @@ -// oauth/atproto/http.go +// Package atproto oauth/atproto/http.go package atproto import ( @@ -11,7 +11,7 @@ func strPtr(raw string) *string { return &raw } -func (a *ATprotoAuthService) HandleJwks(w http.ResponseWriter, r *http.Request) { +func (a *AuthService) HandleJwks(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") body := a.clientApp.Config.PublicJWKS() if err := json.NewEncoder(w).Encode(body); err != nil { @@ -20,7 +20,7 @@ func (a *ATprotoAuthService) HandleJwks(w http.ResponseWriter, r *http.Request) } } -func (a *ATprotoAuthService) HandleClientMetadata(w http.ResponseWriter, r *http.Request, serverUrlRoot, serverMetadataUrl, serverCallbackUrl string) { +func (a *AuthService) HandleClientMetadata(w http.ResponseWriter, r *http.Request, serverUrlRoot string) { meta := a.clientApp.Config.ClientMetadata() if a.clientApp.Config.IsConfidential() { diff --git a/oauth/atproto/resolve.go b/oauth/atproto/resolve.go deleted file mode 100644 index c6940da..0000000 --- a/oauth/atproto/resolve.go +++ /dev/null @@ -1,151 +0,0 @@ -package atproto - -// Stolen from https://github.com/haileyok/atproto-oauth-golang/blob/f780d3716e2b8a06c87271a2930894319526550e/cmd/web_server_demo/resolution.go - -import ( - "context" - "encoding/json" - "fmt" - "io" - "net" - "net/http" - "strings" - - "github.com/bluesky-social/indigo/atproto/syntax" -) - -// user information struct -type UserInformation struct { - AuthService string `json:"authService"` - AuthServer string `json:"authServer"` - // do NOT save the current handle permanently! - Handle string `json:"handle"` - DID string `json:"did"` -} - -type Identity struct { - AlsoKnownAs []string `json:"alsoKnownAs"` - Service []struct { - ID string `json:"id"` - Type string `json:"type"` - ServiceEndpoint string `json:"serviceEndpoint"` - } `json:"service"` -} - -func resolveHandle(ctx context.Context, handle string) (string, error) { - var did string - - _, err := syntax.ParseHandle(handle) - if err != nil { - return "", err - } - - recs, err := net.LookupTXT(fmt.Sprintf("_atproto.%s", handle)) - if err == nil { - for _, rec := range recs { - if strings.HasPrefix(rec, "did=") { - did = strings.Split(rec, "did=")[1] - break - } - } - } - - if did == "" { - req, err := http.NewRequestWithContext( - ctx, - "GET", - fmt.Sprintf("https://%s/.well-known/atproto-did", handle), - nil, - ) - if err != nil { - return "", err - } - - resp, err := http.DefaultClient.Do(req) - if err != nil { - return "", err - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - io.Copy(io.Discard, resp.Body) - return "", fmt.Errorf("unable to resolve handle") - } - - b, err := io.ReadAll(resp.Body) - if err != nil { - return "", err - } - - maybeDid := string(b) - - if _, err := syntax.ParseDID(maybeDid); err != nil { - return "", fmt.Errorf("unable to resolve handle") - } - - did = maybeDid - } - - return did, nil -} - -// Get the Identity document for a given DID -func getIdentityDocument(ctx context.Context, did string) (*Identity, error) { - var ustr string - if strings.HasPrefix(did, "did:plc:") { - ustr = fmt.Sprintf("https://plc.directory/%s", did) - } else if strings.HasPrefix(did, "did:web:") { - ustr = fmt.Sprintf("https://%s/.well-known/did.json", strings.TrimPrefix(did, "did:web:")) - } else { - return nil, fmt.Errorf("did was not a supported did type") - } - - req, err := http.NewRequestWithContext(ctx, "GET", ustr, nil) - if err != nil { - return nil, err - } - - resp, err := http.DefaultClient.Do(req) - if err != nil { - return nil, err - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - io.Copy(io.Discard, resp.Body) - return nil, fmt.Errorf("could not find identity in plc registry") - } - - var identity Identity - if err := json.NewDecoder(resp.Body).Decode(&identity); err != nil { - return nil, err - } - - return &identity, nil -} - -// Get the atproto PDS service endpoint from an Identity document -func getAtprotoPdsService(identity *Identity) (string, error) { - var service string - for _, svc := range identity.Service { - if svc.ID == "#atproto_pds" { - service = svc.ServiceEndpoint - break - } - } - - if service == "" { - return "", fmt.Errorf("could not find atproto_pds service in identity services") - } - - return service, nil -} - -func resolveServiceFromDoc(identity *Identity) (string, error) { - service, err := getAtprotoPdsService(identity) - if err != nil { - return "", err - } - - return service, nil -} diff --git a/oauth/oauth2.go b/oauth/oauth2.go index c13a4e1..e4adeea 100644 --- a/oauth/oauth2.go +++ b/oauth/oauth2.go @@ -16,7 +16,7 @@ import ( "golang.org/x/oauth2/spotify" ) -type OAuth2Service struct { +type Service struct { config oauth2.Config state string codeVerifier string @@ -30,7 +30,7 @@ func GenerateRandomState() string { return base64.URLEncoding.EncodeToString(b) } -func NewOAuth2Service(clientID, clientSecret, redirectURI string, scopes []string, provider string, tokenReceiver TokenReceiver) *OAuth2Service { +func NewOAuth2Service(clientID, clientSecret, redirectURI string, scopes []string, provider string, tokenReceiver TokenReceiver) *Service { var endpoint oauth2.Endpoint switch strings.ToLower(provider) { @@ -48,7 +48,7 @@ func NewOAuth2Service(clientID, clientSecret, redirectURI string, scopes []strin codeVerifier := GenerateCodeVerifier() codeChallenge := GenerateCodeChallenge(codeVerifier) - return &OAuth2Service{ + return &Service{ config: oauth2.Config{ ClientID: clientID, ClientSecret: clientSecret, @@ -63,21 +63,21 @@ func NewOAuth2Service(clientID, clientSecret, redirectURI string, scopes []strin } } -// generate a random code verifier, for PKCE +// GenerateCodeVerifier generate a random code verifier, for PKCE func GenerateCodeVerifier() string { b := make([]byte, 64) rand.Read(b) return base64.RawURLEncoding.EncodeToString(b) } -// generate a code challenge for verification later +// GenerateCodeChallenge generate a code challenge for verification later func GenerateCodeChallenge(verifier string) string { h := sha256.New() h.Write([]byte(verifier)) return base64.RawURLEncoding.EncodeToString(h.Sum(nil)) } -func (o *OAuth2Service) HandleLogin(w http.ResponseWriter, r *http.Request) { +func (o *Service) HandleLogin(w http.ResponseWriter, r *http.Request) { opts := []oauth2.AuthCodeOption{ oauth2.SetAuthURLParam("code_challenge", o.codeChallenge), oauth2.SetAuthURLParam("code_challenge_method", "S256"), @@ -86,12 +86,12 @@ func (o *OAuth2Service) HandleLogin(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, authURL, http.StatusSeeOther) } -func (o *OAuth2Service) HandleLogout(w http.ResponseWriter, r *http.Request) { +func (o *Service) HandleLogout(w http.ResponseWriter, r *http.Request) { //TODO not implemented yet. not sure what the api call is for this package http.Redirect(w, r, "/", http.StatusSeeOther) } -func (o *OAuth2Service) HandleCallback(w http.ResponseWriter, r *http.Request) (int64, error) { +func (o *Service) HandleCallback(w http.ResponseWriter, r *http.Request) (int64, error) { state := r.URL.Query().Get("state") if state != o.state { log.Printf("OAuth2 Callback Error: State mismatch. Expected '%s', got '%s'", o.state, state) @@ -131,30 +131,30 @@ func (o *OAuth2Service) HandleCallback(w http.ResponseWriter, r *http.Request) ( // store token and get uid userID, err := o.tokenReceiver.SetAccessToken(token.AccessToken, token.RefreshToken, userId, hasSession) if err != nil { - log.Printf("OAuth2 Callback Info: TokenReceiver did not return a valid user ID for token: %s...", token.AccessToken[:min(10, len(token.AccessToken))]) + log.Printf("OAuth2 Callback Info: TokenReceiver did not return a valid user ID for token: %s...", token.AccessToken[:customMin(10, len(token.AccessToken))]) } log.Printf("OAuth2 Callback Success: Exchanged code for token, UserID: %d", userID) return userID, nil } -func (o *OAuth2Service) GetToken(code string) (*oauth2.Token, error) { +func (o *Service) GetToken(code string) (*oauth2.Token, error) { opts := []oauth2.AuthCodeOption{ oauth2.SetAuthURLParam("code_verifier", o.codeVerifier), } return o.config.Exchange(context.Background(), code, opts...) } -func (o *OAuth2Service) GetClient(token *oauth2.Token) *http.Client { +func (o *Service) GetClient(token *oauth2.Token) *http.Client { return o.config.Client(context.Background(), token) } -func (o *OAuth2Service) RefreshToken(token *oauth2.Token) (*oauth2.Token, error) { +func (o *Service) RefreshToken(token *oauth2.Token) (*oauth2.Token, error) { source := o.config.TokenSource(context.Background(), token) return oauth2.ReuseTokenSource(token, source).Token() } -func min(a, b int) int { +func customMin(a, b int) int { if a < b { return a } diff --git a/oauth/oauth_manager.go b/oauth/oauth_manager.go index 1ee20f4..9707972 100644 --- a/oauth/oauth_manager.go +++ b/oauth/oauth_manager.go @@ -1,4 +1,4 @@ -// Modify piper/oauth/oauth_manager.go +// Package oauth Modify piper/oauth/oauth_manager.go package oauth import ( @@ -8,37 +8,37 @@ import ( "sync" ) -// manages multiple oauth client services -type OAuthServiceManager struct { +// ServiceManager OAuthServiceManager manages multiple oauth client services +type ServiceManager struct { services map[string]AuthService mu sync.RWMutex logger *log.Logger } -func NewOAuthServiceManager() *OAuthServiceManager { - return &OAuthServiceManager{ +func NewOAuthServiceManager() *ServiceManager { + return &ServiceManager{ services: make(map[string]AuthService), logger: log.New(log.Writer(), "oauth: ", log.LstdFlags|log.Lmsgprefix), } } -// registers any service that impls AuthService -func (m *OAuthServiceManager) RegisterService(name string, service AuthService) { +// RegisterService registers any service that impls AuthService +func (m *ServiceManager) RegisterService(name string, service AuthService) { m.mu.Lock() defer m.mu.Unlock() m.services[name] = service m.logger.Printf("Registered auth service: %s", name) } -// get an AuthService by registered name -func (m *OAuthServiceManager) GetService(name string) (AuthService, bool) { +// GetService get an AuthService by registered name +func (m *ServiceManager) GetService(name string) (AuthService, bool) { m.mu.RLock() defer m.mu.RUnlock() service, exists := m.services[name] return service, exists } -func (m *OAuthServiceManager) HandleLogin(serviceName string) http.HandlerFunc { +func (m *ServiceManager) HandleLogin(serviceName string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { m.mu.RLock() service, exists := m.services[serviceName] @@ -54,7 +54,7 @@ func (m *OAuthServiceManager) HandleLogin(serviceName string) http.HandlerFunc { } } -func (m *OAuthServiceManager) HandleLogout(serviceName string) http.HandlerFunc { +func (m *ServiceManager) HandleLogout(serviceName string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { m.mu.RLock() service, exists := m.services[serviceName] @@ -70,7 +70,7 @@ func (m *OAuthServiceManager) HandleLogout(serviceName string) http.HandlerFunc } } -func (m *OAuthServiceManager) HandleCallback(serviceName string) http.HandlerFunc { +func (m *ServiceManager) HandleCallback(serviceName string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { m.mu.RLock() service, exists := m.services[serviceName] diff --git a/oauth/service.go b/oauth/service.go index 732c74d..fff06aa 100644 --- a/oauth/service.go +++ b/oauth/service.go @@ -5,18 +5,18 @@ import ( ) type AuthService interface { - // inits the login flow for the service + // HandleLogin inits the login flow for the service HandleLogin(w http.ResponseWriter, r *http.Request) - // handles the callback for the provider. is responsible for inserting + // HandleCallback handles the callback for the provider. is responsible for inserting // sessions in the db HandleCallback(w http.ResponseWriter, r *http.Request) (int64, error) HandleLogout(w http.ResponseWriter, r *http.Request) } -// optional but recommended +// TokenReceiver optional but recommended type TokenReceiver interface { - // stores the access token in the db + // SetAccessToken stores the access token in the db // if there is a session, will associate the token with the session SetAccessToken(token string, refreshToken string, currentId int64, hasSession bool) (int64, error) } diff --git a/service/apikey/apikey.go b/service/apikey/apikey.go index eb618c3..7abcd7e 100644 --- a/service/apikey/apikey.go +++ b/service/apikey/apikey.go @@ -8,17 +8,17 @@ import ( "time" "github.com/teal-fm/piper/db" - db_apikey "github.com/teal-fm/piper/db/apikey" // Assuming this is the package for ApiKey struct + dbapikey "github.com/teal-fm/piper/db/apikey" // Assuming this is the package for ApiKey struct "github.com/teal-fm/piper/pages" "github.com/teal-fm/piper/session" ) type Service struct { db *db.DB - sessions *session.SessionManager + sessions *session.Manager } -func NewAPIKeyService(database *db.DB, sessionManager *session.SessionManager) *Service { +func NewAPIKeyService(database *db.DB, sessionManager *session.Manager) *Service { return &Service{ db: database, sessions: sessionManager, @@ -204,8 +204,8 @@ func (s *Service) HandleAPIKeyManagement(database *db.DB, pg *pages.Pages) http. } data := struct { - Keys []*db_apikey.ApiKey // Assuming GetUserApiKeys returns this type - NewKeyID string // Changed from NewKey for clarity as it's an ID + Keys []*dbapikey.ApiKey // Assuming GetUserApiKeys returns this type + NewKeyID string // Changed from NewKey for clarity as it's an ID NavBar pages.NavBar }{ Keys: keys, diff --git a/service/applemusic/applemusic.go b/service/applemusic/applemusic.go index 605ea60..a4d2961 100644 --- a/service/applemusic/applemusic.go +++ b/service/applemusic/applemusic.go @@ -42,8 +42,8 @@ type Service struct { // ingestion deps DB *db.DB - atprotoService *atprotoauth.ATprotoAuthService - mbService *musicbrainz.MusicBrainzService + atprotoService *atprotoauth.AuthService + mbService *musicbrainz.Service playingNowService interface { PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error ClearPlayingNow(ctx context.Context, userID int64) error @@ -73,7 +73,7 @@ func (s *Service) WithPersistence( } // WithDeps wires services needed for ingestion -func (s *Service) WithDeps(database *db.DB, atproto *atprotoauth.ATprotoAuthService, mb *musicbrainz.MusicBrainzService, playingNowService interface { +func (s *Service) WithDeps(database *db.DB, atproto *atprotoauth.AuthService, mb *musicbrainz.Service, playingNowService interface { PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error ClearPlayingNow(ctx context.Context, userID int64) error }) *Service { @@ -331,7 +331,7 @@ func (s *Service) FetchRecentPlayedTracks(ctx context.Context, userToken string, } // toTrack converts appleRecentTrack to internal models.Track -func (s *Service) toTrack(t appleRecentTrack, userID int64) *models.Track { +func (s *Service) toTrack(t appleRecentTrack) *models.Track { var duration int64 if t.Attributes.DurationInMillis != nil { duration = *t.Attributes.DurationInMillis @@ -427,7 +427,7 @@ func (s *Service) ProcessUser(ctx context.Context, user *models.User) error { } // Convert to internal track format - track := s.toTrack(*currentAppleTrack, user.ID) + track := s.toTrack(*currentAppleTrack) if track == nil || strings.TrimSpace(track.Name) == "" || len(track.Artist) == 0 { s.logger.Printf("invalid track data for user %d", user.ID) return nil diff --git a/service/atproto/submission.go b/service/atproto/submission.go index f30d98d..7dfa819 100644 --- a/service/atproto/submission.go +++ b/service/atproto/submission.go @@ -15,7 +15,7 @@ import ( ) // SubmitPlayToPDS submits a track play to the ATProto PDS as a feed.play record -func SubmitPlayToPDS(ctx context.Context, did string, mostRecentAtProtoSessionID string, track *models.Track, atprotoService *atprotoauth.ATprotoAuthService) error { +func SubmitPlayToPDS(ctx context.Context, did string, mostRecentAtProtoSessionID string, track *models.Track, atprotoService *atprotoauth.AuthService) error { if did == "" { return fmt.Errorf("DID cannot be empty") } diff --git a/service/lastfm/lastfm.go b/service/lastfm/lastfm.go index 4444dbd..f4b21f1 100644 --- a/service/lastfm/lastfm.go +++ b/service/lastfm/lastfm.go @@ -23,17 +23,16 @@ import ( const ( lastfmAPIBaseURL = "https://ws.audioscrobbler.com/2.0/" - defaultLimit = 1 // Default number of tracks to fetch per user ) -type LastFMService struct { +type Service struct { db *db.DB httpClient *http.Client limiter *rate.Limiter apiKey string Usernames []string - musicBrainzService *musicbrainz.MusicBrainzService - atprotoService *atprotoauth.ATprotoAuthService + musicBrainzService *musicbrainz.Service + atprotoService *atprotoauth.AuthService playingNowService interface { PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error ClearPlayingNow(ctx context.Context, userID int64) error @@ -43,13 +42,13 @@ type LastFMService struct { logger *log.Logger } -func NewLastFMService(db *db.DB, apiKey string, musicBrainzService *musicbrainz.MusicBrainzService, atprotoService *atprotoauth.ATprotoAuthService, playingNowService interface { +func NewLastFMService(db *db.DB, apiKey string, musicBrainzService *musicbrainz.Service, atprotoService *atprotoauth.AuthService, playingNowService interface { PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error ClearPlayingNow(ctx context.Context, userID int64) error -}) *LastFMService { +}) *Service { logger := log.New(os.Stdout, "lastfm: ", log.LstdFlags|log.Lmsgprefix) - return &LastFMService{ + return &Service{ db: db, httpClient: &http.Client{ Timeout: 10 * time.Second, @@ -67,7 +66,7 @@ func NewLastFMService(db *db.DB, apiKey string, musicBrainzService *musicbrainz. } } -func (l *LastFMService) loadUsernames() error { +func (l *Service) loadUsernames() error { u, err := l.db.GetAllUsersWithLastFM() if err != nil { l.logger.Printf("Error loading users with Last.fm from DB: %v", err) @@ -98,7 +97,7 @@ func (l *LastFMService) loadUsernames() error { } // getRecentTracks fetches the most recent tracks for a given Last.fm user. -func (l *LastFMService) getRecentTracks(ctx context.Context, username string, limit int) (*RecentTracksResponse, error) { +func (l *Service) getRecentTracks(ctx context.Context, username string, limit int) (*RecentTracksResponse, error) { if username == "" { return nil, fmt.Errorf("username cannot be empty") } @@ -159,7 +158,7 @@ func (l *LastFMService) getRecentTracks(ctx context.Context, username string, li return &recentTracksResp, nil } -func (l *LastFMService) StartListeningTracker(interval time.Duration) { +func (l *Service) StartListeningTracker(interval time.Duration) { if err := l.loadUsernames(); err != nil { l.logger.Printf("Failed to perform initial username load: %v", err) // Decide if we should proceed without initial load or return error @@ -207,7 +206,7 @@ func (l *LastFMService) StartListeningTracker(interval time.Duration) { } // fetchAllUserTracks iterates through users and fetches their tracks. -func (l *LastFMService) fetchAllUserTracks(ctx context.Context) { +func (l *Service) fetchAllUserTracks(ctx context.Context) { l.logger.Printf("Starting fetch cycle for %d users...", len(l.Usernames)) var wg sync.WaitGroup // Use WaitGroup to fetch concurrently (optional) fetchErrors := make(chan error, len(l.Usernames)) // Channel for errors @@ -266,7 +265,7 @@ func (l *LastFMService) fetchAllUserTracks(ctx context.Context) { } } -func (l *LastFMService) processTracks(ctx context.Context, username string, tracks []Track) error { +func (l *Service) processTracks(ctx context.Context, username string, tracks []Track) error { if l.db == nil { return fmt.Errorf("database connection is nil") } @@ -413,13 +412,13 @@ func (l *LastFMService) processTracks(ctx context.Context, username string, trac return nil } -func (l *LastFMService) SubmitTrackToPDS(did string, mostRecentAtProtoSessionID string, track *models.Track, ctx context.Context) error { +func (l *Service) SubmitTrackToPDS(did string, mostRecentAtProtoSessionID string, track *models.Track, ctx context.Context) error { // Use shared atproto service for submission return atprotoservice.SubmitPlayToPDS(ctx, did, mostRecentAtProtoSessionID, track, l.atprotoService) } // convertLastFMTrackToModelsTrack converts a Last.fm Track to models.Track format -func (l *LastFMService) convertLastFMTrackToModelsTrack(track Track) *models.Track { +func (l *Service) convertLastFMTrackToModelsTrack(track Track) *models.Track { // Create artist array artists := []models.Artist{ { diff --git a/service/lastfm/model.go b/service/lastfm/model.go index 24b1d77..1e2a036 100644 --- a/service/lastfm/model.go +++ b/service/lastfm/model.go @@ -6,7 +6,7 @@ import ( "time" ) -// Structs to represent the Last.fm API response for user.getrecenttracks +// RecentTracksResponse Structs to represent the Last.fm API response for user.getrecenttracks type RecentTracksResponse struct { RecentTracks RecentTracks `json:"recenttracks"` } @@ -25,7 +25,7 @@ type Track struct { Name string `json:"name"` URL string `json:"url"` Date *TrackDate `json:"date,omitempty"` // Use pointer for optional fields - Attr *struct { // Custom handling for @attr.nowplaying + Attr *struct { // Custom handling for @attr.nowplaying NowPlaying string `json:"nowplaying"` // Field name corrected to match struct tag } `json:"@attr,omitempty"` // This captures the @attr object within the track } diff --git a/service/musicbrainz/musicbrainz.go b/service/musicbrainz/musicbrainz.go index 4bf691b..2c502c4 100644 --- a/service/musicbrainz/musicbrainz.go +++ b/service/musicbrainz/musicbrainz.go @@ -19,8 +19,8 @@ import ( "golang.org/x/time/rate" ) -// MusicBrainz API Types -type MusicBrainzArtistCredit struct { +// ArtistCredit API Types +type ArtistCredit struct { Artist struct { ID string `json:"id"` Name string `json:"name"` @@ -30,7 +30,7 @@ type MusicBrainzArtistCredit struct { Name string `json:"name"` } -type MusicBrainzRelease struct { +type Release struct { ID string `json:"id"` Title string `json:"title"` Status string `json:"status,omitempty"` @@ -40,20 +40,20 @@ type MusicBrainzRelease struct { TrackCount int `json:"track-count,omitempty"` } -type MusicBrainzRecording struct { - ID string `json:"id"` - Title string `json:"title"` - Length int `json:"length,omitempty"` // milliseconds - ISRCs []string `json:"isrcs,omitempty"` - ArtistCredit []MusicBrainzArtistCredit `json:"artist-credit,omitempty"` - Releases []MusicBrainzRelease `json:"releases,omitempty"` +type Recording struct { + ID string `json:"id"` + Title string `json:"title"` + Length int `json:"length,omitempty"` // milliseconds + ISRCs []string `json:"isrcs,omitempty"` + ArtistCredit []ArtistCredit `json:"artist-credit,omitempty"` + Releases []Release `json:"releases,omitempty"` } -type MusicBrainzSearchResponse struct { - Created time.Time `json:"created"` - Count int `json:"count"` - Offset int `json:"offset"` - Recordings []MusicBrainzRecording `json:"recordings"` +type SearchResponse struct { + Created time.Time `json:"created"` + Count int `json:"count"` + Offset int `json:"offset"` + Recordings []Recording `json:"recordings"` } type SearchParams struct { @@ -64,11 +64,11 @@ type SearchParams struct { // cacheEntry holds the cached data and its expiration time. type cacheEntry struct { - recordings []MusicBrainzRecording + recordings []Recording expiresAt time.Time } -type MusicBrainzService struct { +type Service struct { db *db.DB httpClient *http.Client limiter *rate.Limiter @@ -80,13 +80,13 @@ type MusicBrainzService struct { } // NewMusicBrainzService creates a new service instance with rate limiting and caching. -func NewMusicBrainzService(db *db.DB) *MusicBrainzService { +func NewMusicBrainzService(db *db.DB) *Service { // MusicBrainz allows 1 request per second limiter := rate.NewLimiter(rate.Every(time.Second), 1) // Set a default cache TTL (e.g., 1 hour) defaultCacheTTL := 1 * time.Hour logger := log.New(os.Stdout, "musicbrainz: ", log.LstdFlags|log.Lmsgprefix) - return &MusicBrainzService{ + return &Service{ db: db, httpClient: &http.Client{ Timeout: 10 * time.Second, @@ -111,7 +111,7 @@ func generateCacheKey(params SearchParams) string { } // SearchMusicBrainz searches the MusicBrainz API for recordings, using an in-memory cache. -func (s *MusicBrainzService) SearchMusicBrainz(ctx context.Context, params SearchParams) ([]MusicBrainzRecording, error) { +func (s *Service) SearchMusicBrainz(ctx context.Context, params SearchParams) ([]Recording, error) { // Validate parameters first if params.Track == "" && params.Artist == "" && params.Release == "" { return nil, fmt.Errorf("at least one search parameter (Track, Artist, Release) must be provided") @@ -142,7 +142,7 @@ func (s *MusicBrainzService) SearchMusicBrainz(ctx context.Context, params Searc } // --- Proceed with API call --- - queryParts := []string{} + var queryParts []string if params.Track != "" { queryParts = append(queryParts, fmt.Sprintf(`recording:"%s"`, params.Track)) } @@ -182,7 +182,7 @@ func (s *MusicBrainzService) SearchMusicBrainz(ctx context.Context, params Searc return nil, fmt.Errorf("MusicBrainz API request to %s returned status %d", endpoint, resp.StatusCode) } - var result MusicBrainzSearchResponse + var result SearchResponse if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { return nil, fmt.Errorf("failed to decode response from %s: %w", endpoint, err) } @@ -201,7 +201,7 @@ func (s *MusicBrainzService) SearchMusicBrainz(ctx context.Context, params Searc } // GetBestRelease selects the 'best' release from a list based on specific criteria. -func (s *MusicBrainzService) GetBestRelease(releases []MusicBrainzRelease, trackTitle string) *MusicBrainzRelease { +func (s *Service) GetBestRelease(releases []Release, trackTitle string) *Release { if len(releases) == 0 { return nil } @@ -259,7 +259,7 @@ func (s *MusicBrainzService) GetBestRelease(releases []MusicBrainzRelease, track return &r } -func HydrateTrack(mb *MusicBrainzService, track models.Track) (*models.Track, error) { +func HydrateTrack(mb *Service, track models.Track) (*models.Track, error) { ctx := context.Background() // array of strings artistArray := make([]string, len(track.Artist)) // Assuming Name is string type diff --git a/service/playingnow/playingnow.go b/service/playingnow/playingnow.go index d610a2e..d2c18e0 100644 --- a/service/playingnow/playingnow.go +++ b/service/playingnow/playingnow.go @@ -2,6 +2,7 @@ package playingnow import ( "context" + "errors" "fmt" "log" "os" @@ -20,20 +21,20 @@ import ( atprotoauth "github.com/teal-fm/piper/oauth/atproto" ) -// PlayingNowService handles publishing current playing status to ATProto -type PlayingNowService struct { +// Service handles publishing current playing status to ATProto +type Service struct { db *db.DB - atprotoService *atprotoauth.ATprotoAuthService + atprotoService *atprotoauth.AuthService logger *log.Logger mu sync.RWMutex clearedStatus map[int64]bool // tracks if a user's status has been cleared on their repo } // NewPlayingNowService creates a new playing now service -func NewPlayingNowService(database *db.DB, atprotoService *atprotoauth.ATprotoAuthService) *PlayingNowService { +func NewPlayingNowService(database *db.DB, atprotoService *atprotoauth.AuthService) *Service { logger := log.New(os.Stdout, "playingnow: ", log.LstdFlags|log.Lmsgprefix) - return &PlayingNowService{ + return &Service{ db: database, atprotoService: atprotoService, logger: logger, @@ -42,12 +43,15 @@ func NewPlayingNowService(database *db.DB, atprotoService *atprotoauth.ATprotoAu } // PublishPlayingNow publishes a currently playing track as actor status -func (p *PlayingNowService) PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error { +func (p *Service) PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error { // Get user information to find their DID user, err := p.db.GetUserByID(userID) if err != nil { return fmt.Errorf("failed to get user: %w", err) } + if user == nil { + return fmt.Errorf("user not found user ID: %d", userID) + } if user.ATProtoDID == nil { p.logger.Printf("User %d has no ATProto DID, skipping playing now", userID) @@ -117,7 +121,7 @@ func (p *PlayingNowService) PublishPlayingNow(ctx context.Context, userID int64, } // ClearPlayingNow removes the current playing status by setting an expired status -func (p *PlayingNowService) ClearPlayingNow(ctx context.Context, userID int64) error { +func (p *Service) ClearPlayingNow(ctx context.Context, userID int64) error { // Check if status is already cleared to avoid clearing on the users repo over and over p.mu.RLock() alreadyCleared := p.clearedStatus[userID] @@ -133,6 +137,10 @@ func (p *PlayingNowService) ClearPlayingNow(ctx context.Context, userID int64) e return fmt.Errorf("failed to get user: %w", err) } + if user == nil { + return fmt.Errorf("user not found user ID: %d", userID) + } + if user.ATProtoDID == nil { p.logger.Printf("User %d has no ATProto DID, skipping clear playing now", userID) return nil @@ -200,7 +208,7 @@ func (p *PlayingNowService) ClearPlayingNow(ctx context.Context, userID int64) e } // trackToPlayView converts a models.Track to teal.AlphaFeedDefs_PlayView -func (p *PlayingNowService) trackToPlayView(track *models.Track) (*teal.AlphaFeedDefs_PlayView, error) { +func (p *Service) trackToPlayView(track *models.Track) (*teal.AlphaFeedDefs_PlayView, error) { if track.Name == "" { return nil, fmt.Errorf("track name cannot be empty") } @@ -273,11 +281,12 @@ func (p *PlayingNowService) trackToPlayView(track *models.Track) (*teal.AlphaFee // getStatusSwapRecord retrieves the current swap record (CID) for the actor status record. // Returns (nil, nil) if the record does not exist yet. -func (p *PlayingNowService) getStatusSwapRecord(ctx context.Context, atApiClient *client.APIClient) (*comatproto.RepoGetRecord_Output, error) { +func (p *Service) getStatusSwapRecord(ctx context.Context, atApiClient *client.APIClient) (*comatproto.RepoGetRecord_Output, error) { result, err := comatproto.RepoGetRecord(ctx, atApiClient, "", "fm.teal.alpha.actor.status", atApiClient.AccountDID.String(), "self") if err != nil { - xErr, ok := err.(*client.APIError) + var xErr *client.APIError + ok := errors.As(err, &xErr) if !ok { return nil, fmt.Errorf("error getting the record: %w", err) } diff --git a/service/playingnow/playingnow_test.go b/service/playingnow/playingnow_test.go index df7d253..c4ff740 100644 --- a/service/playingnow/playingnow_test.go +++ b/service/playingnow/playingnow_test.go @@ -21,7 +21,7 @@ func TestTrackToPlayView(t *testing.T) { } // Mock ATProto service (we'll just test the conversion, not the actual submission) - service := &PlayingNowService{ + service := &Service{ db: database, logger: nil, // We'll skip logging in tests } @@ -97,7 +97,7 @@ func TestTrackToPlayView(t *testing.T) { } func TestTrackToPlayViewEmptyTrack(t *testing.T) { - service := &PlayingNowService{} + service := &Service{} // Test with empty track name (should fail) track := &models.Track{ @@ -112,7 +112,7 @@ func TestTrackToPlayViewEmptyTrack(t *testing.T) { } func TestTrackToPlayViewMinimal(t *testing.T) { - service := &PlayingNowService{} + service := &Service{} // Test with minimal track data track := &models.Track{ diff --git a/service/spotify/playlists.go b/service/spotify/playlists.go index e5968ca..8aa10c2 100644 --- a/service/spotify/playlists.go +++ b/service/spotify/playlists.go @@ -22,7 +22,7 @@ type PlaylistResponse struct { Items []Playlist `json:"items"` } -func (s *SpotifyService) getUserPlaylists(userID int64) (*PlaylistResponse, error) { +func (s *Service) getUserPlaylists(userID int64) (*PlaylistResponse, error) { s.mu.RLock() token, exists := s.userTokens[userID] s.mu.RUnlock() diff --git a/service/spotify/spotify.go b/service/spotify/spotify.go index 3ad6f63..b6c8cdb 100644 --- a/service/spotify/spotify.go +++ b/service/spotify/spotify.go @@ -28,10 +28,10 @@ import ( "github.com/teal-fm/piper/session" ) -type SpotifyService struct { +type Service struct { DB *db.DB - atprotoService *atprotoauth.ATprotoAuthService // Added field - mb *musicbrainz.MusicBrainzService // Added field + atprotoService *atprotoauth.AuthService // Added field + mb *musicbrainz.Service // Added field playingNowService interface { PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error ClearPlayingNow(ctx context.Context, userID int64) error @@ -42,13 +42,13 @@ type SpotifyService struct { logger *log.Logger } -func NewSpotifyService(database *db.DB, atprotoService *atprotoauth.ATprotoAuthService, musicBrainzService *musicbrainz.MusicBrainzService, playingNowService interface { +func NewSpotifyService(database *db.DB, atprotoService *atprotoauth.AuthService, musicBrainzService *musicbrainz.Service, playingNowService interface { PublishPlayingNow(ctx context.Context, userID int64, track *models.Track) error ClearPlayingNow(ctx context.Context, userID int64) error -}) *SpotifyService { +}) *Service { logger := log.New(os.Stdout, "spotify: ", log.LstdFlags|log.Lmsgprefix) - return &SpotifyService{ + return &Service{ DB: database, atprotoService: atprotoService, mb: musicBrainzService, @@ -59,7 +59,7 @@ func NewSpotifyService(database *db.DB, atprotoService *atprotoauth.ATprotoAuthS } } -func (s *SpotifyService) SubmitTrackToPDS(did string, mostRecentAtProtoSessionID string, track *models.Track, ctx context.Context) error { +func (s *Service) SubmitTrackToPDS(did string, mostRecentAtProtoSessionID string, track *models.Track, ctx context.Context) error { //Had a empty feed.play get submitted not sure why. Tracking here if track.Name == "" { s.logger.Println("Track name is empty. Skipping submission. Please record the logs before and send to the teal.fm Discord") @@ -70,7 +70,7 @@ func (s *SpotifyService) SubmitTrackToPDS(did string, mostRecentAtProtoSessionID return atprotoservice.SubmitPlayToPDS(ctx, did, mostRecentAtProtoSessionID, track, s.atprotoService) } -func (s *SpotifyService) SetAccessToken(token string, refreshToken string, userId int64, hasSession bool) (int64, error) { +func (s *Service) SetAccessToken(token string, refreshToken string, userId int64, hasSession bool) (int64, error) { userID, err := s.identifyAndStoreUser(token, refreshToken, userId, hasSession) if err != nil { s.logger.Printf("Error identifying and storing user: %v", err) @@ -79,7 +79,7 @@ func (s *SpotifyService) SetAccessToken(token string, refreshToken string, userI return userID, nil } -func (s *SpotifyService) identifyAndStoreUser(token string, refreshToken string, userId int64, hasSession bool) (int64, error) { +func (s *Service) identifyAndStoreUser(token string, refreshToken string, userId int64, hasSession bool) (int64, error) { userProfile, err := s.fetchSpotifyProfile(token) if err != nil { s.logger.Printf("Error fetching Spotify profile: %v", err) @@ -102,13 +102,13 @@ func (s *SpotifyService) identifyAndStoreUser(token string, refreshToken string, if !hasSession { s.logger.Printf("User does not seem to exist") return 0, fmt.Errorf("user does not seem to exist") - } else { - // overwrite prev user - user, err = s.DB.AddSpotifySession(userId, userProfile.DisplayName, userProfile.Email, userProfile.ID, token, refreshToken, tokenExpiryTime) - if err != nil { - s.logger.Printf("Error adding Spotify session for user ID %d: %v", userId, err) - return 0, err - } + } + + // overwrite prev user + user, err = s.DB.AddSpotifySession(userId, userProfile.DisplayName, userProfile.Email, userProfile.ID, token, refreshToken, tokenExpiryTime) + if err != nil { + s.logger.Printf("Error adding Spotify session for user ID %d: %v", userId, err) + return 0, err } } else { err = s.DB.UpdateUserToken(user.ID, token, refreshToken, tokenExpiryTime) @@ -119,6 +119,10 @@ func (s *SpotifyService) identifyAndStoreUser(token string, refreshToken string, s.logger.Printf("Updated token for existing user: %s (ID: %d)", *user.Username, user.ID) } } + if user == nil { + return 0, fmt.Errorf("user does not seem to exist") + } + user.AccessToken = &token user.TokenExpiry = &tokenExpiryTime @@ -136,7 +140,7 @@ type spotifyProfile struct { Email string `json:"email"` } -func (s *SpotifyService) LoadAllUsers() error { +func (s *Service) LoadAllUsers() error { users, err := s.DB.GetAllActiveUsers() if err != nil { return fmt.Errorf("error loading users: %v", err) @@ -170,7 +174,7 @@ func (s *SpotifyService) LoadAllUsers() error { return nil } -func (s *SpotifyService) UnloadAllUsers() error { +func (s *Service) UnloadAllUsers() error { s.mu.Lock() defer s.mu.Unlock() s.userTokens = make(map[int64]string) @@ -179,7 +183,7 @@ func (s *SpotifyService) UnloadAllUsers() error { // refreshTokenInner handles the actual Spotify token refresh logic. // It returns the new access token or an error. -func (s *SpotifyService) refreshTokenInner(userID int64) (string, error) { +func (s *Service) refreshTokenInner(userID int64) (string, error) { user, err := s.DB.GetUserByID(userID) if err != nil { return "", fmt.Errorf("error loading user %d for refresh: %w", userID, err) @@ -276,13 +280,13 @@ func (s *SpotifyService) refreshTokenInner(userID int64) (string, error) { // RefreshToken attempts to refresh the token for a given user ID. // It's less commonly needed now refreshTokenInner handles fetching the user. -func (s *SpotifyService) RefreshToken(userID int64) error { +func (s *Service) RefreshToken(userID int64) error { _, err := s.refreshTokenInner(userID) return err } -// attempt to refresh expired tokens -func (s *SpotifyService) RefreshExpiredTokens() { +// RefreshExpiredTokens attempt to refresh expired tokens +func (s *Service) RefreshExpiredTokens() { users, err := s.DB.GetUsersWithExpiredTokens() if err != nil { s.logger.Printf("Error fetching users with expired tokens: %v", err) @@ -311,7 +315,7 @@ func (s *SpotifyService) RefreshExpiredTokens() { } } -func (s *SpotifyService) fetchSpotifyProfile(token string) (*spotifyProfile, error) { +func (s *Service) fetchSpotifyProfile(token string) (*spotifyProfile, error) { req, err := http.NewRequest("GET", "https://api.spotify.com/v1/me", nil) if err != nil { return nil, err @@ -338,7 +342,7 @@ func (s *SpotifyService) fetchSpotifyProfile(token string) (*spotifyProfile, err return &profile, nil } -func (s *SpotifyService) HandleCurrentTrack(w http.ResponseWriter, r *http.Request) { +func (s *Service) HandleCurrentTrack(w http.ResponseWriter, r *http.Request) { userID, ok := session.GetUserID(r.Context()) if !ok { http.Error(w, "User not authenticated", http.StatusUnauthorized) @@ -358,7 +362,7 @@ func (s *SpotifyService) HandleCurrentTrack(w http.ResponseWriter, r *http.Reque json.NewEncoder(w).Encode(track) } -func (s *SpotifyService) HandleTrackHistory(w http.ResponseWriter, r *http.Request) { +func (s *Service) HandleTrackHistory(w http.ResponseWriter, r *http.Request) { userID, ok := session.GetUserID(r.Context()) if !ok { http.Error(w, "User not authenticated", http.StatusUnauthorized) @@ -376,7 +380,7 @@ func (s *SpotifyService) HandleTrackHistory(w http.ResponseWriter, r *http.Reque json.NewEncoder(w).Encode(tracks) } -func (s *SpotifyService) FetchCurrentTrack(userID int64) (*models.Track, error) { +func (s *Service) FetchCurrentTrack(userID int64) (*models.Track, error) { s.mu.RLock() token, exists := s.userTokens[userID] s.mu.RUnlock() @@ -512,7 +516,7 @@ func (s *SpotifyService) FetchCurrentTrack(userID int64) (*models.Track, error) return track, nil } -func (s *SpotifyService) fetchAllUserTracks(ctx context.Context) { +func (s *Service) fetchAllUserTracks(ctx context.Context) { // copy userIDs to avoid holding the lock too long s.mu.RLock() userIDs := make([]int64, 0, len(s.userTokens)) @@ -568,7 +572,7 @@ func (s *SpotifyService) fetchAllUserTracks(ctx context.Context) { currentTrack.DurationMs > 30000 // just log when we stamp tracks - if isNewTrack && isLastTrackStamped && !currentTrack.HasStamped { + if isNewTrack && isLastTrackStamped && currentTrack != nil && !currentTrack.HasStamped { artistName := "Unknown Artist" if len(currentTrack.Artist) > 0 { artistName = currentTrack.Artist[0].Name @@ -576,8 +580,11 @@ func (s *SpotifyService) fetchAllUserTracks(ctx context.Context) { s.logger.Printf("User %d stamped (previous) track: %s by %s", userID, currentTrack.Name, artistName) currentTrack.HasStamped = true if currentTrack.PlayID != 0 { - s.DB.UpdateTrack(currentTrack.PlayID, currentTrack) - + err := s.DB.UpdateTrack(currentTrack.PlayID, currentTrack) + if err != nil { + s.logger.Printf("Error updating track %d in DB: %v", currentTrack.PlayID, err) + return + } s.logger.Printf("Updated!") } } @@ -591,7 +598,11 @@ func (s *SpotifyService) fetchAllUserTracks(ctx context.Context) { track.HasStamped = true // if currenttrack has a playid and the last track is the same as the current track if !isNewTrack && currentTrack.PlayID != 0 { - s.DB.UpdateTrack(currentTrack.PlayID, track) + err := s.DB.UpdateTrack(currentTrack.PlayID, track) + if err != nil { + s.logger.Printf("Error updating track %d in DB: %v", currentTrack.PlayID, err) + return + } // Update in memory s.mu.Lock() @@ -640,8 +651,8 @@ func (s *SpotifyService) fetchAllUserTracks(ctx context.Context) { s.logger.Printf("User %d (%d): ATProto DID not set. Skipping PDS submission for track '%s'.", userID, dbUser.ATProtoDID, track.Name) } else { // User has a DID, proceed with hydration and submission - var trackToSubmitToPDS *models.Track = track // Default to the original track (already *models.Track) - if s.mb != nil { // Check if MusicBrainz service is available + var trackToSubmitToPDS = track // Default to the original track (already *models.Track) + if s.mb != nil { // Check if MusicBrainz service is available // musicbrainz.HydrateTrack expects models.Track as second argument, so we pass *track // and it returns *models.Track hydratedTrack, errHydrate := musicbrainz.HydrateTrack(s.mb, *track) @@ -679,7 +690,7 @@ func (s *SpotifyService) fetchAllUserTracks(ctx context.Context) { } } -func (s *SpotifyService) StartListeningTracker(interval time.Duration) { +func (s *Service) StartListeningTracker(interval time.Duration) { ticker := time.NewTicker(interval) go func() { diff --git a/session/session.go b/session/session.go index 99c772e..7399ca3 100644 --- a/session/session.go +++ b/session/session.go @@ -15,12 +15,8 @@ import ( "github.com/teal-fm/piper/db/apikey" ) -// session/session.go +// Session session/session.go. Web session for Piper type Session struct { - - //need to re work this. May add onto it for atproto oauth. But need to be careful about that expiresd - //Maybe a speerate oauth session store table and it has a created date? yeah do that then can look it up by session id from this table for user actions - ID string UserID int64 ATProtoSessionID string @@ -28,14 +24,14 @@ type Session struct { ExpiresAt time.Time } -type SessionManager struct { +type Manager struct { db *db.DB sessions map[string]*Session // use in memory cache if necessary - ApiKeyMgr *apikey.ApiKeyManager + ApiKeyMgr *apikey.Manager mu sync.RWMutex } -func NewSessionManager(database *db.DB) *SessionManager { +func NewSessionManager(database *db.DB) *Manager { _, err := database.Exec(` CREATE TABLE IF NOT EXISTS sessions ( @@ -53,15 +49,15 @@ func NewSessionManager(database *db.DB) *SessionManager { apiKeyMgr := apikey.NewApiKeyManager(database) - return &SessionManager{ + return &Manager{ db: database, sessions: make(map[string]*Session), ApiKeyMgr: apiKeyMgr, } } -// create a new session for a user -func (sm *SessionManager) CreateSession(userID int64, atProtoSessionId string) *Session { +// CreateSession create a new session for a user +func (sm *Manager) CreateSession(userID int64, atProtoSessionId string) *Session { sm.mu.Lock() defer sm.mu.Unlock() @@ -99,8 +95,8 @@ func (sm *SessionManager) CreateSession(userID int64, atProtoSessionId string) * return session } -// retrieve a session by ID -func (sm *SessionManager) GetSession(sessionID string) (*Session, bool) { +// GetSession retrieve a session by ID +func (sm *Manager) GetSession(sessionID string) (*Session, bool) { // First check in-memory cache sm.mu.RLock() session, exists := sm.sessions[sessionID] @@ -144,8 +140,8 @@ func (sm *SessionManager) GetSession(sessionID string) (*Session, bool) { return nil, false } -// remove a session -func (sm *SessionManager) DeleteSession(sessionID string) { +// DeleteSession remove a session +func (sm *Manager) DeleteSession(sessionID string) { sm.mu.Lock() delete(sm.sessions, sessionID) sm.mu.Unlock() @@ -158,8 +154,8 @@ func (sm *SessionManager) DeleteSession(sessionID string) { } } -// set a session cookie for the user -func (sm *SessionManager) SetSessionCookie(w http.ResponseWriter, session *Session) { +// SetSessionCookie set a session cookie for the user +func (sm *Manager) SetSessionCookie(w http.ResponseWriter, session *Session) { cookie := &http.Cookie{ Name: "session", Value: session.ID, @@ -172,7 +168,7 @@ func (sm *SessionManager) SetSessionCookie(w http.ResponseWriter, session *Sessi } // ClearSessionCookie clears the session cookie -func (sm *SessionManager) ClearSessionCookie(w http.ResponseWriter) { +func (sm *Manager) ClearSessionCookie(w http.ResponseWriter) { cookie := &http.Cookie{ Name: "session", Value: "", @@ -184,16 +180,16 @@ func (sm *SessionManager) ClearSessionCookie(w http.ResponseWriter) { http.SetCookie(w, cookie) } -func (sm *SessionManager) GetAPIKeyManager() *apikey.ApiKeyManager { +func (sm *Manager) GetAPIKeyManager() *apikey.Manager { return sm.ApiKeyMgr } -func (sm *SessionManager) CreateAPIKey(userID int64, name string, validityDays int) (*apikey.ApiKey, error) { +func (sm *Manager) CreateAPIKey(userID int64, name string, validityDays int) (*apikey.ApiKey, error) { return sm.ApiKeyMgr.CreateApiKey(userID, name, validityDays) } -// middleware that checks if a user is authenticated via cookies or API key -func WithAuth(handler http.HandlerFunc, sm *SessionManager) http.HandlerFunc { +// WithAuth middleware that checks if a user is authenticated via cookies or API key +func WithAuth(handler http.HandlerFunc, sm *Manager) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { // first: check API keys apiKeyStr, apiKeyErr := apikey.ExtractApiKey(r) @@ -232,8 +228,8 @@ func WithAuth(handler http.HandlerFunc, sm *SessionManager) http.HandlerFunc { } } -// middleware that checks if a user is authenticated but doesn't error out if not -func WithPossibleAuth(handler http.HandlerFunc, sm *SessionManager) http.HandlerFunc { +// WithPossibleAuth middleware that checks if a user is authenticated but doesn't error out if not +func WithPossibleAuth(handler http.HandlerFunc, sm *Manager) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() authenticated := false @@ -267,8 +263,8 @@ func WithPossibleAuth(handler http.HandlerFunc, sm *SessionManager) http.Handler } } -// middleware that only accepts API keys -func WithAPIAuth(handler http.HandlerFunc, sm *SessionManager) http.HandlerFunc { +// WithAPIAuth middleware that only accepts API keys +func WithAPIAuth(handler http.HandlerFunc, sm *Manager) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { apiKeyStr, apiKeyErr := apikey.ExtractApiKey(r) if apiKeyErr != nil || apiKeyStr == "" { @@ -294,7 +290,7 @@ func WithAPIAuth(handler http.HandlerFunc, sm *SessionManager) http.HandlerFunc } } -func (sm *SessionManager) HandleDebug(w http.ResponseWriter, r *http.Request) { +func (sm *Manager) HandleDebug(w http.ResponseWriter, r *http.Request) { ctx := r.Context() userID, ok := GetUserID(ctx) if !ok {