package integration_test import ( "context" "encoding/json" "errors" "sync/atomic" "testing" "time" "gitea.maximumdirect.net/eric/notarius/internal/core/config" "gitea.maximumdirect.net/eric/notarius/internal/core/source" "gitea.maximumdirect.net/eric/notarius/internal/framework/contracts" frameworkllm "gitea.maximumdirect.net/eric/notarius/internal/framework/llm" "gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline" "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd" ) const ( concurrentChunkerKey = "test/concurrent-chunks" concurrentExtractorKey = "test/concurrent-extractor" concurrentValidatorKey = "test/concurrent-validator" ) type integrationConcurrencyTracker struct { extractActive atomic.Int32 extractMaximum atomic.Int32 providerActive atomic.Int32 providerMaximum atomic.Int32 providerCalls atomic.Int32 validatorCalls atomic.Int32 } func updateMaximum(maximum *atomic.Int32, current int32) { for { seen := maximum.Load() if current <= seen || maximum.CompareAndSwap(seen, current) { return } } } type instrumentedProvider struct { tracker *integrationConcurrencyTracker } func (p instrumentedProvider) CompleteStructured(ctx context.Context, _ contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) { p.tracker.providerCalls.Add(1) current := p.tracker.providerActive.Add(1) defer p.tracker.providerActive.Add(-1) updateMaximum(&p.tracker.providerMaximum, current) select { case <-time.After(5 * time.Millisecond): case <-ctx.Done(): return contracts.StructuredCompletionResponse{}, ctx.Err() } content := []byte(`{"approved":true}`) if out != nil { if err := json.Unmarshal(content, out); err != nil { return contracts.StructuredCompletionResponse{}, err } } return contracts.StructuredCompletionResponse{Content: content}, nil } type concurrentChunker struct{} func (concurrentChunker) Key() string { return concurrentChunkerKey } func (concurrentChunker) ReferenceSlots() []contracts.ReferenceSlot { return nil } func (concurrentChunker) Plan(_ context.Context, request contracts.ChunkRequest) (contracts.ChunkPlanResult, error) { ranges := make([]source.ChunkRange, len(request.Source.Units)) for i, unit := range request.Source.Units { ranges[i] = source.ChunkRange{StartUnitID: unit.ID, EndUnitID: unit.ID} } return contracts.ChunkPlanResult{Plan: source.ChunkPlan{SourceDigest: request.Source.Digest, Ranges: ranges}}, nil } type concurrentExtractor struct { client contracts.StructuredLLMClient tracker *integrationConcurrencyTracker failed atomic.Bool } func (*concurrentExtractor) Key() string { return concurrentExtractorKey } func (*concurrentExtractor) ReferenceSlots() []contracts.ReferenceSlot { return nil } func (e *concurrentExtractor) Extract(ctx context.Context, request contracts.TypedExtractionRequest) (contracts.TypedExtractionResult[dnd.SpellList], error) { current := e.tracker.extractActive.Add(1) defer e.tracker.extractActive.Add(-1) updateMaximum(&e.tracker.extractMaximum, current) var response map[string]any if _, err := e.client.CompleteStructured(ctx, contracts.StructuredCompletionRequest{StageName: concurrentExtractorKey}, &response); err != nil { return contracts.TypedExtractionResult[dnd.SpellList]{}, err } if e.failed.CompareAndSwap(false, true) { return contracts.TypedExtractionResult[dnd.SpellList]{}, errors.New("retry requested") } return contracts.TypedExtractionResult[dnd.SpellList]{Value: dnd.SpellList{SpellCasts: []dnd.SpellCast{}}}, nil } type concurrentValidator struct { client contracts.StructuredLLMClient tracker *integrationConcurrencyTracker } func (*concurrentValidator) Name() string { return concurrentValidatorKey } func (*concurrentValidator) ExecutionClass() contracts.ExecutionClass { return contracts.ExecutionClassLLMBacked } func (v *concurrentValidator) Validate(ctx context.Context, _ contracts.TypedValidationRequest[dnd.SpellList]) (contracts.ValidationResult, error) { v.tracker.validatorCalls.Add(1) var response map[string]any if _, err := v.client.CompleteStructured(ctx, contracts.StructuredCompletionRequest{StageName: concurrentValidatorKey}, &response); err != nil { return contracts.ValidationResult{}, err } return contracts.ValidationResult{Approved: true}, nil } func TestRunnerIndependentlyBoundsWorkersAndProviderCallsAcrossRegisteredModules(t *testing.T) { tracker := &integrationConcurrencyTracker{} scheduler, err := frameworkllm.NewScheduler(2) if err != nil { t.Fatalf("NewScheduler() error = %v", err) } client := frameworkllm.NewScheduledClient(instrumentedProvider{tracker: tracker}, scheduler) catalog := dndSpellsTestCatalog(t, dndSpellsCatalogSpecs{}) if err := catalog.Chunkers.RegisterWithSpec(pipeline.ModuleSpec{Key: concurrentChunkerKey, Stage: pipeline.StageChunk, Requires: []string{"source.transcript"}, Provides: []string{"chunks"}}, func() (contracts.Chunker, error) { return concurrentChunker{}, nil }); err != nil { t.Fatalf("register chunker: %v", err) } validateOptions := func(options map[string]any) error { return pipeline.RejectUnknownOptions(options) } if err := pipeline.RegisterExtractorBuilder(catalog.Extractors, pipeline.ModuleSpec{Key: concurrentExtractorKey, Stage: pipeline.StageExtract, ArtifactKind: dnd.SpellListKind, Requires: []string{"chunks", "source.transcript"}, Provides: []string{"dnd.spell_casts"}}, validateOptions, func(request pipeline.BuildRequest) (contracts.Extractor[dnd.SpellList], error) { return &concurrentExtractor{client: request.Dependencies.LLM, tracker: tracker}, nil }); err != nil { t.Fatalf("register extractor: %v", err) } validators := pipeline.NewValidatorRegistry() catalog.Validators = validators if err := pipeline.RegisterTypedValidatorBuilder(validators, dnd.SpellListKind, pipeline.ValidatorSpec{Key: concurrentValidatorKey, ExecutionClass: contracts.ExecutionClassLLMBacked}, validateOptions, func(request pipeline.BuildRequest) (contracts.TypedValidator[dnd.SpellList], error) { return &concurrentValidator{client: request.Dependencies.LLM, tracker: tracker}, nil }); err != nil { t.Fatalf("register validator: %v", err) } cfg := loadDNDSpellsPipelineConfig(t) profile := cfg.Pipelines["dnd-spells-fixture"] profile.Chunk = pipeline.Binding(concurrentChunkerKey) validatorOverride := pipeline.ValidatorOverride{Set: true, Validators: []pipeline.ModuleBinding{pipeline.Binding(concurrentValidatorKey)}} profile.Artifacts = map[string]pipeline.ArtifactLaneProfile{ "alpha": {Extract: pipeline.ModuleBinding{Module: concurrentExtractorKey, Retries: 1, Validators: validatorOverride}, Merge: pipeline.Binding(pipeline.DefaultMergeModule), Normalize: pipeline.Binding(pipeline.DefaultNormalizeModule)}, "beta": {Extract: pipeline.ModuleBinding{Module: concurrentExtractorKey, Retries: 1, Validators: validatorOverride}, Merge: pipeline.Binding(pipeline.DefaultMergeModule), Normalize: pipeline.Binding(pipeline.DefaultNormalizeModule)}, } cfg.Pipelines["dnd-spells-fixture"] = profile resolved, err := cfg.Resolve(config.ResolveInput{PipelineID: "dnd-spells-fixture", Catalog: catalog}) if err != nil { t.Fatalf("Resolve() error = %v", err) } registries := pipeline.Registries{Inputs: catalog.Inputs, Chunkers: catalog.Chunkers, ArtifactCodecs: catalog.ArtifactCodecs, ArtifactEvidence: catalog.ArtifactEvidence, Extractors: catalog.Extractors, Mergers: catalog.Mergers, Normalizers: catalog.Normalizers, Validators: catalog.Validators, ValidatorChains: catalog.ValidatorChains, Outputs: catalog.Outputs} output, err := runPreparedPipeline(t, registries, resolved.ResolvedPipeline, client, pipeline.RunInput{RawInput: readDNDSpellsFixture(t), ExtractWorkers: 3}) if err != nil { t.Fatalf("Run() error = %v", err) } if len(output.NormalizeOutputs) != 2 { t.Fatalf("normalize outputs = %d, want 2", len(output.NormalizeOutputs)) } if got := tracker.extractMaximum.Load(); got != 3 { t.Fatalf("maximum extract jobs = %d, want 3", got) } if got := tracker.providerMaximum.Load(); got != 2 { t.Fatalf("maximum provider calls = %d, want 2", got) } if got := tracker.validatorCalls.Load(); got == 0 { t.Fatal("LLM-backed validator was not called") } if got := tracker.providerCalls.Load(); got <= tracker.validatorCalls.Load() { t.Fatalf("provider calls = %d, validator calls = %d; want extractor, retry, and validator calls", got, tracker.validatorCalls.Load()) } }