package query import ( "context" "errors" "fmt" "log/slog" "queryorchestration/internal/database" "queryorchestration/internal/database/repository" resultprocessor "queryorchestration/internal/query/result/processor" contextfull "queryorchestration/internal/query/types/contextFull" jsonextractor "queryorchestration/internal/query/types/jsonExtractor" "queryorchestration/internal/validation" "github.com/google/uuid" ) func (s *Service) Update(ctx context.Context, entity *resultprocessor.Update) error { current, err := s.Get(ctx, entity.ID) if err != nil { return err } err = s.normalizeUpdate(ctx, current, entity) if err != nil { return err } err = s.submitUpdate(ctx, current, entity) if err != nil { return err } return nil } func (s *Service) normalizeUpdateRequiredQueryIDs(ctx context.Context, current *Query, entity RequiredQueryIDs) error { err := s.NormalizeQueryIDs(ctx, entity) if err != nil { return err } if entity.GetRequiredQueryIDs() == nil { return nil } createsloop, err := s.cfg.GetDBQueries().IsQueryInDependencyTree(ctx, &repository.IsQueryInDependencyTreeParams{ Requiredqueryid: database.MustToDBUUID(current.ID), ID: database.MustToDBUUIDArray(*entity.GetRequiredQueryIDs()), }) if err != nil { return err } else if createsloop { return errors.New("required ids create a loop") } addIDs := getSetDifference(entity.GetRequiredQueryIDs(), current.RequiredQueryIDs) removeIDs := getSetDifference(current.RequiredQueryIDs, entity.GetRequiredQueryIDs()) if len(addIDs) == 0 && len(removeIDs) == 0 { entity.SetRequiredQueryIDs(nil) } return nil } func (s *Service) normalizeUpdate(ctx context.Context, current *Query, entity *resultprocessor.Update) error { err := s.normalizeActiveVersion(current, entity) if err != nil { return err } err = s.normalizeUpdateRequiredQueryIDs(ctx, current, entity) if err != nil { return err } err = s.normalizeConfig(entity) if err != nil { return err } activeVersionName, err := validation.GetFieldName(entity, entity.ActiveVersion) if err != nil { return err } if (entity.ActiveVersion == nil || *entity.ActiveVersion == current.LatestVersion+1) && validation.AreAllPointersNilExcept(entity, activeVersionName) { return errors.New("no changes") } validator, err := s.getUpdator(current.Type) if err != nil { return err } err = validator.Validate(ctx, ParseQuery(current), entity) if err != nil { return err } return nil } func (s *Service) submitUpdate(ctx context.Context, current *Query, entity *resultprocessor.Update) error { err := s.cfg.ExecuteDBTransaction(ctx, func(ctx context.Context, q *repository.Queries) error { latestVersion := current.LatestVersion + 1 id := database.MustToDBUUID(entity.ID) addIDs := getSetDifference(entity.RequiredQueryIDs, current.RequiredQueryIDs) for _, qID := range addIDs { err := q.AddRequiredQuery(ctx, &repository.AddRequiredQueryParams{ Queryid: id, Requiredqueryid: database.MustToDBUUID(qID), Addedversion: latestVersion, }) if err != nil { return err } } removeIDs := getSetDifference(current.RequiredQueryIDs, entity.RequiredQueryIDs) for _, qID := range removeIDs { err := q.RemoveRequiredQuery(ctx, &repository.RemoveRequiredQueryParams{ Queryid: id, Requiredqueryid: database.MustToDBUUID(qID), Removedversion: &latestVersion, }) if err != nil { return err } } if entity.Config != nil && *entity.Config != "" { err := q.RemoveQueryConfig(ctx, &repository.RemoveQueryConfigParams{ Queryid: id, Removedversion: &latestVersion, }) if err != nil { return err } err = q.AddQueryConfig(ctx, &repository.AddQueryConfigParams{ Queryid: id, Config: []byte(*entity.Config), Addedversion: latestVersion, }) if err != nil { return err } } activeVersion := entity.ActiveVersion if activeVersion == nil { activeVersion = ¤t.ActiveVersion } err := q.UpdateQuery(ctx, &repository.UpdateQueryParams{ Latestversion: latestVersion, Activeversion: *activeVersion, ID: id, }) if err != nil { return err } slog.Debug("query updated", "update", *entity) return nil }) if err != nil { return err } return nil } func getSetDifference(setA *[]uuid.UUID, setB *[]uuid.UUID) []uuid.UUID { if setA == nil { return []uuid.UUID{} } else if setB == nil { return *setA } diff := []uuid.UUID{} for _, q := range *setA { isFound := false for _, eq := range *setB { if q == eq { isFound = true break } } if !isFound { diff = append(diff, q) } } return diff } func (s *Service) getUpdator(qType resultprocessor.Type) (resultprocessor.Updator, error) { switch qType { case resultprocessor.TypeJsonExtractor: return jsonextractor.NewUpdator(), nil case resultprocessor.TypeContextFull: return contextfull.NewUpdator(), nil default: return nil, fmt.Errorf("attempting to process invalid query type") } } func ParseQuery(q *Query) *resultprocessor.Query { return &resultprocessor.Query{ ID: q.ID, Type: q.Type, Version: q.ActiveVersion, RequiredQueryIDs: q.RequiredQueryIDs, Config: q.Config, } }