Construct universal modules with decoded options
This commit is contained in:
@@ -21,10 +21,17 @@ const (
|
||||
|
||||
var _ contracts.Chunker = (*Chunker)(nil)
|
||||
|
||||
type Chunker struct{}
|
||||
type Options struct {
|
||||
MaxUnits int
|
||||
OverlapUnits int
|
||||
}
|
||||
|
||||
func New() *Chunker {
|
||||
return &Chunker{}
|
||||
type Chunker struct {
|
||||
options Options
|
||||
}
|
||||
|
||||
func New(options Options) *Chunker {
|
||||
return &Chunker{options: options}
|
||||
}
|
||||
|
||||
func (c *Chunker) Key() string {
|
||||
@@ -39,6 +46,9 @@ func (c *Chunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contra
|
||||
if c == nil {
|
||||
return contracts.ChunkResult{}, chunkerErrorf("chunker must not be nil")
|
||||
}
|
||||
if c.options.MaxUnits <= 0 || c.options.OverlapUnits < 0 || c.options.OverlapUnits >= c.options.MaxUnits {
|
||||
return contracts.ChunkResult{}, chunkerErrorf("chunker options must be initialized by construction")
|
||||
}
|
||||
if ctx == nil {
|
||||
return contracts.ChunkResult{}, chunkerErrorf("context must not be nil")
|
||||
}
|
||||
@@ -55,15 +65,10 @@ func (c *Chunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contra
|
||||
return contracts.ChunkResult{}, chunkerErrorf("validate source document: %w", err)
|
||||
}
|
||||
|
||||
opts, err := chunkOptionsFrom(req.Options)
|
||||
if err != nil {
|
||||
return contracts.ChunkResult{}, err
|
||||
}
|
||||
|
||||
step := opts.maxUnits - opts.overlapUnits
|
||||
step := c.options.MaxUnits - c.options.OverlapUnits
|
||||
chunks := make([]source.Chunk, 0, (len(req.Source.Units)+step-1)/step)
|
||||
for start := 0; start < len(req.Source.Units); start += step {
|
||||
end := start + opts.maxUnits
|
||||
end := start + c.options.MaxUnits
|
||||
if end > len(req.Source.Units) {
|
||||
end = len(req.Source.Units)
|
||||
}
|
||||
@@ -119,36 +124,43 @@ func ModuleSpec() pipeline.ModuleSpec {
|
||||
}
|
||||
|
||||
func Register(registry *pipeline.ChunkerRegistry) error {
|
||||
return registry.RegisterWithSpec(ModuleSpec(), func() (contracts.Chunker, error) {
|
||||
return New(), nil
|
||||
return registry.RegisterBuilderWithSpec(ModuleSpec(), validateOptions, func(request pipeline.BuildRequest) (contracts.Chunker, error) {
|
||||
options, err := DecodeOptions(request.Options)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return New(options), nil
|
||||
})
|
||||
}
|
||||
|
||||
type chunkOptions struct {
|
||||
maxUnits int
|
||||
overlapUnits int
|
||||
func validateOptions(options map[string]any) error {
|
||||
_, err := DecodeOptions(options)
|
||||
return err
|
||||
}
|
||||
|
||||
func chunkOptionsFrom(options map[string]any) (chunkOptions, error) {
|
||||
opts := chunkOptions{
|
||||
maxUnits: defaultMaxUnits,
|
||||
overlapUnits: defaultOverlapUnits,
|
||||
func DecodeOptions(options map[string]any) (Options, error) {
|
||||
if err := pipeline.RejectUnknownOptions(options, "max_units", "overlap_units"); err != nil {
|
||||
return Options{}, chunkerErrorf("%w", err)
|
||||
}
|
||||
opts := Options{
|
||||
MaxUnits: defaultMaxUnits,
|
||||
OverlapUnits: defaultOverlapUnits,
|
||||
}
|
||||
var err error
|
||||
if value, ok := options["max_units"]; ok {
|
||||
opts.maxUnits, err = positiveIntOption("max_units", value)
|
||||
opts.MaxUnits, err = positiveIntOption("max_units", value)
|
||||
if err != nil {
|
||||
return chunkOptions{}, err
|
||||
return Options{}, err
|
||||
}
|
||||
}
|
||||
if value, ok := options["overlap_units"]; ok {
|
||||
opts.overlapUnits, err = nonNegativeIntOption("overlap_units", value)
|
||||
opts.OverlapUnits, err = nonNegativeIntOption("overlap_units", value)
|
||||
if err != nil {
|
||||
return chunkOptions{}, err
|
||||
return Options{}, err
|
||||
}
|
||||
}
|
||||
if opts.overlapUnits >= opts.maxUnits {
|
||||
return chunkOptions{}, chunkerErrorf("overlap_units must be less than max_units")
|
||||
if opts.OverlapUnits >= opts.MaxUnits {
|
||||
return Options{}, chunkerErrorf("overlap_units must be less than max_units")
|
||||
}
|
||||
return opts, nil
|
||||
}
|
||||
|
||||
@@ -44,10 +44,18 @@ func TestModuleSpecAndRegister(t *testing.T) {
|
||||
if slots := chunker.ReferenceSlots(); len(slots) != 0 {
|
||||
t.Fatalf("ReferenceSlots() = %#v, want none", slots)
|
||||
}
|
||||
configured, err := registry.BuildWithRequest(Key, pipeline.BuildRequest{Options: map[string]any{"max_units": 1}})
|
||||
if err != nil {
|
||||
t.Fatalf("BuildWithRequest(%q) error = %v, want nil", Key, err)
|
||||
}
|
||||
result, err := configured.Chunk(context.Background(), contracts.ChunkRequest{Source: testSource(2)})
|
||||
if err != nil || len(result.Chunks) != 2 {
|
||||
t.Fatalf("constructed chunker result = %#v, %v; want two chunks", result, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChunkUsesDefaultsForSingleChunk(t *testing.T) {
|
||||
result, err := New().Chunk(context.Background(), contracts.ChunkRequest{Source: testSource(3), Options: nil})
|
||||
result, err := newChunker(t, nil).Chunk(context.Background(), contracts.ChunkRequest{Source: testSource(3)})
|
||||
if err != nil {
|
||||
t.Fatalf("Chunk() error = %v, want nil", err)
|
||||
}
|
||||
@@ -74,10 +82,7 @@ func TestChunkUsesDefaultsForSingleChunk(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestChunkExactBoundaries(t *testing.T) {
|
||||
result, err := New().Chunk(context.Background(), contracts.ChunkRequest{
|
||||
Source: testSource(6),
|
||||
Options: map[string]any{"max_units": 2},
|
||||
})
|
||||
result, err := newChunker(t, map[string]any{"max_units": 2}).Chunk(context.Background(), contracts.ChunkRequest{Source: testSource(6)})
|
||||
if err != nil {
|
||||
t.Fatalf("Chunk() error = %v, want nil", err)
|
||||
}
|
||||
@@ -93,10 +98,7 @@ func TestChunkExactBoundaries(t *testing.T) {
|
||||
}
|
||||
|
||||
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},
|
||||
})
|
||||
result, err := newChunker(t, map[string]any{"max_units": 3, "overlap_units": 1}).Chunk(context.Background(), contracts.ChunkRequest{Source: testSource(7)})
|
||||
if err != nil {
|
||||
t.Fatalf("Chunk() error = %v, want nil", err)
|
||||
}
|
||||
@@ -123,19 +125,17 @@ func TestChunkRejectsInvalidOptions(t *testing.T) {
|
||||
{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"},
|
||||
{name: "unknown", options: map[string]any{"unexpected": true}, want: "unknown option"},
|
||||
}
|
||||
|
||||
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,
|
||||
})
|
||||
_, err := DecodeOptions(test.options)
|
||||
if err == nil {
|
||||
t.Fatal("Chunk() error = nil, want error")
|
||||
t.Fatal("DecodeOptions() 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)
|
||||
t.Fatalf("DecodeOptions() error = %q, want module context and %q", err.Error(), test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -145,7 +145,7 @@ func TestChunkRejectsEmptySource(t *testing.T) {
|
||||
doc := testSource(1)
|
||||
doc.Units = nil
|
||||
|
||||
_, err := New().Chunk(context.Background(), contracts.ChunkRequest{Source: doc})
|
||||
_, err := newChunker(t, nil).Chunk(context.Background(), contracts.ChunkRequest{Source: doc})
|
||||
if err == nil {
|
||||
t.Fatal("Chunk() error = nil, want empty source error")
|
||||
}
|
||||
@@ -157,10 +157,7 @@ func TestChunkRejectsEmptySource(t *testing.T) {
|
||||
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},
|
||||
})
|
||||
result, err := newChunker(t, map[string]any{"max_units": 1}).Chunk(context.Background(), contracts.ChunkRequest{Source: doc})
|
||||
if err != nil {
|
||||
t.Fatalf("Chunk() error = %v, want nil", err)
|
||||
}
|
||||
@@ -183,6 +180,15 @@ func TestChunkDefensivelyCopiesUnits(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func newChunker(t *testing.T, rawOptions map[string]any) *Chunker {
|
||||
t.Helper()
|
||||
options, err := DecodeOptions(rawOptions)
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeOptions() error = %v, want nil", err)
|
||||
}
|
||||
return New(options)
|
||||
}
|
||||
|
||||
func testSource(count int) *source.SourceDocument {
|
||||
units := make([]source.SourceUnit, 0, count)
|
||||
for i := 1; i <= count; i++ {
|
||||
|
||||
@@ -21,10 +21,15 @@ var safeOutputFileChar = regexp.MustCompile(`[^A-Za-z0-9._-]`)
|
||||
|
||||
var _ contracts.OutputEncoder = (*Encoder)(nil)
|
||||
|
||||
type Encoder struct{}
|
||||
type Options struct{}
|
||||
|
||||
type Encoder struct {
|
||||
options Options
|
||||
}
|
||||
|
||||
func New() *Encoder {
|
||||
return &Encoder{}
|
||||
options, _ := DecodeOptions(nil)
|
||||
return &Encoder{options: options}
|
||||
}
|
||||
|
||||
func (e *Encoder) Key() string {
|
||||
@@ -59,11 +64,27 @@ func ModuleSpec() pipeline.ModuleSpec {
|
||||
}
|
||||
|
||||
func Register(registry *pipeline.OutputEncoderRegistry) error {
|
||||
return registry.RegisterWithSpec(ModuleSpec(), func() (contracts.OutputEncoder, error) {
|
||||
return New(), nil
|
||||
return registry.RegisterBuilderWithSpec(ModuleSpec(), validateOptions, func(request pipeline.BuildRequest) (contracts.OutputEncoder, error) {
|
||||
options, err := DecodeOptions(request.Options)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Encoder{options: options}, nil
|
||||
})
|
||||
}
|
||||
|
||||
func validateOptions(options map[string]any) error {
|
||||
_, err := DecodeOptions(options)
|
||||
return err
|
||||
}
|
||||
|
||||
func DecodeOptions(options map[string]any) (Options, error) {
|
||||
if err := pipeline.RejectUnknownOptions(options); err != nil {
|
||||
return Options{}, encoderErrorf("%w", err)
|
||||
}
|
||||
return Options{}, nil
|
||||
}
|
||||
|
||||
type indexFile struct {
|
||||
ManifestFile string `json:"manifest_file"`
|
||||
OutputFiles []outputFileIndex `json:"output_files"`
|
||||
|
||||
@@ -34,6 +34,9 @@ func TestModuleSpecAndRegister(t *testing.T) {
|
||||
if !reflect.DeepEqual(spec, want) {
|
||||
t.Fatalf("registered spec = %#v, want %#v", spec, want)
|
||||
}
|
||||
if err := registry.ValidateOptions(Key, map[string]any{"unexpected": true}); err == nil || !strings.Contains(err.Error(), "unknown option") {
|
||||
t.Fatalf("ValidateOptions() error = %v, want unknown option error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncodeReturnsLogicalFilesForNormalizedOutputs(t *testing.T) {
|
||||
|
||||
Reference in New Issue
Block a user