Files
notarius/internal/modules/chunk/generic/chunker_test.go

214 lines
6.5 KiB
Go

package generic
import (
"context"
"encoding/json"
"fmt"
"reflect"
"strings"
"testing"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
)
func TestModuleSpecAndRegister(t *testing.T) {
want := pipeline.ModuleSpec{
Key: Key,
Stage: pipeline.StageChunk,
Provides: []string{"chunks"},
}
if got := ModuleSpec(); !reflect.DeepEqual(got, want) {
t.Fatalf("ModuleSpec() = %#v, want %#v", got, want)
}
registry := pipeline.NewChunkerRegistry()
if err := Register(registry); err != nil {
t.Fatalf("Register() error = %v, want nil", err)
}
spec, ok := registry.Spec(Key)
if !ok {
t.Fatalf("Spec(%q) ok = false, want true", Key)
}
if !reflect.DeepEqual(spec, want) {
t.Fatalf("registered spec = %#v, want %#v", spec, want)
}
chunker, err := registry.Build(Key)
if err != nil {
t.Fatalf("Build(%q) error = %v, want nil", Key, err)
}
if chunker.Key() != Key {
t.Fatalf("Key() = %q, want %q", chunker.Key(), Key)
}
}
func TestChunkUsesDefaultsForSingleChunk(t *testing.T) {
result, err := New().Chunk(context.Background(), contracts.ChunkRequest{Source: testSource(3), Options: nil})
if err != nil {
t.Fatalf("Chunk() error = %v, want nil", err)
}
if got := chunkIDs(result.Chunks); !reflect.DeepEqual(got, []string{"chunk-000001"}) {
t.Fatalf("chunk IDs = %#v, want one stable ID", got)
}
chunk := result.Chunks[0]
if chunk.Index != 0 || chunk.SourceID != "source-1" {
t.Fatalf("chunk = %#v, want source and index fields", chunk)
}
if got := unitIDs(chunk.Units); !reflect.DeepEqual(got, []string{"u001", "u002", "u003"}) {
t.Fatalf("unit IDs = %#v, want all units", got)
}
if chunk.Metadata["start_unit_id"] != "u001" || chunk.Metadata["end_unit_id"] != "u003" || chunk.Metadata["unit_count"] != 3 {
t.Fatalf("metadata = %#v, want chunk bounds", chunk.Metadata)
}
}
func TestChunkExactBoundaries(t *testing.T) {
result, err := New().Chunk(context.Background(), contracts.ChunkRequest{
Source: testSource(6),
Options: map[string]any{"max_units": 2},
})
if err != nil {
t.Fatalf("Chunk() error = %v, want nil", err)
}
if got := chunkIDs(result.Chunks); !reflect.DeepEqual(got, []string{"chunk-000001", "chunk-000002", "chunk-000003"}) {
t.Fatalf("chunk IDs = %#v, want stable IDs", got)
}
gotUnits := [][]string{unitIDs(result.Chunks[0].Units), unitIDs(result.Chunks[1].Units), unitIDs(result.Chunks[2].Units)}
wantUnits := [][]string{{"u001", "u002"}, {"u003", "u004"}, {"u005", "u006"}}
if !reflect.DeepEqual(gotUnits, wantUnits) {
t.Fatalf("chunk units = %#v, want %#v", gotUnits, wantUnits)
}
}
func TestChunkOverlap(t *testing.T) {
result, err := New().Chunk(context.Background(), contracts.ChunkRequest{
Source: testSource(7),
Options: map[string]any{"max_units": 3, "overlap_units": 1},
})
if err != nil {
t.Fatalf("Chunk() error = %v, want nil", err)
}
gotUnits := make([][]string, 0, len(result.Chunks))
for _, chunk := range result.Chunks {
gotUnits = append(gotUnits, unitIDs(chunk.Units))
}
wantUnits := [][]string{{"u001", "u002", "u003"}, {"u003", "u004", "u005"}, {"u005", "u006", "u007"}}
if !reflect.DeepEqual(gotUnits, wantUnits) {
t.Fatalf("chunk units = %#v, want %#v", gotUnits, wantUnits)
}
}
func TestChunkRejectsInvalidOptions(t *testing.T) {
tests := []struct {
name string
options map[string]any
want string
}{
{name: "max wrong type", options: map[string]any{"max_units": "2"}, want: "max_units"},
{name: "max fractional", options: map[string]any{"max_units": 1.5}, want: "integer"},
{name: "max zero", options: map[string]any{"max_units": 0}, want: "positive"},
{name: "overlap negative", options: map[string]any{"overlap_units": -1}, want: "non-negative"},
{name: "overlap too large", options: map[string]any{"max_units": 2, "overlap_units": 2}, want: "less than"},
{name: "json number", options: map[string]any{"max_units": json.Number("bad")}, want: "integer"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
_, err := New().Chunk(context.Background(), contracts.ChunkRequest{
Source: testSource(3),
Options: test.options,
})
if err == nil {
t.Fatal("Chunk() error = nil, want error")
}
if !strings.Contains(err.Error(), "generic chunker") || !strings.Contains(err.Error(), test.want) {
t.Fatalf("Chunk() error = %q, want module context and %q", err.Error(), test.want)
}
})
}
}
func TestChunkRejectsEmptySource(t *testing.T) {
doc := testSource(1)
doc.Units = nil
_, err := New().Chunk(context.Background(), contracts.ChunkRequest{Source: doc})
if err == nil {
t.Fatal("Chunk() error = nil, want empty source error")
}
if !strings.Contains(err.Error(), "generic chunker") || !strings.Contains(err.Error(), "units") {
t.Fatalf("Chunk() error = %q, want empty source context", err.Error())
}
}
func TestChunkDefensivelyCopiesUnits(t *testing.T) {
doc := testSource(2)
result, err := New().Chunk(context.Background(), contracts.ChunkRequest{
Source: doc,
Options: map[string]any{"max_units": 1},
})
if err != nil {
t.Fatalf("Chunk() error = %v, want nil", err)
}
if len(result.Chunks) != 2 {
t.Fatalf("len(Chunks) = %d, want 2", len(result.Chunks))
}
doc.Units[0].ID = "changed"
doc.Units[0].Metadata["speaker"] = "changed"
if result.Chunks[0].Units[0].ID != "u001" {
t.Fatalf("chunk unit ID changed after source mutation: %#v", result.Chunks[0].Units[0])
}
if result.Chunks[0].Units[0].Metadata["speaker"] != "speaker-001" {
t.Fatalf("chunk unit metadata changed after source mutation: %#v", result.Chunks[0].Units[0].Metadata)
}
}
func testSource(count int) *source.SourceDocument {
units := make([]source.SourceUnit, 0, count)
for i := 1; i <= count; i++ {
id := "u" + zeroPad3(i)
units = append(units, source.SourceUnit{
ID: id,
Kind: "unit",
Text: "Text for " + id,
Metadata: map[string]any{
"speaker": "speaker-" + zeroPad3(i),
},
})
}
return &source.SourceDocument{
ID: "source-1",
Kind: "document",
Format: "text/plain",
Digest: "sha256:source",
Units: units,
}
}
func zeroPad3(value int) string {
return fmt.Sprintf("%03d", value)
}
func chunkIDs(chunks []contracts.SourceChunk) []string {
ids := make([]string, 0, len(chunks))
for _, chunk := range chunks {
ids = append(ids, chunk.ID)
}
return ids
}
func unitIDs(units []source.SourceUnit) []string {
ids := make([]string, 0, len(units))
for _, unit := range units {
ids = append(ids, unit.ID)
}
return ids
}