Add typed spell validation strategies
This commit is contained in:
@@ -6,12 +6,17 @@ import (
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/llm"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/chunk/scenes"
|
||||
spellcodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/spells"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/extract/spells"
|
||||
spellshape "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/validate/spells/shape"
|
||||
spellsourcerefs "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/validate/spells/source_refs"
|
||||
spellrelatedness "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/validate/spells/source_relatedness"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/generic/merge/appendorder"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/generic/normalize/noop"
|
||||
alwaysaccept "gitea.maximumdirect.net/eric/notarius/internal/modules/generic/validate/always_accept"
|
||||
alwaysreject "gitea.maximumdirect.net/eric/notarius/internal/modules/generic/validate/always_reject"
|
||||
validjson "gitea.maximumdirect.net/eric/notarius/internal/modules/generic/validate/valid_json"
|
||||
validjsonschema "gitea.maximumdirect.net/eric/notarius/internal/modules/generic/validate/valid_json_schema"
|
||||
)
|
||||
@@ -21,16 +26,30 @@ func Register(registries pipeline.Registries, assets *llm.AssetRegistry) error {
|
||||
if err := validateRegistries(registries, assets); err != nil {
|
||||
return err
|
||||
}
|
||||
codec := spellcodec.New()
|
||||
registrations := []struct {
|
||||
name string
|
||||
register func() error
|
||||
}{
|
||||
{name: "spells codec", register: func() error { return pipeline.RegisterArtifactCodec(registries.ArtifactCodecs, spellcodec.New()) }},
|
||||
{name: "spells codec", register: func() error { return pipeline.RegisterArtifactCodec(registries.ArtifactCodecs, codec) }},
|
||||
{name: "scenes chunker", register: func() error { return scenes.Register(registries.Chunkers) }},
|
||||
{name: "spells extractor", register: func() error { return spells.RegisterWithRawAdapter(registries.Extractors, spellcodec.New()) }},
|
||||
{name: "spells extractor", register: func() error { return spells.RegisterWithRawAdapter(registries.Extractors, codec) }},
|
||||
{name: "spell-list appendorder merger", register: func() error {
|
||||
return appendorder.RegisterTyped(registries.Mergers, dnd.SpellListKind, appendSpellLists)
|
||||
}},
|
||||
{name: "spell-list noop normalizer", register: func() error { return noop.RegisterTyped[dnd.SpellList](registries.Normalizers, dnd.SpellListKind) }},
|
||||
{name: "spell shape validator", register: func() error { return spellshape.Register(registries.Validators) }},
|
||||
{name: "legacy spell shape validator", register: func() error { return spellshape.RegisterLegacy(registries.Validators, codec) }},
|
||||
{name: "spell source references validator", register: func() error { return spellsourcerefs.Register(registries.Validators) }},
|
||||
{name: "legacy spell source references validator", register: func() error { return spellsourcerefs.RegisterLegacy(registries.Validators, codec) }},
|
||||
{name: "spell source relatedness validator", register: func() error { return spellrelatedness.Register(registries.Validators) }},
|
||||
{name: "legacy spell source relatedness validator", register: func() error { return spellrelatedness.RegisterLegacy(registries.Validators, codec) }},
|
||||
{name: "spell-list always accept validator", register: func() error {
|
||||
return alwaysaccept.RegisterTyped[dnd.SpellList](registries.Validators, dnd.SpellListKind)
|
||||
}},
|
||||
{name: "spell-list always reject validator", register: func() error {
|
||||
return alwaysreject.RegisterTyped[dnd.SpellList](registries.Validators, dnd.SpellListKind)
|
||||
}},
|
||||
{name: "scenes prompt assets", register: func() error { return scenes.RegisterPromptAssets(assets) }},
|
||||
{name: "spells prompt assets", register: func() error { return spells.RegisterPromptAssets(assets) }},
|
||||
}
|
||||
@@ -55,6 +74,18 @@ func Register(registries pipeline.Registries, assets *llm.AssetRegistry) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func appendSpellLists(values []dnd.SpellList) (dnd.SpellList, error) {
|
||||
count := 0
|
||||
for _, value := range values {
|
||||
count += len(value.SpellCasts)
|
||||
}
|
||||
combined := dnd.SpellList{SpellCasts: make([]dnd.SpellCast, 0, count)}
|
||||
for _, value := range values {
|
||||
combined.SpellCasts = append(combined.SpellCasts, value.SpellCasts...)
|
||||
}
|
||||
return combined, nil
|
||||
}
|
||||
|
||||
func validateRegistries(registries pipeline.Registries, assets *llm.AssetRegistry) error {
|
||||
switch {
|
||||
case registries.Chunkers == nil:
|
||||
@@ -63,6 +94,10 @@ func validateRegistries(registries pipeline.Registries, assets *llm.AssetRegistr
|
||||
return fmt.Errorf("dnd registrar: artifact codec registry must not be nil")
|
||||
case registries.Extractors == nil:
|
||||
return fmt.Errorf("dnd registrar: extractor registry must not be nil")
|
||||
case registries.Mergers == nil:
|
||||
return fmt.Errorf("dnd registrar: merger registry must not be nil")
|
||||
case registries.Normalizers == nil:
|
||||
return fmt.Errorf("dnd registrar: normalizer registry must not be nil")
|
||||
case registries.Validators == nil:
|
||||
return fmt.Errorf("dnd registrar: validator registry must not be nil")
|
||||
case registries.ValidatorChains == nil:
|
||||
|
||||
@@ -28,6 +28,8 @@ func TestRegisterAddsDNDFamily(t *testing.T) {
|
||||
"extract/dnd/spells/shape",
|
||||
"extract/dnd/spells/source_refs",
|
||||
"extract/dnd/spells/source_relatedness",
|
||||
"generic/always_accept",
|
||||
"generic/always_reject",
|
||||
})
|
||||
wantChain := []pipeline.ModuleBinding{
|
||||
pipeline.Binding("generic/valid_json"),
|
||||
@@ -68,6 +70,8 @@ func TestRegisterRejectsMissingDNDDependenciesBeforeMutation(t *testing.T) {
|
||||
{name: "chunkers", remove: func(r *pipeline.Registries, _ **llm.AssetRegistry) { r.Chunkers = nil }, wantErr: "chunker registry"},
|
||||
{name: "artifact codecs", remove: func(r *pipeline.Registries, _ **llm.AssetRegistry) { r.ArtifactCodecs = nil }, wantErr: "artifact codec registry"},
|
||||
{name: "extractors", remove: func(r *pipeline.Registries, _ **llm.AssetRegistry) { r.Extractors = nil }, wantErr: "extractor registry"},
|
||||
{name: "mergers", remove: func(r *pipeline.Registries, _ **llm.AssetRegistry) { r.Mergers = nil }, wantErr: "merger registry"},
|
||||
{name: "normalizers", remove: func(r *pipeline.Registries, _ **llm.AssetRegistry) { r.Normalizers = nil }, wantErr: "normalizer registry"},
|
||||
{name: "validators", remove: func(r *pipeline.Registries, _ **llm.AssetRegistry) { r.Validators = nil }, wantErr: "validator registry"},
|
||||
{name: "validator chains", remove: func(r *pipeline.Registries, _ **llm.AssetRegistry) { r.ValidatorChains = nil }, wantErr: "validator chain registry"},
|
||||
{name: "assets", remove: func(_ *pipeline.Registries, assets **llm.AssetRegistry) { *assets = nil }, wantErr: "asset registry"},
|
||||
|
||||
Reference in New Issue
Block a user