package queryupdate import ( "context" "errors" "fmt" "log/slog" queryversionsyncrunner "queryorchestration/api/queryVersionSyncRunner" "queryorchestration/internal/database" "queryorchestration/internal/database/repository" "queryorchestration/internal/query" resultprocessor "queryorchestration/internal/query/result/processor" contextfull "queryorchestration/internal/query/types/contextFull" jsonextractor "queryorchestration/internal/query/types/jsonExtractor" "queryorchestration/internal/serviceconfig/queue" "queryorchestration/internal/validation" "github.com/google/uuid" ) func (s *Service) Update(ctx context.Context, entity *resultprocessor.Update) error { current, err := s.svc.Query.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 } err = s.informUpdate(ctx, entity) if err != nil { return err } return nil } func (s *Service) informUpdate(ctx context.Context, entity *resultprocessor.Update) error { if entity.ActiveVersion == nil { return nil } return s.cfg.SendToQueue(ctx, &queue.SendParams{ QueueURL: s.cfg.GetQueryVersionSyncURL(), Body: queryversionsyncrunner.Body{ ID: entity.ID, }, }) } func (s *Service) normalizeUpdateRequiredQueryIDs(ctx context.Context, current *query.Query, entity query.RequiredQueryIDs) error { err := s.svc.Query.NormalizeQueryIDs(ctx, entity) if err != nil { return err } if entity.GetRequiredQueryIDs() == nil { return nil } createsloop, err := s.cfg.GetDBQueries().IsQueryInDependencyTree(ctx, &repository.IsQueryInDependencyTreeParams{ Queryid: database.MustToDBUUID(current.ID), Requiredqueryids: 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.Query, entity *resultprocessor.Update) error { err := s.svc.Query.NormalizeActiveVersion(current, entity) if err != nil { return err } err = s.normalizeUpdateRequiredQueryIDs(ctx, current, entity) if err != nil { return err } err = s.svc.Query.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, query.ParseQuery(current), entity) if err != nil { return err } return nil } func (s *Service) submitUpdate(ctx context.Context, current *query.Query, entity *resultprocessor.Update) error { err := s.cfg.ExecuteDBTransaction(ctx, func(ctx context.Context, q *repository.Queries) error { id := database.MustToDBUUID(entity.ID) activeName, err := validation.GetFieldName(entity, entity.ActiveVersion) if err != nil { return err } onlyactive := validation.AreAllPointersNilExcept(entity, activeName) var latestVersion int32 if !onlyactive { latestVersion, err = q.AddLatestQueryVersion(ctx, id) if err != nil { return err } } if entity.ActiveVersion != nil { err := q.AddActiveQueryVersion(ctx, &repository.AddActiveQueryVersionParams{ Queryid: id, Versionid: *entity.ActiveVersion, }) if err != nil { return err } } if entity.RequiredQueryIDs != nil { 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.SetQueryConfig(ctx, &repository.SetQueryConfigParams{ Queryid: id, Config: []byte(*entity.Config), Addedversion: latestVersion, }) 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") } }