Files
notarius/internal/framework/pipeline/merger_registry.go

156 lines
5.3 KiB
Go

package pipeline
import (
"context"
"fmt"
"reflect"
"strings"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
)
type artifactVariantKey struct {
module string
kind contracts.ArtifactKind
}
type MergerRegistry struct {
typedEntries map[artifactVariantKey]typedMergerEntry
}
type typedMergerEntry struct {
spec ModuleSpec
valueType reflect.Type
validateOptions OptionValidator
builder func(BuildRequest) (any, error)
merge typedMergeOperation
}
func NewMergerRegistry() *MergerRegistry {
return &MergerRegistry{
typedEntries: make(map[artifactVariantKey]typedMergerEntry),
}
}
func RegisterMerger[T any](registry *MergerRegistry, spec ModuleSpec, constructor func() (contracts.Merger[T], error)) error {
if constructor == nil {
return fmt.Errorf("merger constructor for %q must not be nil", strings.TrimSpace(spec.Key))
}
return RegisterMergerBuilder(registry, spec, rejectUnconfiguredOptions, func(BuildRequest) (contracts.Merger[T], error) {
return constructor()
})
}
func RegisterMergerBuilder[T any](registry *MergerRegistry, spec ModuleSpec, validateOptions OptionValidator, builder func(BuildRequest) (contracts.Merger[T], error)) error {
if registry == nil {
return fmt.Errorf("merger registry must not be nil")
}
normalizedSpec := normalizeModuleSpec(spec)
if err := validateModuleSpec("merger", StageMerge, normalizedSpec); err != nil {
return err
}
if normalizedSpec.ArtifactKind == "" {
return fmt.Errorf("typed merger %q artifact kind must not be empty", normalizedSpec.Key)
}
if validateOptions == nil {
return fmt.Errorf("merger option validator for %q must not be nil", normalizedSpec.Key)
}
if builder == nil {
return fmt.Errorf("merger builder for %q must not be nil", normalizedSpec.Key)
}
key := artifactVariantKey{module: normalizedSpec.Key, kind: normalizedSpec.ArtifactKind}
if _, ok := registry.typedEntries[key]; ok {
return fmt.Errorf("merger %q variant for artifact kind %q is already registered", key.module, key.kind)
}
if registry.typedEntries == nil {
registry.typedEntries = make(map[artifactVariantKey]typedMergerEntry)
}
registry.typedEntries[key] = typedMergerEntry{
spec: cloneModuleSpec(normalizedSpec),
valueType: reflect.TypeFor[T](),
validateOptions: validateOptions,
builder: func(request BuildRequest) (any, error) {
return builder(cloneBuildRequest(request))
},
merge: func(ctx context.Context, implementation any, request contracts.TypedMergeRequest[any]) (erasedTypedResult, error) {
merger, ok := implementation.(contracts.Merger[T])
if !ok {
return erasedTypedResult{}, fmt.Errorf("merger %q has incompatible implementation %T", normalizedSpec.Key, implementation)
}
outputs := make([]contracts.ExtractArtifact[T], len(request.ExtractOutputs))
for i, output := range request.ExtractOutputs {
value, err := exactTypedValue[T]("merge extract value", output.Value)
if err != nil {
return erasedTypedResult{}, err
}
outputs[i] = contracts.ExtractArtifact[T]{LaneID: output.LaneID, ExtractorKey: output.ExtractorKey, SourceID: output.SourceID, ChunkID: output.ChunkID, ChunkIndex: output.ChunkIndex, ChunkRef: output.ChunkRef, Value: value}
}
result, err := merger.Merge(ctx, contracts.TypedMergeRequest[T]{Source: request.Source, LaneID: request.LaneID, ExtractOutputs: outputs, SourceInput: request.SourceInput, SessionID: request.SessionID, References: request.References, LLMProfile: request.LLMProfile, Metadata: request.Metadata})
if err != nil {
return erasedTypedResult{}, err
}
return erasedTypedResult{Value: result.Value, Warnings: result.Warnings}, nil
},
}
return nil
}
func (r *MergerRegistry) validateOptions(key string, kind contracts.ArtifactKind, options map[string]any) error {
if r == nil {
return fmt.Errorf("merger registry must not be nil")
}
normalizedKey := strings.TrimSpace(key)
entry, ok := r.typedEntry(normalizedKey, kind)
if !ok {
return fmt.Errorf("merger %q variant for artifact kind %q is not registered", normalizedKey, kind)
}
return validateRegisteredOptions(entry.validateOptions, options)
}
func (r *MergerRegistry) Spec(key string) (ModuleSpec, bool) {
if r == nil {
return ModuleSpec{}, false
}
module := strings.TrimSpace(key)
for variant, entry := range r.typedEntries {
if variant.module == module {
return cloneModuleSpec(entry.spec), true
}
}
return ModuleSpec{}, false
}
func (r *MergerRegistry) typedEntry(key string, kind contracts.ArtifactKind) (typedMergerEntry, bool) {
if r == nil {
return typedMergerEntry{}, false
}
entry, ok := r.typedEntries[artifactVariantKey{module: strings.TrimSpace(key), kind: normalizeArtifactKind(kind)}]
return entry, ok
}
func (r *MergerRegistry) registeredKinds(key string) []contracts.ArtifactKind {
if r == nil {
return nil
}
module := strings.TrimSpace(key)
kinds := make([]contracts.ArtifactKind, 0)
for variant := range r.typedEntries {
if variant.module == module {
kinds = append(kinds, variant.kind)
}
}
sortArtifactKinds(kinds)
return kinds
}
func (r *MergerRegistry) RegisteredKeys() []string {
if r == nil {
return nil
}
keys := make(map[string]struct{}, len(r.typedEntries))
for key := range r.typedEntries {
keys[key.module] = struct{}{}
}
return sortedRegistryKeys(keys)
}