111 lines
3.3 KiB
Go
111 lines
3.3 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]string, len(binding.References))
|
|
for key, value := range binding.References {
|
|
references[key] = value
|
|
}
|
|
binding.References = references
|
|
}
|
|
binding.Validators = cloneValidatorOverride(binding.Validators)
|
|
return binding
|
|
}
|
|
|
|
func cloneValidatorOverride(override ValidatorOverride) ValidatorOverride {
|
|
return ValidatorOverride{
|
|
Set: override.Set,
|
|
Validators: cloneModuleBindings(override.Validators),
|
|
}
|
|
}
|