Migrate production modules to raw outputs

This commit is contained in:
2026-07-07 19:23:35 +00:00
parent cc6b050367
commit aa14faa3cb
8 changed files with 247 additions and 78 deletions

View File

@@ -4,6 +4,9 @@ import (
"context"
"encoding/json"
"fmt"
"mime"
"sort"
"strings"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
@@ -34,20 +37,24 @@ func (m *Merger) Merge(ctx context.Context, req contracts.MergeRequest) (contrac
return contracts.MergeResult{}, mergerErrorf("context error before merge: %w", err)
}
if len(req.ExtractOutputs) == 1 {
payload := cloneRawPayload(req.ExtractOutputs[0].Payload)
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: req.ExtractOutputs[0].SourceID,
Schema: req.ExtractOutputs[0].Schema,
SourceID: outputs[0].SourceID,
Schema: outputs[0].Schema,
Payload: payload,
},
}, nil
}
content, err := orderedContent(req.ExtractOutputs)
content, err := mergedContent(outputs)
if err != nil {
return contracts.MergeResult{}, err
}
@@ -55,7 +62,8 @@ func (m *Merger) Merge(ctx context.Context, req contracts.MergeRequest) (contrac
Output: contracts.MergeOutput{
LaneID: req.LaneID,
MergerKey: Key,
SourceID: sourceID(req.ExtractOutputs),
SourceID: sourceID(outputs),
Schema: commonSchema(outputs),
Payload: contracts.RawPayload{
Content: content,
MediaType: "application/json",
@@ -78,30 +86,94 @@ func Register(registry *pipeline.MergerRegistry) error {
})
}
func orderedContent(outputs []contracts.ExtractOutput) ([]byte, error) {
items := make([]map[string]any, 0, len(outputs))
func orderedOutputs(outputs []contracts.ExtractOutput) ([]contracts.ExtractOutput, error) {
ordered := make([]contracts.ExtractOutput, 0, len(outputs))
for _, output := range outputs {
item := map[string]any{
"chunk_id": output.ChunkID,
"chunk_index": output.ChunkIndex,
"media_type": output.Payload.MediaType,
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) {
item["content"] = json.RawMessage(append([]byte(nil), output.Payload.Content...))
} else {
item["content"] = string(output.Payload.Content)
if !json.Valid(output.Payload.Content) {
return nil, mergerErrorf("extract output for chunk %q contains invalid JSON", output.ChunkID)
}
items = append(items, item)
ordered = append(ordered, cloneExtractOutput(output))
}
content, err := json.Marshal(struct {
Outputs []map[string]any `json:"outputs"`
}{Outputs: items})
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 != "" {
@@ -111,6 +183,24 @@ func sourceID(outputs []contracts.ExtractOutput) string {
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...),