package repositories import ( "clapclap/internal/models" "clapclap/internal/utils" "clapclap/internal/uuid" "gorm.io/gorm" "gorm.io/gorm/clause" ) // It creates a temporary clone of RatingRepo that uses the transaction DB. func (r *Repo) RunInTx(fn func(txRepo *Repo) error) error { return r.db.Transaction(func(tx *gorm.DB) error { txRepo := &Repo{db: tx} return fn(txRepo) }) } // --- CRUD --- func (r *Repo) GetExistingRating(authorID, trackID uuid.UUID) (*models.Rating, error) { var rating models.Rating result := r.db.Where("author_id = ? AND track_id = ?", authorID, trackID).Limit(1).Find(&rating) if result.Error != nil { return nil, result.Error } if result.RowsAffected == 0 { return nil, nil } return &rating, nil } func (r *Repo) Upsert(rating *models.Rating) error { return r.db.Clauses(clause.OnConflict{ Columns: []clause.Column{{Name: "author_id"}, {Name: "track_id"}}, DoUpdates: clause.AssignmentColumns([]string{"value", "updated_at"}), }).Create(rating).Error } func (r *Repo) Delete(id uuid.UUID) (*models.Rating, error) { var rating models.Rating err := r.db.Clauses(clause.Returning{}).Where("id = ?", id).Delete(&rating).Error return &rating, err } func (r *Repo) GetTrackStats(trackID uuid.UUID) (float64, int, error) { var track models.Track // This forces other concurrent requests to wait in line until this transaction finishes! err := r.db.Clauses(clause.Locking{Strength: "UPDATE"}). Select("average_rating", "rating_count"). First(&track, "id = ?", trackID).Error return track.AverageRating, track.RatingCount, err } func (r *Repo) UpdateTrackStats(trackID uuid.UUID, newAvg float64, newCount int) error { // Using a map forces GORM to update the fields no matter what, // avoiding the silent struct-skipping bug. err := r.db.Model(&models.Track{}). Where("id = ?", trackID). Updates(map[string]any{ "average_rating": newAvg, "rating_count": newCount, }).Error if err != nil { utils.LogError(err) } return err } func (r *Repo) GetByTrackID(trackID string, preloads []string) ([]models.Rating, error) { var ratings []models.Rating err := r.db.Where("track_id = ?", trackID).Preload("Author").Find(&ratings).Error return ratings, err } func (r *Repo) GetByID(id string, preloads []string) (*models.Rating, error) { var rating models.Rating uid, err := uuid.Parse(id) if err != nil { return nil, err } query := r.db for _, p := range preloads { query = query.Preload(p) } err = query.First(&rating, "id = ?", uid).Error return &rating, err }