Refactoring of API

- Add request context to all database calls
- Update deprecated gqlgen functions
- Update go.mod dependencies
This commit is contained in:
viktorstrate
2021-11-11 18:57:02 +01:00
parent 12085698c8
commit c5a307283d
15 changed files with 271 additions and 703 deletions

View File

@@ -0,0 +1,48 @@
package graphql_endpoint
import (
"time"
graphql_handler "github.com/99designs/gqlgen/graphql/handler"
"github.com/99designs/gqlgen/graphql/handler/extension"
"github.com/99designs/gqlgen/graphql/handler/lru"
"github.com/99designs/gqlgen/graphql/handler/transport"
photoview_graphql "github.com/photoview/photoview/api/graphql"
"github.com/photoview/photoview/api/graphql/resolvers"
"github.com/photoview/photoview/api/utils"
"gorm.io/gorm"
)
func GraphqlEndpoint(db *gorm.DB) *graphql_handler.Server {
graphqlResolver := resolvers.NewRootResolver(db)
graphqlDirective := photoview_graphql.DirectiveRoot{}
graphqlDirective.IsAdmin = photoview_graphql.IsAdmin
graphqlDirective.IsAuthorized = photoview_graphql.IsAuthorized
graphqlConfig := photoview_graphql.Config{
Resolvers: &graphqlResolver,
Directives: graphqlDirective,
}
graphqlServer := graphql_handler.New(photoview_graphql.NewExecutableSchema(graphqlConfig))
graphqlServer.AddTransport(transport.Websocket{
KeepAlivePingInterval: 10 * time.Second,
})
graphqlServer.AddTransport(transport.Options{})
graphqlServer.AddTransport(transport.GET{})
graphqlServer.AddTransport(transport.POST{})
graphqlServer.AddTransport(transport.MultipartForm{})
graphqlServer.SetQueryCache(lru.New(1000))
graphqlServer.Use(extension.Introspection{})
graphqlServer.Use(extension.AutomaticPersistedQuery{
Cache: lru.New(100),
})
if utils.DevelopmentMode() {
graphqlServer.Use(extension.Introspection{})
}
return graphqlServer
}

View File

@@ -17,10 +17,11 @@ func (r *queryResolver) MyAlbums(ctx context.Context, order *models.Ordering, pa
return nil, auth.ErrUnauthorized
}
return actions.MyAlbums(r.Database, user, order, paginate, onlyRoot, showEmpty, onlyWithFavorites)
return actions.MyAlbums(r.DB(ctx), user, order, paginate, onlyRoot, showEmpty, onlyWithFavorites)
}
func (r *queryResolver) Album(ctx context.Context, id int, tokenCredentials *models.ShareTokenCredentials) (*models.Album, error) {
db := r.DB(ctx)
if tokenCredentials != nil {
shareToken, err := r.ShareToken(ctx, *tokenCredentials)
@@ -33,7 +34,7 @@ func (r *queryResolver) Album(ctx context.Context, id int, tokenCredentials *mod
return shareToken.Album, nil
}
subAlbum, err := shareToken.Album.GetChildren(r.Database, func(query *gorm.DB) *gorm.DB { return query.Where("sub_albums.id = ?", id) })
subAlbum, err := shareToken.Album.GetChildren(db, func(query *gorm.DB) *gorm.DB { return query.Where("sub_albums.id = ?", id) })
if err != nil {
return nil, errors.Wrapf(err, "find sub album of share token (%s)", tokenCredentials.Token)
}
@@ -49,7 +50,7 @@ func (r *queryResolver) Album(ctx context.Context, id int, tokenCredentials *mod
return nil, auth.ErrUnauthorized
}
return actions.Album(r.Database, user, id)
return actions.Album(db, user, id)
}
func (r *Resolver) Album() api.AlbumResolver {
@@ -59,10 +60,11 @@ func (r *Resolver) Album() api.AlbumResolver {
type albumResolver struct{ *Resolver }
func (r *albumResolver) Media(ctx context.Context, album *models.Album, order *models.Ordering, paginate *models.Pagination, onlyFavorites *bool) ([]*models.Media, error) {
db := r.DB(ctx)
query := r.Database.
query := db.
Where("media.album_id = ?", album.ID).
Where("media.id IN (?)", r.Database.Model(&models.MediaURL{}).Select("media_urls.media_id").Where("media_urls.media_id = media.id"))
Where("media.id IN (?)", db.Model(&models.MediaURL{}).Select("media_urls.media_id").Where("media_urls.media_id = media.id"))
if onlyFavorites != nil && *onlyFavorites == true {
user := auth.UserFromContext(ctx)
@@ -70,7 +72,7 @@ func (r *albumResolver) Media(ctx context.Context, album *models.Album, order *m
return nil, errors.New("cannot get favorite media without being authorized")
}
favoriteQuery := r.Database.Model(&models.UserMediaData{
favoriteQuery := db.Model(&models.UserMediaData{
UserID: user.ID,
}).Where("user_media_data.media_id = media.id").Where("user_media_data.favorite = true")
@@ -88,14 +90,14 @@ func (r *albumResolver) Media(ctx context.Context, album *models.Album, order *m
}
func (r *albumResolver) Thumbnail(ctx context.Context, album *models.Album) (*models.Media, error) {
return album.Thumbnail(r.Database)
return album.Thumbnail(r.DB(ctx))
}
func (r *albumResolver) SubAlbums(ctx context.Context, parent *models.Album, order *models.Ordering, paginate *models.Pagination) ([]*models.Album, error) {
var albums []*models.Album
query := r.Database.Where("parent_album_id = ?", parent.ID)
query := r.DB(ctx).Where("parent_album_id = ?", parent.ID)
query = models.FormatSQL(query, order, paginate)
if err := query.Find(&albums).Error; err != nil {
@@ -116,7 +118,7 @@ func (r *albumResolver) Owner(ctx context.Context, obj *models.Album) (*models.U
func (r *albumResolver) Shares(ctx context.Context, album *models.Album) ([]*models.ShareToken, error) {
var shareTokens []*models.ShareToken
if err := r.Database.Where("album_id = ?", album.ID).Find(&shareTokens).Error; err != nil {
if err := r.DB(ctx).Where("album_id = ?", album.ID).Find(&shareTokens).Error; err != nil {
return nil, err
}
@@ -131,7 +133,7 @@ func (r *albumResolver) Path(ctx context.Context, obj *models.Album) ([]*models.
return empty, nil
}
return actions.AlbumPath(r.Database, user, obj)
return actions.AlbumPath(r.DB(ctx), user, obj)
}
// Takes album_id, resets album.cover_id to 0 (null)
@@ -141,7 +143,7 @@ func (r *mutationResolver) ResetAlbumCover(ctx context.Context, albumID int) (*m
return nil, errors.New("unauthorized")
}
return actions.ResetAlbumCover(r.Database, user, albumID)
return actions.ResetAlbumCover(r.DB(ctx), user, albumID)
}
func (r *mutationResolver) SetAlbumCover(ctx context.Context, mediaID int) (*models.Album, error) {
@@ -150,5 +152,5 @@ func (r *mutationResolver) SetAlbumCover(ctx context.Context, mediaID int) (*mod
return nil, errors.New("unauthorized")
}
return actions.SetAlbumCover(r.Database, user, mediaID)
return actions.SetAlbumCover(r.DB(ctx), user, mediaID)
}

View File

@@ -37,7 +37,7 @@ func (r imageFaceResolver) FaceGroup(ctx context.Context, obj *models.ImageFace)
}
var faceGroup models.FaceGroup
if err := r.Database.Model(&obj).Association("FaceGroup").Find(&faceGroup); err != nil {
if err := r.DB(ctx).Model(&obj).Association("FaceGroup").Find(&faceGroup); err != nil {
return nil, err
}
@@ -47,7 +47,7 @@ func (r imageFaceResolver) FaceGroup(ctx context.Context, obj *models.ImageFace)
}
func (r imageFaceResolver) Media(ctx context.Context, obj *models.ImageFace) (*models.Media, error) {
if err := obj.FillMedia(r.Database); err != nil {
if err := obj.FillMedia(r.DB(ctx)); err != nil {
return nil, err
}
@@ -55,6 +55,7 @@ func (r imageFaceResolver) Media(ctx context.Context, obj *models.ImageFace) (*m
}
func (r faceGroupResolver) ImageFaces(ctx context.Context, obj *models.FaceGroup, paginate *models.Pagination) ([]*models.ImageFace, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -64,7 +65,7 @@ func (r faceGroupResolver) ImageFaces(ctx context.Context, obj *models.FaceGroup
return nil, errors.New("face detector not initialized")
}
if err := user.FillAlbums(r.Database); err != nil {
if err := user.FillAlbums(db); err != nil {
return nil, err
}
@@ -73,7 +74,7 @@ func (r faceGroupResolver) ImageFaces(ctx context.Context, obj *models.FaceGroup
userAlbumIDs[i] = album.ID
}
query := r.Database.
query := db.
Joins("Media").
Where("face_group_id = ?", obj.ID).
Where("album_id IN (?)", userAlbumIDs)
@@ -89,6 +90,7 @@ func (r faceGroupResolver) ImageFaces(ctx context.Context, obj *models.FaceGroup
}
func (r faceGroupResolver) ImageFaceCount(ctx context.Context, obj *models.FaceGroup) (int, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return -1, errors.New("unauthorized")
@@ -98,7 +100,7 @@ func (r faceGroupResolver) ImageFaceCount(ctx context.Context, obj *models.FaceG
return -1, errors.New("face detector not initialized")
}
if err := user.FillAlbums(r.Database); err != nil {
if err := user.FillAlbums(db); err != nil {
return -1, err
}
@@ -107,7 +109,7 @@ func (r faceGroupResolver) ImageFaceCount(ctx context.Context, obj *models.FaceG
userAlbumIDs[i] = album.ID
}
query := r.Database.
query := db.
Model(&models.ImageFace{}).
Joins("Media").
Where("face_group_id = ?", obj.ID).
@@ -122,6 +124,7 @@ func (r faceGroupResolver) ImageFaceCount(ctx context.Context, obj *models.FaceG
}
func (r *queryResolver) FaceGroup(ctx context.Context, id int) (*models.FaceGroup, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -131,7 +134,7 @@ func (r *queryResolver) FaceGroup(ctx context.Context, id int) (*models.FaceGrou
return nil, errors.New("face detector not initialized")
}
if err := user.FillAlbums(r.Database); err != nil {
if err := user.FillAlbums(db); err != nil {
return nil, err
}
@@ -140,10 +143,10 @@ func (r *queryResolver) FaceGroup(ctx context.Context, id int) (*models.FaceGrou
userAlbumIDs[i] = album.ID
}
faceGroupQuery := r.Database.
faceGroupQuery := db.
Joins("LEFT JOIN image_faces ON image_faces.face_group_id = face_groups.id").
Where("face_groups.id = ?", id).
Where("image_faces.media_id IN (?)", r.Database.Select("media_id").Table("media").Where("media.album_id IN (?)", userAlbumIDs))
Where("image_faces.media_id IN (?)", db.Select("media_id").Table("media").Where("media.album_id IN (?)", userAlbumIDs))
var faceGroup models.FaceGroup
if err := faceGroupQuery.Find(&faceGroup).Error; err != nil {
@@ -154,6 +157,7 @@ func (r *queryResolver) FaceGroup(ctx context.Context, id int) (*models.FaceGrou
}
func (r *queryResolver) MyFaceGroups(ctx context.Context, paginate *models.Pagination) ([]*models.FaceGroup, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -163,7 +167,7 @@ func (r *queryResolver) MyFaceGroups(ctx context.Context, paginate *models.Pagin
return nil, errors.New("face detector not initialized")
}
if err := user.FillAlbums(r.Database); err != nil {
if err := user.FillAlbums(db); err != nil {
return nil, err
}
@@ -172,9 +176,9 @@ func (r *queryResolver) MyFaceGroups(ctx context.Context, paginate *models.Pagin
userAlbumIDs[i] = album.ID
}
faceGroupQuery := r.Database.
faceGroupQuery := db.
Joins("JOIN image_faces ON image_faces.face_group_id = face_groups.id").
Where("image_faces.media_id IN (?)", r.Database.Select("media.id").Table("media").Where("media.album_id IN (?)", userAlbumIDs)).
Where("image_faces.media_id IN (?)", db.Select("media.id").Table("media").Where("media.album_id IN (?)", userAlbumIDs)).
Group("image_faces.face_group_id").
Group("face_groups.id").
Order("CASE WHEN label IS NULL THEN 1 ELSE 0 END").
@@ -191,6 +195,7 @@ func (r *queryResolver) MyFaceGroups(ctx context.Context, paginate *models.Pagin
}
func (r *mutationResolver) SetFaceGroupLabel(ctx context.Context, faceGroupID int, label *string) (*models.FaceGroup, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -200,12 +205,12 @@ func (r *mutationResolver) SetFaceGroupLabel(ctx context.Context, faceGroupID in
return nil, errors.New("face detector not initialized")
}
faceGroup, err := userOwnedFaceGroup(r.Database, user, faceGroupID)
faceGroup, err := userOwnedFaceGroup(db, user, faceGroupID)
if err != nil {
return nil, err
}
if err := r.Database.Model(faceGroup).Update("label", label).Error; err != nil {
if err := db.Model(faceGroup).Update("label", label).Error; err != nil {
return nil, err
}
@@ -213,6 +218,7 @@ func (r *mutationResolver) SetFaceGroupLabel(ctx context.Context, faceGroupID in
}
func (r *mutationResolver) CombineFaceGroups(ctx context.Context, destinationFaceGroupID int, sourceFaceGroupID int) (*models.FaceGroup, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -222,17 +228,17 @@ func (r *mutationResolver) CombineFaceGroups(ctx context.Context, destinationFac
return nil, errors.New("face detector not initialized")
}
destinationFaceGroup, err := userOwnedFaceGroup(r.Database, user, destinationFaceGroupID)
destinationFaceGroup, err := userOwnedFaceGroup(db, user, destinationFaceGroupID)
if err != nil {
return nil, err
}
sourceFaceGroup, err := userOwnedFaceGroup(r.Database, user, sourceFaceGroupID)
sourceFaceGroup, err := userOwnedFaceGroup(db, user, sourceFaceGroupID)
if err != nil {
return nil, err
}
updateError := r.Database.Transaction(func(tx *gorm.DB) error {
updateError := db.Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&models.ImageFace{}).Where("face_group_id = ?", sourceFaceGroup.ID).Update("face_group_id", destinationFaceGroup.ID).Error; err != nil {
return err
}
@@ -254,6 +260,7 @@ func (r *mutationResolver) CombineFaceGroups(ctx context.Context, destinationFac
}
func (r *mutationResolver) MoveImageFaces(ctx context.Context, imageFaceIDs []int, destinationFaceGroupID int) (*models.FaceGroup, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -266,7 +273,7 @@ func (r *mutationResolver) MoveImageFaces(ctx context.Context, imageFaceIDs []in
userOwnedImageFaceIDs := make([]int, 0)
var destFaceGroup *models.FaceGroup
transErr := r.Database.Transaction(func(tx *gorm.DB) error {
transErr := db.Transaction(func(tx *gorm.DB) error {
var err error
destFaceGroup, err = userOwnedFaceGroup(tx, user, destinationFaceGroupID)
@@ -325,6 +332,7 @@ func (r *mutationResolver) MoveImageFaces(ctx context.Context, imageFaceIDs []in
}
func (r *mutationResolver) RecognizeUnlabeledFaces(ctx context.Context) ([]*models.ImageFace, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -336,7 +344,7 @@ func (r *mutationResolver) RecognizeUnlabeledFaces(ctx context.Context) ([]*mode
var updatedImageFaces []*models.ImageFace
transactionError := r.Database.Transaction(func(tx *gorm.DB) error {
transactionError := db.Transaction(func(tx *gorm.DB) error {
var err error
updatedImageFaces, err = face_detection.GlobalFaceDetector.RecognizeUnlabeledFaces(tx, user)
@@ -351,6 +359,7 @@ func (r *mutationResolver) RecognizeUnlabeledFaces(ctx context.Context) ([]*mode
}
func (r *mutationResolver) DetachImageFaces(ctx context.Context, imageFaceIDs []int) (*models.FaceGroup, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
@@ -363,7 +372,7 @@ func (r *mutationResolver) DetachImageFaces(ctx context.Context, imageFaceIDs []
userOwnedImageFaceIDs := make([]int, 0)
newFaceGroup := models.FaceGroup{}
transactionError := r.Database.Transaction(func(tx *gorm.DB) error {
transactionError := db.Transaction(func(tx *gorm.DB) error {
userOwnedImageFaces, err := getUserOwnedImageFaces(tx, user, imageFaceIDs)
if err != nil {

View File

@@ -19,10 +19,11 @@ func (r *queryResolver) MyMedia(ctx context.Context, order *models.Ordering, pag
return nil, errors.New("unauthorized")
}
return actions.MyMedia(r.Database, user, order, paginate)
return actions.MyMedia(r.DB(ctx), user, order, paginate)
}
func (r *queryResolver) Media(ctx context.Context, id int, tokenCredentials *models.ShareTokenCredentials) (*models.Media, error) {
db := r.DB(ctx)
if tokenCredentials != nil {
shareToken, err := r.ShareToken(ctx, *tokenCredentials)
@@ -42,11 +43,11 @@ func (r *queryResolver) Media(ctx context.Context, id int, tokenCredentials *mod
var media models.Media
err := r.Database.
err := db.
Joins("Album").
Where("media.id = ?", id).
Where("EXISTS (SELECT * FROM user_albums WHERE user_albums.album_id = media.album_id AND user_albums.user_id = ?)", user.ID).
Where("media.id IN (?)", r.Database.Model(&models.MediaURL{}).Select("media_id").Where("media_urls.media_id = media.id")).
Where("media.id IN (?)", db.Model(&models.MediaURL{}).Select("media_id").Where("media_urls.media_id = media.id")).
First(&media).Error
if err != nil {
@@ -57,6 +58,7 @@ func (r *queryResolver) Media(ctx context.Context, id int, tokenCredentials *mod
}
func (r *queryResolver) MediaList(ctx context.Context, ids []int) ([]*models.Media, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, auth.ErrUnauthorized
@@ -67,7 +69,7 @@ func (r *queryResolver) MediaList(ctx context.Context, ids []int) ([]*models.Med
}
var media []*models.Media
err := r.Database.Model(&media).
err := db.Model(&media).
Joins("LEFT JOIN user_albums ON user_albums.album_id = media.album_id").
Where("media.id IN ?", ids).
Where("user_albums.user_id = ?", user.ID).
@@ -95,7 +97,7 @@ func (r *mediaResolver) Type(ctx context.Context, media *models.Media) (models.M
func (r *mediaResolver) Album(ctx context.Context, obj *models.Media) (*models.Album, error) {
var album models.Album
err := r.Database.Find(&album, obj.AlbumID).Error
err := r.DB(ctx).Find(&album, obj.AlbumID).Error
if err != nil {
return nil, err
}
@@ -104,7 +106,7 @@ func (r *mediaResolver) Album(ctx context.Context, obj *models.Media) (*models.A
func (r *mediaResolver) Shares(ctx context.Context, media *models.Media) ([]*models.ShareToken, error) {
var shareTokens []*models.ShareToken
if err := r.Database.Where("media_id = ?", media.ID).Find(&shareTokens).Error; err != nil {
if err := r.DB(ctx).Where("media_id = ?", media.ID).Find(&shareTokens).Error; err != nil {
return nil, errors.Wrapf(err, "get shares for media (%s)", media.Path)
}
@@ -114,7 +116,7 @@ func (r *mediaResolver) Shares(ctx context.Context, media *models.Media) ([]*mod
func (r *mediaResolver) Downloads(ctx context.Context, media *models.Media) ([]*models.MediaDownload, error) {
var mediaUrls []*models.MediaURL
if err := r.Database.Where("media_id = ?", media.ID).Find(&mediaUrls).Error; err != nil {
if err := r.DB(ctx).Where("media_id = ?", media.ID).Find(&mediaUrls).Error; err != nil {
return nil, errors.Wrapf(err, "get downloads for media (%s)", media.Path)
}
@@ -171,7 +173,7 @@ func (r *mediaResolver) Exif(ctx context.Context, media *models.Media) (*models.
}
var exif models.MediaEXIF
if err := r.Database.Model(&media).Association("Exif").Find(&exif); err != nil {
if err := r.DB(ctx).Model(&media).Association("Exif").Find(&exif); err != nil {
return nil, err
}
@@ -197,7 +199,7 @@ func (r *mutationResolver) FavoriteMedia(ctx context.Context, mediaID int, favor
return nil, auth.ErrUnauthorized
}
return user.FavoriteMedia(r.Database, mediaID, favorite)
return user.FavoriteMedia(r.DB(ctx), mediaID, favorite)
}
func (r *mediaResolver) Faces(ctx context.Context, media *models.Media) ([]*models.ImageFace, error) {
@@ -210,7 +212,7 @@ func (r *mediaResolver) Faces(ctx context.Context, media *models.Media) ([]*mode
}
var faces []*models.ImageFace
if err := r.Database.Model(&media).Association("Faces").Find(&faces); err != nil {
if err := r.DB(ctx).Model(&media).Association("Faces").Find(&faces); err != nil {
return nil, err
}

View File

@@ -77,7 +77,7 @@ func (r *queryResolver) MyMediaGeoJSON(ctx context.Context) (interface{}, error)
var media []*geoMedia
err := r.Database.Table("media").
err := r.DB(ctx).Table("media").
Select(
"media.id AS media_id, media.title AS media_title, "+
"media_urls.media_name AS thumbnail_name, media_urls.width AS thumbnail_width, media_urls.height AS thumbnail_height, "+

View File

@@ -1,6 +1,8 @@
package resolvers
import (
"context"
api "github.com/photoview/photoview/api/graphql"
"gorm.io/gorm"
)
@@ -8,7 +10,18 @@ import (
//go:generate go run github.com/99designs/gqlgen
type Resolver struct {
Database *gorm.DB
database *gorm.DB
}
func NewRootResolver(db *gorm.DB) Resolver {
return Resolver{
database: db,
}
}
// DB returns a database instance that is tied to the given context
func (r *Resolver) DB(ctx context.Context) *gorm.DB {
return r.database.WithContext(ctx)
}
func (r *Resolver) Mutation() api.MutationResolver {

View File

@@ -27,9 +27,8 @@ func (r *mutationResolver) ScanAll(ctx context.Context) (*models.ScannerResult,
}
func (r *mutationResolver) ScanUser(ctx context.Context, userID int) (*models.ScannerResult, error) {
var user models.User
if err := r.Database.First(&user, userID).Error; err != nil {
if err := r.DB(ctx).First(&user, userID).Error; err != nil {
return nil, errors.Wrap(err, "get user from database")
}
@@ -44,16 +43,17 @@ func (r *mutationResolver) ScanUser(ctx context.Context, userID int) (*models.Sc
}
func (r *mutationResolver) SetPeriodicScanInterval(ctx context.Context, interval int) (int, error) {
db := r.DB(ctx)
if interval < 0 {
return 0, errors.New("interval must be 0 or above")
}
if err := r.Database.Session(&gorm.Session{AllowGlobalUpdate: true}).Model(&models.SiteInfo{}).Update("periodic_scan_interval", interval).Error; err != nil {
if err := db.Session(&gorm.Session{AllowGlobalUpdate: true}).Model(&models.SiteInfo{}).Update("periodic_scan_interval", interval).Error; err != nil {
return 0, err
}
var siteInfo models.SiteInfo
if err := r.Database.First(&siteInfo).Error; err != nil {
if err := db.First(&siteInfo).Error; err != nil {
return 0, err
}
@@ -63,6 +63,7 @@ func (r *mutationResolver) SetPeriodicScanInterval(ctx context.Context, interval
}
func (r *mutationResolver) SetScannerConcurrentWorkers(ctx context.Context, workers int) (int, error) {
db := r.DB(ctx)
if workers < 1 {
return 0, errors.New("concurrent workers must at least be 1")
}
@@ -71,12 +72,12 @@ func (r *mutationResolver) SetScannerConcurrentWorkers(ctx context.Context, work
return 0, errors.New("multiple workers not supported for SQLite databases")
}
if err := r.Database.Session(&gorm.Session{AllowGlobalUpdate: true}).Model(&models.SiteInfo{}).Update("concurrent_workers", workers).Error; err != nil {
if err := db.Session(&gorm.Session{AllowGlobalUpdate: true}).Model(&models.SiteInfo{}).Update("concurrent_workers", workers).Error; err != nil {
return 0, err
}
var siteInfo models.SiteInfo
if err := r.Database.First(&siteInfo).Error; err != nil {
if err := db.First(&siteInfo).Error; err != nil {
return 0, err
}

View File

@@ -15,5 +15,5 @@ func (r *Resolver) Search(ctx context.Context, query string, limitMedia *int, li
return nil, auth.ErrUnauthorized
}
return actions.Search(r.Database, query, user.ID, limitMedia, limitAlbums)
return actions.Search(r.DB(ctx), query, user.ID, limitMedia, limitAlbums)
}

View File

@@ -43,7 +43,7 @@ func (r *shareTokenResolver) HasPassword(ctx context.Context, obj *models.ShareT
func (r *queryResolver) ShareToken(ctx context.Context, credentials models.ShareTokenCredentials) (*models.ShareToken, error) {
var token models.ShareToken
if err := r.Database.Preload(clause.Associations).Where("value = ?", credentials.Token).First(&token).Error; err != nil {
if err := r.DB(ctx).Preload(clause.Associations).Where("value = ?", credentials.Token).First(&token).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("share not found")
} else {
@@ -66,7 +66,7 @@ func (r *queryResolver) ShareToken(ctx context.Context, credentials models.Share
func (r *queryResolver) ShareTokenValidatePassword(ctx context.Context, credentials models.ShareTokenCredentials) (bool, error) {
var token models.ShareToken
if err := r.Database.Where("value = ?", credentials.Token).First(&token).Error; err != nil {
if err := r.DB(ctx).Where("value = ?", credentials.Token).First(&token).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, errors.New("share not found")
} else {
@@ -99,7 +99,7 @@ func (r *mutationResolver) ShareAlbum(ctx context.Context, albumID int, expire *
return nil, auth.ErrUnauthorized
}
return actions.AddAlbumShare(r.Database, user, albumID, expire, password)
return actions.AddAlbumShare(r.DB(ctx), user, albumID, expire, password)
}
func (r *mutationResolver) ShareMedia(ctx context.Context, mediaID int, expire *time.Time, password *string) (*models.ShareToken, error) {
@@ -108,7 +108,7 @@ func (r *mutationResolver) ShareMedia(ctx context.Context, mediaID int, expire *
return nil, auth.ErrUnauthorized
}
return actions.AddMediaShare(r.Database, user, mediaID, expire, password)
return actions.AddMediaShare(r.DB(ctx), user, mediaID, expire, password)
}
func (r *mutationResolver) DeleteShareToken(ctx context.Context, tokenValue string) (*models.ShareToken, error) {
@@ -117,7 +117,7 @@ func (r *mutationResolver) DeleteShareToken(ctx context.Context, tokenValue stri
return nil, auth.ErrUnauthorized
}
return actions.DeleteShareToken(r.Database, user.ID, tokenValue)
return actions.DeleteShareToken(r.DB(ctx), user.ID, tokenValue)
}
func (r *mutationResolver) ProtectShareToken(ctx context.Context, tokenValue string, password *string) (*models.ShareToken, error) {
@@ -126,5 +126,5 @@ func (r *mutationResolver) ProtectShareToken(ctx context.Context, tokenValue str
return nil, auth.ErrUnauthorized
}
return actions.ProtectShareToken(r.Database, user.ID, tokenValue, password)
return actions.ProtectShareToken(r.DB(ctx), user.ID, tokenValue, password)
}

View File

@@ -9,7 +9,7 @@ import (
)
func (r *queryResolver) SiteInfo(ctx context.Context) (*models.SiteInfo, error) {
return models.GetSiteInfo(r.Database)
return models.GetSiteInfo(r.DB(ctx))
}
type SiteInfoResolver struct {

View File

@@ -15,5 +15,5 @@ func (r *queryResolver) MyTimeline(ctx context.Context, paginate *models.Paginat
return nil, auth.ErrUnauthorized
}
return actions.MyTimeline(r.Database, user, paginate, onlyFavorites, fromDate)
return actions.MyTimeline(r.DB(ctx), user, paginate, onlyFavorites, fromDate)
}

View File

@@ -30,7 +30,7 @@ func (r *queryResolver) User(ctx context.Context, order *models.Ordering, pagina
var users []*models.User
if err := models.FormatSQL(r.Database.Model(models.User{}), order, paginate).Find(&users).Error; err != nil {
if err := models.FormatSQL(r.DB(ctx).Model(models.User{}), order, paginate).Find(&users).Error; err != nil {
return nil, err
}
@@ -38,7 +38,7 @@ func (r *queryResolver) User(ctx context.Context, order *models.Ordering, pagina
}
func (r *userResolver) Albums(ctx context.Context, user *models.User) ([]*models.Album, error) {
user.FillAlbums(r.Database)
user.FillAlbums(r.DB(ctx))
pointerAlbums := make([]*models.Album, len(user.Albums))
for i, album := range user.Albums {
@@ -49,10 +49,11 @@ func (r *userResolver) Albums(ctx context.Context, user *models.User) ([]*models
}
func (r *userResolver) RootAlbums(ctx context.Context, user *models.User) (albums []*models.Album, err error) {
db := r.DB(ctx)
err = r.Database.Model(&user).
err = db.Model(&user).
Where("albums.parent_album_id NOT IN (?)",
r.Database.Table("user_albums").
db.Table("user_albums").
Select("albums.id").
Joins("JOIN albums ON albums.id = user_albums.album_id AND user_albums.user_id = ?", user.ID),
).Or("albums.parent_album_id IS NULL").
@@ -72,7 +73,8 @@ func (r *queryResolver) MyUser(ctx context.Context) (*models.User, error) {
}
func (r *mutationResolver) AuthorizeUser(ctx context.Context, username string, password string) (*models.AuthorizeResult, error) {
user, err := models.AuthorizeUser(r.Database, username, password)
db := r.DB(ctx)
user, err := models.AuthorizeUser(db, username, password)
if err != nil {
return &models.AuthorizeResult{
Success: false,
@@ -82,7 +84,7 @@ func (r *mutationResolver) AuthorizeUser(ctx context.Context, username string, p
var token *models.AccessToken
transactionError := r.Database.Transaction(func(tx *gorm.DB) error {
transactionError := db.Transaction(func(tx *gorm.DB) error {
token, err = user.GenerateAccessToken(tx)
if err != nil {
return err
@@ -103,7 +105,8 @@ func (r *mutationResolver) AuthorizeUser(ctx context.Context, username string, p
}
func (r *mutationResolver) InitialSetupWizard(ctx context.Context, username string, password string, rootPath string) (*models.AuthorizeResult, error) {
siteInfo, err := models.GetSiteInfo(r.Database)
db := r.DB(ctx)
siteInfo, err := models.GetSiteInfo(db)
if err != nil {
return nil, err
}
@@ -116,7 +119,7 @@ func (r *mutationResolver) InitialSetupWizard(ctx context.Context, username stri
var token *models.AccessToken
transactionError := r.Database.Transaction(func(tx *gorm.DB) error {
transactionError := db.Transaction(func(tx *gorm.DB) error {
if err := tx.Exec("UPDATE site_info SET initial_setup = false").Error; err != nil {
return err
}
@@ -162,7 +165,7 @@ func (r *queryResolver) MyUserPreferences(ctx context.Context) (*models.UserPref
userPref := models.UserPreferences{
UserID: user.ID,
}
if err := r.Database.Where("user_id = ?", user.ID).FirstOrCreate(&userPref).Error; err != nil {
if err := r.DB(ctx).Where("user_id = ?", user.ID).FirstOrCreate(&userPref).Error; err != nil {
return nil, err
}
@@ -170,6 +173,7 @@ func (r *queryResolver) MyUserPreferences(ctx context.Context) (*models.UserPref
}
func (r *mutationResolver) ChangeUserPreferences(ctx context.Context, language *string) (*models.UserPreferences, error) {
db := r.DB(ctx)
user := auth.UserFromContext(ctx)
if user == nil {
return nil, auth.ErrUnauthorized
@@ -182,14 +186,14 @@ func (r *mutationResolver) ChangeUserPreferences(ctx context.Context, language *
}
var userPref models.UserPreferences
if err := r.Database.Where("user_id = ?", user.ID).FirstOrInit(&userPref).Error; err != nil {
if err := db.Where("user_id = ?", user.ID).FirstOrInit(&userPref).Error; err != nil {
return nil, err
}
userPref.UserID = user.ID
userPref.Language = langTrans
if err := r.Database.Save(&userPref).Error; err != nil {
if err := db.Save(&userPref).Error; err != nil {
return nil, err
}
@@ -198,13 +202,14 @@ func (r *mutationResolver) ChangeUserPreferences(ctx context.Context, language *
// Admin queries
func (r *mutationResolver) UpdateUser(ctx context.Context, id int, username *string, password *string, admin *bool) (*models.User, error) {
db := r.DB(ctx)
if username == nil && password == nil && admin == nil {
return nil, errors.New("no updates requested")
}
var user models.User
if err := r.Database.First(&user, id).Error; err != nil {
if err := db.First(&user, id).Error; err != nil {
return nil, err
}
@@ -226,7 +231,7 @@ func (r *mutationResolver) UpdateUser(ctx context.Context, id int, username *str
user.Admin = *admin
}
if err := r.Database.Save(&user).Error; err != nil {
if err := db.Save(&user).Error; err != nil {
return nil, errors.Wrap(err, "failed to update user")
}
@@ -237,7 +242,7 @@ func (r *mutationResolver) CreateUser(ctx context.Context, username string, pass
var user *models.User
transactionError := r.Database.Transaction(func(tx *gorm.DB) error {
transactionError := r.DB(ctx).Transaction(func(tx *gorm.DB) error {
var err error
user, err = models.RegisterUser(tx, username, password, admin)
if err != nil {
@@ -255,19 +260,20 @@ func (r *mutationResolver) CreateUser(ctx context.Context, username string, pass
}
func (r *mutationResolver) DeleteUser(ctx context.Context, id int) (*models.User, error) {
return actions.DeleteUser(r.Database, id)
return actions.DeleteUser(r.DB(ctx), id)
}
func (r *mutationResolver) UserAddRootPath(ctx context.Context, id int, rootPath string) (*models.Album, error) {
db := r.DB(ctx)
rootPath = path.Clean(rootPath)
var user models.User
if err := r.Database.First(&user, id).Error; err != nil {
if err := db.First(&user, id).Error; err != nil {
return nil, err
}
newAlbum, err := scanner.NewRootAlbum(r.Database, rootPath, &user)
newAlbum, err := scanner.NewRootAlbum(db, rootPath, &user)
if err != nil {
return nil, err
}
@@ -276,15 +282,16 @@ func (r *mutationResolver) UserAddRootPath(ctx context.Context, id int, rootPath
}
func (r *mutationResolver) UserRemoveRootAlbum(ctx context.Context, userID int, albumID int) (*models.Album, error) {
db := r.DB(ctx)
var album models.Album
if err := r.Database.First(&album, albumID).Error; err != nil {
if err := db.First(&album, albumID).Error; err != nil {
return nil, err
}
var deletedAlbumIDs []int = nil
transactionError := r.Database.Transaction(func(tx *gorm.DB) error {
transactionError := db.Transaction(func(tx *gorm.DB) error {
if err := tx.Raw("DELETE FROM user_albums WHERE user_id = ? AND album_id = ?", userID, albumID).Error; err != nil {
return err
}
@@ -344,7 +351,7 @@ func (r *mutationResolver) UserRemoveRootAlbum(ctx context.Context, userID int,
// Reload faces as media might have been deleted
if face_detection.GlobalFaceDetector != nil {
if err := face_detection.GlobalFaceDetector.ReloadFacesFromDatabase(r.Database); err != nil {
if err := face_detection.GlobalFaceDetector.ReloadFacesFromDatabase(db); err != nil {
return nil, err
}
}