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

120 lines
3.5 KiB
Go

package pipeline
import (
"fmt"
"strings"
)
type validatorChainKey struct {
stage ModuleStage
module string
}
type ValidatorChainMapping struct {
Stage ModuleStage `json:"stage"`
Module string `json:"module"`
Validators []ModuleBinding `json:"validators,omitempty"`
}
type ValidatorChainRegistry struct {
chains map[validatorChainKey][]ModuleBinding
}
func NewValidatorChainRegistry() *ValidatorChainRegistry {
return &ValidatorChainRegistry{
chains: make(map[validatorChainKey][]ModuleBinding),
}
}
func (r *ValidatorChainRegistry) Register(mapping ValidatorChainMapping) error {
if r == nil {
return fmt.Errorf("validator chain registry must not be nil")
}
normalized, err := normalizeValidatorChainMapping(mapping)
if err != nil {
return err
}
if r.chains == nil {
r.chains = make(map[validatorChainKey][]ModuleBinding)
}
key := validatorChainKey{stage: normalized.Stage, module: normalized.Module}
if _, exists := r.chains[key]; exists {
return fmt.Errorf("validator chain for %q %q is already registered", normalized.Stage, normalized.Module)
}
r.chains[key] = cloneModuleBindings(normalized.Validators)
return nil
}
func (r *ValidatorChainRegistry) Validators(stage ModuleStage, module string) []ModuleBinding {
if r == nil {
return nil
}
chain := r.chains[validatorChainKey{stage: stage, module: strings.TrimSpace(module)}]
return cloneModuleBindings(chain)
}
func normalizeValidatorChainMapping(mapping ValidatorChainMapping) (ValidatorChainMapping, error) {
normalized := ValidatorChainMapping{
Stage: mapping.Stage,
Module: strings.TrimSpace(mapping.Module),
Validators: cloneModuleBindings(mapping.Validators),
}
switch normalized.Stage {
case StageChunk, StageExtract, StageMerge, StageNormalize:
default:
return ValidatorChainMapping{}, fmt.Errorf("validator chain stage %q is not supported", normalized.Stage)
}
if normalized.Module == "" {
return ValidatorChainMapping{}, fmt.Errorf("validator chain module key must not be empty")
}
for i, validator := range normalized.Validators {
if strings.TrimSpace(validator.Module) == "" {
return ValidatorChainMapping{}, fmt.Errorf("validator chain for %q %q has empty validator key at index %d", normalized.Stage, normalized.Module, i)
}
normalized.Validators[i] = resolveBinding(validator, "")
}
return normalized, nil
}
func cloneModuleBindings(bindings []ModuleBinding) []ModuleBinding {
if len(bindings) == 0 {
return nil
}
out := make([]ModuleBinding, len(bindings))
for i, binding := range bindings {
out[i] = cloneModuleBinding(binding)
}
return out
}
func cloneModuleBinding(binding ModuleBinding) ModuleBinding {
binding.Module = strings.TrimSpace(binding.Module)
binding.LLMProfile = strings.TrimSpace(binding.LLMProfile)
binding.Options = cloneOptions(binding.Options)
if len(binding.References) > 0 {
references := make(map[string]ReferenceSource, len(binding.References))
for key, value := range binding.References {
references[key] = cloneReferenceSource(value)
}
binding.References = references
}
binding.Validators = cloneValidatorOverride(binding.Validators)
return binding
}
func cloneReferenceSource(source ReferenceSource) ReferenceSource {
out := source
if source.Artifact != nil {
artifact := *source.Artifact
out.Artifact = &artifact
}
return out
}
func cloneValidatorOverride(override ValidatorOverride) ValidatorOverride {
return ValidatorOverride{
Set: override.Set,
Validators: cloneModuleBindings(override.Validators),
}
}