Review LLM concurrency refactor
This commit is contained in:
@@ -179,6 +179,8 @@ func TestRunnerProposalSectionConcurrencyBoundedByProposalLLMConcurrency(t *test
|
||||
|
||||
var inFlight int32
|
||||
var maxInFlight int32
|
||||
entered := make(chan struct{}, len(transcript.Segments))
|
||||
release := make(chan struct{})
|
||||
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
|
||||
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
|
||||
current := atomic.AddInt32(&inFlight, 1)
|
||||
@@ -188,19 +190,29 @@ func TestRunnerProposalSectionConcurrencyBoundedByProposalLLMConcurrency(t *test
|
||||
break
|
||||
}
|
||||
}
|
||||
time.Sleep(25 * time.Millisecond)
|
||||
entered <- struct{}{}
|
||||
<-release
|
||||
atomic.AddInt32(&inFlight, -1)
|
||||
return nil, nil
|
||||
}},
|
||||
}})
|
||||
|
||||
_, err := r.Run(context.Background(), RunInput{
|
||||
Config: &cfg,
|
||||
Transcript: transcript,
|
||||
ModuleSpecs: []contracts.ModuleRunSpec{
|
||||
{ModuleKey: "m", InstanceName: "m"},
|
||||
},
|
||||
})
|
||||
resultCh := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := r.Run(context.Background(), RunInput{
|
||||
Config: &cfg,
|
||||
Transcript: transcript,
|
||||
ModuleSpecs: []contracts.ModuleRunSpec{
|
||||
{ModuleKey: "m", InstanceName: "m"},
|
||||
},
|
||||
})
|
||||
resultCh <- err
|
||||
}()
|
||||
|
||||
waitForRunnerEntries(t, entered, 2, "proposal workers to enter")
|
||||
close(release)
|
||||
|
||||
err := <-resultCh
|
||||
if err != nil {
|
||||
t.Fatalf("Run error: %v", err)
|
||||
}
|
||||
@@ -225,27 +237,52 @@ func TestRunnerProposalIndexOrderingIsDeterministicAcrossParallelSections(t *tes
|
||||
{ID: 3, Text: "charlie three"},
|
||||
}}
|
||||
|
||||
started := make(chan int, len(transcript.Segments))
|
||||
releaseByID := map[int]chan struct{}{
|
||||
1: make(chan struct{}),
|
||||
2: make(chan struct{}),
|
||||
3: make(chan struct{}),
|
||||
}
|
||||
|
||||
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
|
||||
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
|
||||
if len(req.WorkingTranscript.Segments) != 1 {
|
||||
t.Fatalf("expected one segment per section, got %d", len(req.WorkingTranscript.Segments))
|
||||
}
|
||||
seg := req.WorkingTranscript.Segments[0]
|
||||
// Sleep inversely by ID to force completion order to differ from section order.
|
||||
time.Sleep(time.Duration(4-seg.ID) * 10 * time.Millisecond)
|
||||
started <- seg.ID
|
||||
<-releaseByID[seg.ID]
|
||||
return []proposals.CorrectionProposal{
|
||||
{TargetSegmentID: seg.ID, OriginalText: seg.Text, CorrectedText: strings.ToUpper(seg.Text), Confidence: 1},
|
||||
}, nil
|
||||
}},
|
||||
}})
|
||||
|
||||
out, err := r.Run(context.Background(), RunInput{
|
||||
Config: &cfg,
|
||||
Transcript: transcript,
|
||||
ModuleSpecs: []contracts.ModuleRunSpec{
|
||||
{ModuleKey: "m", InstanceName: "m"},
|
||||
},
|
||||
})
|
||||
resultCh := make(chan struct {
|
||||
out RunOutput
|
||||
err error
|
||||
}, 1)
|
||||
go func() {
|
||||
out, err := r.Run(context.Background(), RunInput{
|
||||
Config: &cfg,
|
||||
Transcript: transcript,
|
||||
ModuleSpecs: []contracts.ModuleRunSpec{
|
||||
{ModuleKey: "m", InstanceName: "m"},
|
||||
},
|
||||
})
|
||||
resultCh <- struct {
|
||||
out RunOutput
|
||||
err error
|
||||
}{out: out, err: err}
|
||||
}()
|
||||
|
||||
waitForRunnerSectionIDs(t, started, map[int]struct{}{1: {}, 2: {}, 3: {}})
|
||||
close(releaseByID[3])
|
||||
close(releaseByID[2])
|
||||
close(releaseByID[1])
|
||||
|
||||
result := <-resultCh
|
||||
out, err := result.out, result.err
|
||||
if err != nil {
|
||||
t.Fatalf("Run error: %v", err)
|
||||
}
|
||||
@@ -269,6 +306,42 @@ func TestRunnerProposalIndexOrderingIsDeterministicAcrossParallelSections(t *tes
|
||||
}
|
||||
}
|
||||
|
||||
func waitForRunnerEntries(t *testing.T, entered <-chan struct{}, want int, label string) {
|
||||
t.Helper()
|
||||
deadline := time.After(350 * time.Millisecond)
|
||||
got := 0
|
||||
for got < want {
|
||||
select {
|
||||
case <-entered:
|
||||
got++
|
||||
case <-deadline:
|
||||
t.Fatalf("timed out waiting for %d entries (%s), got %d", want, label, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func waitForRunnerSectionIDs(t *testing.T, started <-chan int, want map[int]struct{}) {
|
||||
t.Helper()
|
||||
deadline := time.After(350 * time.Millisecond)
|
||||
seen := map[int]struct{}{}
|
||||
for len(seen) < len(want) {
|
||||
select {
|
||||
case id := <-started:
|
||||
seen[id] = struct{}{}
|
||||
case <-deadline:
|
||||
t.Fatalf("timed out waiting for section IDs %v, got %v", mapKeys(want), mapKeys(seen))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func mapKeys(values map[int]struct{}) []int {
|
||||
keys := make([]int, 0, len(values))
|
||||
for key := range values {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func TestRunnerSkippedRecorded(t *testing.T) {
|
||||
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Speaker: "A", Start: 0, End: 1, Text: "word word"}}}
|
||||
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
|
||||
|
||||
Reference in New Issue
Block a user