Files
query-orchestration/internal/query/update/update.go
T
Michael McGuinness 53ea7d34e6 Merged in feature/textextract (pull request #108)
Start adding Textract + UUID changes

* base

* startclient

* ts

* short

* tests
2025-03-20 11:06:41 +00:00

237 lines
5.5 KiB
Go

package queryupdate
import (
"context"
"errors"
"fmt"
"log/slog"
queryversionsyncrunner "queryorchestration/api/queryVersionSyncRunner"
"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: &current.ID,
Requiredqueryids: *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 := 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: 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: 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")
}
}