From 3eb3435f08f43caaa7bf636cb622d632211351d6 Mon Sep 17 00:00:00 2001 From: viktorstrate Date: Mon, 22 Feb 2021 18:14:31 +0100 Subject: [PATCH] Add DetachImageFaces resolver --- api/graphql/generated.go | 78 ++++++++++++++++++++++++++++++++ api/graphql/resolvers/faces.go | 82 +++++++++++++++++++++++++++------- api/graphql/schema.graphql | 2 + 3 files changed, 147 insertions(+), 15 deletions(-) diff --git a/api/graphql/generated.go b/api/graphql/generated.go index f9ec1221..00734b52 100644 --- a/api/graphql/generated.go +++ b/api/graphql/generated.go @@ -142,6 +142,7 @@ type ComplexityRoot struct { CreateUser func(childComplexity int, username string, password *string, admin bool) int DeleteShareToken func(childComplexity int, token string) int DeleteUser func(childComplexity int, id int) int + DetachImageFaces func(childComplexity int, imageFaceIDs []int) int FavoriteMedia func(childComplexity int, mediaID int, favorite bool) int InitialSetupWizard func(childComplexity int, username string, password string, rootPath string) int MoveImageFaces func(childComplexity int, imageFaceIDs []int, destinationFaceGroupID int) int @@ -297,6 +298,7 @@ type MutationResolver interface { CombineFaceGroups(ctx context.Context, destinationFaceGroupID int, sourceFaceGroupID int) (*models.FaceGroup, error) MoveImageFaces(ctx context.Context, imageFaceIDs []int, destinationFaceGroupID int) (*models.FaceGroup, error) RecognizeUnlabeledFaces(ctx context.Context) ([]*models.ImageFace, error) + DetachImageFaces(ctx context.Context, imageFaceIDs []int) (*models.FaceGroup, error) } type QueryResolver interface { SiteInfo(ctx context.Context) (*models.SiteInfo, error) @@ -803,6 +805,18 @@ func (e *executableSchema) Complexity(typeName, field string, childComplexity in return e.complexity.Mutation.DeleteUser(childComplexity, args["id"].(int)), true + case "Mutation.detachImageFaces": + if e.complexity.Mutation.DetachImageFaces == nil { + break + } + + args, err := ec.field_Mutation_detachImageFaces_args(context.TODO(), rawArgs) + if err != nil { + return 0, false + } + + return e.complexity.Mutation.DetachImageFaces(childComplexity, args["imageFaceIDs"].([]int)), true + case "Mutation.favoriteMedia": if e.complexity.Mutation.FavoriteMedia == nil { break @@ -1655,6 +1669,8 @@ type Mutation { moveImageFaces(imageFaceIDs: [ID!]!, destinationFaceGroupID: ID!): FaceGroup! "Check all unlabeled faces to see if they match a labeled FaceGroup, and move them if they match" recognizeUnlabeledFaces: [ImageFace!]! + "Move a list of ImageFaces to a new face group" + detachImageFaces(imageFaceIDs: [ID!]!): FaceGroup! } type Subscription { @@ -2054,6 +2070,21 @@ func (ec *executionContext) field_Mutation_deleteUser_args(ctx context.Context, return args, nil } +func (ec *executionContext) field_Mutation_detachImageFaces_args(ctx context.Context, rawArgs map[string]interface{}) (map[string]interface{}, error) { + var err error + args := map[string]interface{}{} + var arg0 []int + if tmp, ok := rawArgs["imageFaceIDs"]; ok { + ctx := graphql.WithPathContext(ctx, graphql.NewPathWithField("imageFaceIDs")) + arg0, err = ec.unmarshalNID2ᚕintᚄ(ctx, tmp) + if err != nil { + return nil, err + } + } + args["imageFaceIDs"] = arg0 + return args, nil +} + func (ec *executionContext) field_Mutation_favoriteMedia_args(ctx context.Context, rawArgs map[string]interface{}) (map[string]interface{}, error) { var err error args := map[string]interface{}{} @@ -5525,6 +5556,48 @@ func (ec *executionContext) _Mutation_recognizeUnlabeledFaces(ctx context.Contex return ec.marshalNImageFace2ᚕᚖgithubᚗcomᚋphotoviewᚋphotoviewᚋapiᚋgraphqlᚋmodelsᚐImageFaceᚄ(ctx, field.Selections, res) } +func (ec *executionContext) _Mutation_detachImageFaces(ctx context.Context, field graphql.CollectedField) (ret graphql.Marshaler) { + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + ret = graphql.Null + } + }() + fc := &graphql.FieldContext{ + Object: "Mutation", + Field: field, + Args: nil, + IsMethod: true, + IsResolver: true, + } + + ctx = graphql.WithFieldContext(ctx, fc) + rawArgs := field.ArgumentMap(ec.Variables) + args, err := ec.field_Mutation_detachImageFaces_args(ctx, rawArgs) + if err != nil { + ec.Error(ctx, err) + return graphql.Null + } + fc.Args = args + resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (interface{}, error) { + ctx = rctx // use context from middleware stack in children + return ec.resolvers.Mutation().DetachImageFaces(rctx, args["imageFaceIDs"].([]int)) + }) + if err != nil { + ec.Error(ctx, err) + return graphql.Null + } + if resTmp == nil { + if !graphql.HasFieldError(ctx, fc) { + ec.Errorf(ctx, "must not be null") + } + return graphql.Null + } + res := resTmp.(*models.FaceGroup) + fc.Result = res + return ec.marshalNFaceGroup2ᚖgithubᚗcomᚋphotoviewᚋphotoviewᚋapiᚋgraphqlᚋmodelsᚐFaceGroup(ctx, field.Selections, res) +} + func (ec *executionContext) _Notification_key(ctx context.Context, field graphql.CollectedField, obj *models.Notification) (ret graphql.Marshaler) { defer func() { if r := recover(); r != nil { @@ -9627,6 +9700,11 @@ func (ec *executionContext) _Mutation(ctx context.Context, sel ast.SelectionSet) if out.Values[i] == graphql.Null { invalids++ } + case "detachImageFaces": + out.Values[i] = ec._Mutation_detachImageFaces(ctx, field) + if out.Values[i] == graphql.Null { + invalids++ + } default: panic("unknown field " + strconv.Quote(field.Name)) } diff --git a/api/graphql/resolvers/faces.go b/api/graphql/resolvers/faces.go index f863b514..2ed1fd7d 100644 --- a/api/graphql/resolvers/faces.go +++ b/api/graphql/resolvers/faces.go @@ -155,15 +155,6 @@ func (r *mutationResolver) MoveImageFaces(ctx context.Context, imageFaceIDs []in return nil, errors.New("unauthorized") } - if err := user.FillAlbums(r.Database); err != nil { - return nil, err - } - - userAlbumIDs := make([]int, len(user.Albums)) - for i, album := range user.Albums { - userAlbumIDs[i] = album.ID - } - userOwnedImageFaceIDs := make([]int, 0) var destFaceGroup *models.FaceGroup @@ -175,12 +166,8 @@ func (r *mutationResolver) MoveImageFaces(ctx context.Context, imageFaceIDs []in return err } - var userOwnedImageFaces []*models.ImageFace - if err := tx. - Joins("JOIN media ON media.id = image_faces.media_id"). - Where("media.album_id IN (?)", userAlbumIDs). - Where("image_faces.id IN (?)", imageFaceIDs). - Find(&userOwnedImageFaces).Error; err != nil { + userOwnedImageFaces, err := getUserOwnedImageFaces(tx, user, imageFaceIDs) + if err != nil { return err } @@ -251,6 +238,49 @@ func (r *mutationResolver) RecognizeUnlabeledFaces(ctx context.Context) ([]*mode return updatedImageFaces, nil } +func (r *mutationResolver) DetachImageFaces(ctx context.Context, imageFaceIDs []int) (*models.FaceGroup, error) { + user := auth.UserFromContext(ctx) + if user == nil { + return nil, errors.New("unauthorized") + } + + userOwnedImageFaceIDs := make([]int, 0) + newFaceGroup := models.FaceGroup{} + + transactionError := r.Database.Transaction(func(tx *gorm.DB) error { + + userOwnedImageFaces, err := getUserOwnedImageFaces(tx, user, imageFaceIDs) + if err != nil { + return err + } + + for _, imageFace := range userOwnedImageFaces { + userOwnedImageFaceIDs = append(userOwnedImageFaceIDs, imageFace.ID) + } + + if err := tx.Save(&newFaceGroup).Error; err != nil { + return err + } + + if err := tx. + Model(&models.ImageFace{}). + Where("id IN (?)", userOwnedImageFaceIDs). + Update("face_group_id", newFaceGroup.ID).Error; err != nil { + return err + } + + return nil + }) + + if transactionError != nil { + return nil, transactionError + } + + face_detection.GlobalFaceDetector.MergeImageFaces(userOwnedImageFaceIDs, int32(newFaceGroup.ID)) + + return &newFaceGroup, nil +} + func userOwnedFaceGroup(db *gorm.DB, user *models.User, faceGroupID int) (*models.FaceGroup, error) { if user.Admin { var faceGroup models.FaceGroup @@ -293,3 +323,25 @@ func userOwnedFaceGroup(db *gorm.DB, user *models.User, faceGroupID int) (*model return &faceGroup, nil } + +func getUserOwnedImageFaces(tx *gorm.DB, user *models.User, imageFaceIDs []int) ([]*models.ImageFace, error) { + if err := user.FillAlbums(tx); err != nil { + return nil, err + } + + userAlbumIDs := make([]int, len(user.Albums)) + for i, album := range user.Albums { + userAlbumIDs[i] = album.ID + } + + var userOwnedImageFaces []*models.ImageFace + if err := tx. + Joins("JOIN media ON media.id = image_faces.media_id"). + Where("media.album_id IN (?)", userAlbumIDs). + Where("image_faces.id IN (?)", imageFaceIDs). + Find(&userOwnedImageFaces).Error; err != nil { + return nil, err + } + + return userOwnedImageFaces, nil +} diff --git a/api/graphql/schema.graphql b/api/graphql/schema.graphql index ffcac478..15de4c0a 100644 --- a/api/graphql/schema.graphql +++ b/api/graphql/schema.graphql @@ -124,6 +124,8 @@ type Mutation { moveImageFaces(imageFaceIDs: [ID!]!, destinationFaceGroupID: ID!): FaceGroup! "Check all unlabeled faces to see if they match a labeled FaceGroup, and move them if they match" recognizeUnlabeledFaces: [ImageFace!]! + "Move a list of ImageFaces to a new face group" + detachImageFaces(imageFaceIDs: [ID!]!): FaceGroup! } type Subscription {