package collectorset import ( "context" "errors" "log/slog" clientsyncrunner "queryorchestration/api/clientSyncRunner" "queryorchestration/internal/collector" "queryorchestration/internal/database" "queryorchestration/internal/database/repository" "queryorchestration/internal/serviceconfig/queue" "queryorchestration/internal/validation" "github.com/google/uuid" "github.com/jackc/pgx/v5/pgtype" ) type SetParams struct { ClientID uuid.UUID ActiveVersion *int32 MinCleanVersion *int32 MinTextVersion *int32 Fields *map[string]uuid.UUID } func (s *Service) SetByClientId(ctx context.Context, params *SetParams) error { current, err := s.svc.Collector.GetByClientID(ctx, params.ClientID) if err != nil { return err } dbparams, err := s.getSetParams(ctx, current, params) if err != nil { return err } err = s.submitSet(ctx, current, dbparams) if err != nil { return err } err = s.informSet(ctx, dbparams) if err != nil { return err } return nil } func (s *Service) informSet(ctx context.Context, update *dbSetParams) error { if update.ActiveVersion == nil { return nil } return s.cfg.SendToQueue(ctx, &queue.SendParams{ QueueURL: s.cfg.GetClientSyncURL(), Body: clientsyncrunner.Body{ ID: database.MustToUUID(update.ClientID), }, }) } type dbSetParams struct { ClientID pgtype.UUID ActiveVersion *int32 MinCleanVersion *int32 MinTextVersion *int32 Fields *map[string]pgtype.UUID } func (s *Service) getSetParams(ctx context.Context, current *collector.Collector, params *SetParams) (*dbSetParams, error) { err := s.normalizeCodeVersions(current, params) if err != nil { return nil, err } fields, err := s.normalizeSetFieldsToDB(ctx, current.Fields, params.Fields) if err != nil { return nil, err } err = s.normalizeActiveVersion(current, params) if err != nil { return nil, err } activeVersionName, err := validation.GetFieldName(params, params.ActiveVersion) if err != nil { return nil, err } fieldsName, err := validation.GetFieldName(params, params.Fields) if err != nil { return nil, err } if (params.ActiveVersion == nil || *params.ActiveVersion == current.LatestVersion+1) && validation.AreAllPointersNilExcept(params, activeVersionName, fieldsName) && fields == nil { return nil, errors.New("no changes") } return &dbSetParams{ ClientID: database.MustToDBUUID(params.ClientID), ActiveVersion: params.ActiveVersion, MinCleanVersion: params.MinCleanVersion, MinTextVersion: params.MinTextVersion, Fields: fields, }, nil } func (s *Service) normalizeCodeVersions(current *collector.Collector, params *SetParams) error { if current == nil { return errors.New("current collector required") } if params == nil { return nil } if params.MinCleanVersion != nil && *params.MinCleanVersion == current.MinCleanVersion { params.MinCleanVersion = nil } if params.MinCleanVersion != nil { err := s.svc.CleanVersion.IsValidVersion(*params.MinCleanVersion) if err != nil { return err } } if params.MinTextVersion != nil && *params.MinTextVersion == current.MinTextVersion { params.MinTextVersion = nil } if params.MinTextVersion != nil { err := s.svc.TextVersion.IsValidVersion(*params.MinTextVersion) if err != nil { return err } } return nil } func (s *Service) normalizeActiveVersion(current *collector.Collector, params *SetParams) error { if current == nil { return errors.New("current collector required") } if params == nil || params.ActiveVersion == nil { return nil } err := validation.NormalizeInClosedInterval(¶ms.ActiveVersion, current.ActiveVersion, 1, current.LatestVersion+1) if err != nil { return err } return nil } func (s *Service) normalizeSetFieldsToDB(ctx context.Context, current map[string]uuid.UUID, ofields *map[string]uuid.UUID) (*map[string]pgtype.UUID, error) { if ofields == nil || *ofields == nil { return nil, nil } else if len(*ofields) == 0 { return nil, nil } dbm := map[string]pgtype.UUID{} for name, id := range *ofields { dbm[name] = database.MustToDBUUID(id) } removeIDs := getRemoveFields(current, &dbm) addIDs := getAddFields(current, &dbm) if len(addIDs) == 0 && len(removeIDs) == 0 { return nil, nil } return s.normalizeFieldsToDB(ctx, ofields) } func (s *Service) submitSet(ctx context.Context, current *collector.Collector, params *dbSetParams) error { err := s.cfg.ExecuteDBTransaction(ctx, func(ctx context.Context, qtx *repository.Queries) error { activeName, err := validation.GetFieldName(params, params.ActiveVersion) if err != nil { return err } onlyactive := validation.AreAllPointersNilExcept(params, activeName) var latestVersion int32 if !onlyactive { latestVersion, err = qtx.AddLatestCollectorVersion(ctx, params.ClientID) if err != nil { return err } } if params.ActiveVersion != nil { err := qtx.SetActiveCollectorVersion(ctx, &repository.SetActiveCollectorVersionParams{ Clientid: params.ClientID, Versionid: *params.ActiveVersion, }) if err != nil { return err } } if params.MinCleanVersion != nil { err = qtx.SetCollectorCleanVersion(ctx, &repository.SetCollectorCleanVersionParams{ Clientid: params.ClientID, Versionid: *params.MinCleanVersion, Addedversion: latestVersion, }) if err != nil { return err } } if params.MinTextVersion != nil { err = qtx.SetCollectorTextVersion(ctx, &repository.SetCollectorTextVersionParams{ Clientid: params.ClientID, Versionid: *params.MinTextVersion, Addedversion: latestVersion, }) if err != nil { return err } } if params.Fields != nil { removeIDs := getRemoveFields(current.Fields, params.Fields) for _, field := range removeIDs { err := qtx.RemoveCollectorQuery(ctx, &repository.RemoveCollectorQueryParams{ Clientid: params.ClientID, Queryid: field, Removedversion: &latestVersion, }) if err != nil { return err } } addIDs := getAddFields(current.Fields, params.Fields) for key, field := range addIDs { err := qtx.AddCollectorQuery(ctx, &repository.AddCollectorQueryParams{ Clientid: params.ClientID, Name: key, Queryid: field, Addedversion: latestVersion, }) if err != nil { return err } } } slog.Debug("client collector updated", "update", *params) return nil }) if err != nil { return err } return nil } func getRemoveFields(current map[string]uuid.UUID, update *map[string]pgtype.UUID) []pgtype.UUID { diff := []pgtype.UUID{} if update == nil { return diff } for ckey, cid := range current { found := false for ukey, uid := range *update { if cid == database.MustToUUID(uid) && ckey == ukey { found = true } } if !found { diff = append(diff, database.MustToDBUUID(cid)) } } return diff } func getAddFields(current map[string]uuid.UUID, update *map[string]pgtype.UUID) map[string]pgtype.UUID { diff := map[string]pgtype.UUID{} if update == nil { return diff } for ukey, uid := range *update { if current[ukey] != database.MustToUUID(uid) { diff[ukey] = uid } } return diff } func (s *Service) normalizeFieldsToDB(ctx context.Context, ofields *map[string]uuid.UUID) (*map[string]pgtype.UUID, error) { if ofields == nil || *ofields == nil { return nil, nil } else if len(*ofields) == 0 { return nil, nil } dbm := map[string]pgtype.UUID{} for name, id := range *ofields { dbm[name] = database.MustToDBUUID(id) } dbids := []pgtype.UUID{} for _, id := range dbm { dbids = append(dbids, id) } dedup := validation.DeduplicateArray(dbids) if len(dedup) != len(dbids) { return nil, errors.New("duplicate output fields") } exist, err := s.cfg.GetDBQueries().AllQueriesExist(ctx, dbids) if err != nil { return nil, err } else if !exist { return nil, errors.New("not all required ids are present") } return &dbm, nil }