Files
notarius/internal/modules/merge/appendorder/merger.go

227 lines
5.6 KiB
Go

package appendorder
import (
"context"
"encoding/json"
"fmt"
"mime"
"sort"
"strings"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
)
const Key = "appendorder"
var _ contracts.Merger = (*Merger)(nil)
type Merger struct{}
func New() *Merger {
return &Merger{}
}
func (m *Merger) Key() string {
return Key
}
func (m *Merger) Merge(ctx context.Context, req contracts.MergeRequest) (contracts.MergeResult, error) {
if m == nil {
return contracts.MergeResult{}, mergerErrorf("merger must not be nil")
}
if ctx == nil {
return contracts.MergeResult{}, mergerErrorf("context must not be nil")
}
if err := ctx.Err(); err != nil {
return contracts.MergeResult{}, mergerErrorf("context error before merge: %w", err)
}
outputs, err := orderedOutputs(req.ExtractOutputs)
if err != nil {
return contracts.MergeResult{}, err
}
if len(outputs) == 1 {
payload := cloneRawPayload(outputs[0].Payload)
return contracts.MergeResult{
Output: contracts.MergeOutput{
LaneID: req.LaneID,
MergerKey: Key,
SourceID: outputs[0].SourceID,
Schema: outputs[0].Schema,
Payload: payload,
},
}, nil
}
content, err := mergedContent(outputs)
if err != nil {
return contracts.MergeResult{}, err
}
return contracts.MergeResult{
Output: contracts.MergeOutput{
LaneID: req.LaneID,
MergerKey: Key,
SourceID: sourceID(outputs),
Schema: commonSchema(outputs),
Payload: contracts.RawPayload{
Content: content,
MediaType: "application/json",
},
},
}, nil
}
func ModuleSpec() pipeline.ModuleSpec {
return pipeline.ModuleSpec{
Key: Key,
Stage: pipeline.StageMerge,
Provides: []string{"merged"},
}
}
func Register(registry *pipeline.MergerRegistry) error {
return registry.RegisterWithSpec(ModuleSpec(), func() (contracts.Merger, error) {
return New(), nil
})
}
func orderedOutputs(outputs []contracts.ExtractOutput) ([]contracts.ExtractOutput, error) {
ordered := make([]contracts.ExtractOutput, 0, len(outputs))
for _, output := range outputs {
if !isJSONMediaType(output.Payload.MediaType) {
return nil, mergerErrorf("extract output for chunk %q has unsupported media type %q", output.ChunkID, output.Payload.MediaType)
}
if !json.Valid(output.Payload.Content) {
return nil, mergerErrorf("extract output for chunk %q contains invalid JSON", output.ChunkID)
}
ordered = append(ordered, cloneExtractOutput(output))
}
sort.SliceStable(ordered, func(i, j int) bool {
return ordered[i].ChunkIndex < ordered[j].ChunkIndex
})
return ordered, nil
}
func mergedContent(outputs []contracts.ExtractOutput) ([]byte, error) {
values := make([]any, 0, len(outputs))
objects := make([]map[string]any, 0, len(outputs))
for _, output := range outputs {
var value any
if err := json.Unmarshal(output.Payload.Content, &value); err != nil {
return nil, mergerErrorf("decode extract output for chunk %q: %w", output.ChunkID, err)
}
values = append(values, value)
object, ok := value.(map[string]any)
if !ok {
continue
}
objects = append(objects, object)
}
if len(objects) == len(outputs) {
if field, ok := commonArrayField(objects); ok {
merged := make([]any, 0)
for _, object := range objects {
items := object[field].([]any)
merged = append(merged, items...)
}
return marshalMerged(map[string]any{field: merged})
}
}
return marshalMerged(values)
}
func commonArrayField(objects []map[string]any) (string, bool) {
if len(objects) == 0 {
return "", false
}
candidates := map[string]struct{}{}
for key, value := range objects[0] {
if _, ok := value.([]any); ok {
candidates[key] = struct{}{}
}
}
for _, object := range objects[1:] {
for key := range candidates {
if _, ok := object[key].([]any); !ok {
delete(candidates, key)
}
}
}
if len(candidates) != 1 {
return "", false
}
for key := range candidates {
return key, true
}
return "", false
}
func marshalMerged(value any) ([]byte, error) {
content, err := json.Marshal(value)
if err != nil {
return nil, mergerErrorf("encode merged output: %w", err)
}
return content, nil
}
func isJSONMediaType(mediaType string) bool {
base, _, err := mime.ParseMediaType(strings.TrimSpace(mediaType))
if err != nil {
base = strings.TrimSpace(mediaType)
}
return strings.EqualFold(base, "application/json")
}
func sourceID(outputs []contracts.ExtractOutput) string {
for _, output := range outputs {
if output.SourceID != "" {
return output.SourceID
}
}
return ""
}
func commonSchema(outputs []contracts.ExtractOutput) contracts.ResponseSchema {
if len(outputs) == 0 {
return contracts.ResponseSchema{}
}
schema := outputs[0].Schema
for _, output := range outputs[1:] {
if output.Schema != schema {
return contracts.ResponseSchema{}
}
}
return schema
}
func cloneExtractOutput(output contracts.ExtractOutput) contracts.ExtractOutput {
output.Payload = cloneRawPayload(output.Payload)
return output
}
func cloneRawPayload(payload contracts.RawPayload) contracts.RawPayload {
return contracts.RawPayload{
Content: append([]byte(nil), payload.Content...),
MediaType: payload.MediaType,
Metadata: cloneMetadata(payload.Metadata),
Warnings: append([]contracts.Warning(nil), payload.Warnings...),
}
}
func cloneMetadata(metadata map[string]any) map[string]any {
if len(metadata) == 0 {
return nil
}
out := make(map[string]any, len(metadata))
for key, value := range metadata {
out[key] = value
}
return out
}
func mergerErrorf(format string, args ...any) error {
return fmt.Errorf("appendorder merger: "+format, args...)
}