package units import ( "context" "encoding/json" "fmt" "math" "strconv" "gitea.maximumdirect.net/eric/notarius/internal/core/source" "gitea.maximumdirect.net/eric/notarius/internal/framework/contracts" "gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline" ) const Key = "generic" const ( defaultMaxUnits = 50 defaultOverlapUnits = 0 ) var _ contracts.Chunker = (*Chunker)(nil) type Options struct { MaxUnits int OverlapUnits int } type Chunker struct { options Options } func New(options Options) *Chunker { return &Chunker{options: options} } func (c *Chunker) Key() string { return Key } func (c *Chunker) ReferenceSlots() []contracts.ReferenceSlot { return nil } func (c *Chunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contracts.ChunkResult, error) { 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") } if err := ctx.Err(); err != nil { return contracts.ChunkResult{}, chunkerErrorf("context error before chunking: %w", err) } if req.Source == nil { return contracts.ChunkResult{}, chunkerErrorf("source must not be nil") } if len(req.Source.Units) == 0 { return contracts.ChunkResult{}, chunkerErrorf("source units must not be empty") } if err := source.ValidateDocument(req.Source); err != nil { return contracts.ChunkResult{}, chunkerErrorf("validate source document: %w", err) } 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 + c.options.MaxUnits if end > len(req.Source.Units) { end = len(req.Source.Units) } units := cloneUnits(req.Source.Units[start:end]) content, err := chunkContent(units) if err != nil { return contracts.ChunkResult{}, err } chunks = append(chunks, source.Chunk{ ID: fmt.Sprintf("chunk-%06d", len(chunks)+1), SourceID: req.Source.ID, Index: len(chunks), Ref: source.SourceRef{ SourceID: req.Source.ID, StartUnitID: units[0].Ref.StartUnitID, EndUnitID: units[len(units)-1].Ref.EndUnitID, }, Content: content, MediaType: "application/json", Units: units, Metadata: map[string]any{ "start_unit_id": units[0].ID, "end_unit_id": units[len(units)-1].ID, "unit_count": len(units), }, }) if end == len(req.Source.Units) { break } } return contracts.ChunkResult{Chunks: chunks}, nil } func chunkContent(units []source.SourceUnit) ([]byte, error) { content, err := json.Marshal(struct { Units []source.SourceUnit `json:"units"` }{ Units: units, }) if err != nil { return nil, chunkerErrorf("encode chunk content: %w", err) } return content, nil } func ModuleSpec() pipeline.ModuleSpec { return pipeline.ModuleSpec{ Key: Key, Stage: pipeline.StageChunk, Provides: []string{"chunks"}, } } func Register(registry *pipeline.ChunkerRegistry) error { 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 }) } 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, "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) if err != nil { return Options{}, err } } if value, ok := options["overlap_units"]; ok { opts.OverlapUnits, err = nonNegativeIntOption("overlap_units", value) if err != nil { return Options{}, err } } if opts.OverlapUnits >= opts.MaxUnits { return Options{}, chunkerErrorf("overlap_units must be less than max_units") } return opts, nil } func positiveIntOption(name string, value any) (int, error) { got, err := intOption(name, value) if err != nil { return 0, err } if got <= 0 { return 0, chunkerErrorf("%s must be positive", name) } return got, nil } func nonNegativeIntOption(name string, value any) (int, error) { got, err := intOption(name, value) if err != nil { return 0, err } if got < 0 { return 0, chunkerErrorf("%s must be non-negative", name) } return got, nil } func intOption(name string, value any) (int, error) { switch typed := value.(type) { case int: return typed, nil case int8: return int(typed), nil case int16: return int(typed), nil case int32: return int(typed), nil case int64: if typed > maxInt() || typed < minInt() { return 0, chunkerErrorf("%s is outside supported integer range", name) } return int(typed), nil case uint: if uint64(typed) > uint64(maxInt()) { return 0, chunkerErrorf("%s is outside supported integer range", name) } return int(typed), nil case uint8: return int(typed), nil case uint16: return int(typed), nil case uint32: if uint64(typed) > uint64(maxInt()) { return 0, chunkerErrorf("%s is outside supported integer range", name) } return int(typed), nil case uint64: if typed > uint64(maxInt()) { return 0, chunkerErrorf("%s is outside supported integer range", name) } return int(typed), nil case float64: if typed != math.Trunc(typed) { return 0, chunkerErrorf("%s must be an integer", name) } if typed > float64(maxInt()) || typed < float64(minInt()) { return 0, chunkerErrorf("%s is outside supported integer range", name) } return int(typed), nil case json.Number: parsed, err := typed.Int64() if err != nil { return 0, chunkerErrorf("%s must be an integer", name) } if parsed > maxInt() || parsed < minInt() { return 0, chunkerErrorf("%s is outside supported integer range", name) } return int(parsed), nil default: return 0, chunkerErrorf("%s must be an integer", name) } } func maxInt() int64 { return int64(1<<(strconv.IntSize-1) - 1) } func minInt() int64 { return -maxInt() - 1 } func cloneUnits(units []source.SourceUnit) []source.SourceUnit { out := make([]source.SourceUnit, 0, len(units)) for _, unit := range units { out = append(out, source.SourceUnit{ ID: unit.ID, Kind: unit.Kind, Text: unit.Text, Ref: unit.Ref, Metadata: cloneMetadata(unit.Metadata), }) } return out } func cloneMetadata(metadata map[string]any) map[string]any { if len(metadata) == 0 { return nil } out := make(map[string]any, len(metadata)) for key, value := range metadata { out[key] = value } return out } func chunkerErrorf(format string, args ...any) error { return fmt.Errorf("generic chunker: "+format, args...) }