135 lines
5.3 KiB
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(®istries)
|
|
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)
|
|
}
|
|
}
|