Files
notarius/internal/modules/generic/register/register_test.go

135 lines
5.3 KiB
Go

package register
import (
"io/fs"
"strings"
"testing"
"gitea.maximumdirect.net/eric/notarius/internal/framework/llm"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
)
func TestRegisterAddsGenericFamily(t *testing.T) {
registries := completeRegistries()
assets := llm.NewAssetRegistry()
if err := Register(registries, assets); err != nil {
t.Fatalf("Register() error = %v, want nil", err)
}
assertContainsKeys(t, "chunkers", registries.Chunkers.RegisteredKeys(), []string{"generic"})
assertContainsKeys(t, "validators", registries.Validators.RegisteredKeys(), []string{
"generic/always_accept",
"generic/always_reject",
"generic/valid_json",
"generic/valid_json_schema",
})
assertContainsKeys(t, "outputs", registries.Outputs.RegisteredKeys(), []string{"json"})
assertNoKeys(t, "inputs", registries.Inputs.RegisteredKeys(), "generic registrar to leave inputs unchanged")
assertNoKeys(t, "extractors", registries.Extractors.RegisteredKeys(), "generic registrar to leave extractors unchanged")
assertNoKeys(t, "mergers", registries.Mergers.RegisteredKeys(), "generic registrar to leave mergers for typed family composition")
assertNoKeys(t, "normalizers", registries.Normalizers.RegisteredKeys(), "generic registrar to leave normalizers for typed family composition")
if chunker, err := registries.Chunkers.Build("generic"); err != nil || chunker.Key() != "generic" {
t.Fatalf("build generic chunker = %v, %v; want generic implementation", chunker, err)
}
if output, err := registries.Outputs.Build("json"); err != nil || output.Key() != "json" {
t.Fatalf("build json output = %v, %v; want json implementation", output, err)
}
promptAssets, err := assets.PromptFS()
if err != nil {
t.Fatal(err)
}
if _, err := fs.ReadFile(promptAssets, "generic.semantic_reconciliation/prompt.yaml"); err != nil {
t.Fatalf("registered semantic reconciliation prompt: %v", err)
}
schemaAssets, err := assets.SchemaFS()
if err != nil {
t.Fatal(err)
}
if _, err := fs.ReadFile(schemaAssets, "semantic_reconciliation_llm.v1.json"); err != nil {
t.Fatalf("registered semantic reconciliation schema: %v", err)
}
}
func TestRegisterRejectsNilAssetRegistryBeforeMutation(t *testing.T) {
registries := completeRegistries()
err := Register(registries, nil)
if err == nil || !strings.Contains(err.Error(), "asset registry must not be nil") {
t.Fatalf("Register() error = %v, want nil asset registry error", err)
}
if len(registries.Chunkers.RegisteredKeys()) != 0 {
t.Fatalf("chunker keys = %#v, want validation before mutation", registries.Chunkers.RegisteredKeys())
}
}
func TestRegisterRejectsMissingGenericRegistriesBeforeMutation(t *testing.T) {
tests := []struct {
name string
remove func(*pipeline.Registries)
wantErr string
}{
{name: "chunkers", remove: func(r *pipeline.Registries) { r.Chunkers = nil }, wantErr: "chunker registry"},
{name: "mergers", remove: func(r *pipeline.Registries) { r.Mergers = nil }, wantErr: "merger registry"},
{name: "normalizers", remove: func(r *pipeline.Registries) { r.Normalizers = nil }, wantErr: "normalizer registry"},
{name: "validators", remove: func(r *pipeline.Registries) { r.Validators = nil }, wantErr: "validator registry"},
{name: "outputs", remove: func(r *pipeline.Registries) { r.Outputs = nil }, wantErr: "output registry"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
registries := completeRegistries()
test.remove(&registries)
err := Register(registries, llm.NewAssetRegistry())
if err == nil || !strings.Contains(err.Error(), test.wantErr) {
t.Fatalf("Register() error = %v, want %q", err, test.wantErr)
}
if got := registries.Chunkers; got != nil && len(got.RegisteredKeys()) != 0 {
t.Fatalf("chunker keys = %#v, want validation before mutation", got.RegisteredKeys())
}
})
}
}
func TestRegisterReportsDuplicateGenericRegistration(t *testing.T) {
registries := completeRegistries()
assets := llm.NewAssetRegistry()
if err := Register(registries, assets); err != nil {
t.Fatalf("first Register() error = %v, want nil", err)
}
err := Register(registries, assets)
if err == nil || !strings.Contains(err.Error(), "register semantic reconciliation assets") || !strings.Contains(err.Error(), "already registered") {
t.Fatalf("second Register() error = %v, want contextual duplicate error", err)
}
}
func completeRegistries() pipeline.Registries {
return pipeline.Registries{
Inputs: pipeline.NewInputAdapterRegistry(),
Chunkers: pipeline.NewChunkerRegistry(),
ArtifactCodecs: pipeline.NewArtifactCodecRegistry(),
Extractors: pipeline.NewExtractorRegistry(),
Mergers: pipeline.NewMergerRegistry(),
Normalizers: pipeline.NewNormalizerRegistry(),
Validators: pipeline.NewValidatorRegistry(),
ValidatorChains: pipeline.NewValidatorChainRegistry(),
Outputs: pipeline.NewOutputEncoderRegistry(),
}
}
func assertContainsKeys(t *testing.T, name string, got, want []string) {
t.Helper()
seen := make(map[string]struct{}, len(got))
for _, key := range got {
seen[key] = struct{}{}
}
for _, key := range want {
if _, ok := seen[key]; !ok {
t.Fatalf("%s keys = %#v, want required key %q", name, got, key)
}
}
}
func assertNoKeys(t *testing.T, name string, got []string, reason string) {
t.Helper()
if len(got) != 0 {
t.Fatalf("%s keys = %#v, want %s", name, got, reason)
}
}