Files
notarius/internal/framework/pipeline/runner_concurrent.go

377 lines
14 KiB
Go

package pipeline
import (
"context"
"errors"
"fmt"
"path"
"sort"
"strings"
"sync"
"time"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
)
type laneExtractState struct {
index int
prepared preparedLaneExecutor
deps []CheckpointFingerprint
decision CheckpointDecision
values []erasedExtractArtifact
serialized []CheckpointArtifact
warnings []contracts.Warning
rejected []contracts.RejectedOutput
results map[int]extractJobResult
remaining int
failed bool
}
type extractJob struct {
lane *laneExtractState
chunk source.Chunk
}
type extractJobResult struct {
laneIndex int
chunkIndex int
value erasedExtractArtifact
serialized CheckpointArtifact
warnings []contracts.Warning
rejected *contracts.RejectedOutput
err error
}
type laneCompletion struct {
index int
output RunOutput
err error
}
type orderedRunError struct {
stage int
lane int
chunk int
err error
}
type completedExtractLoader struct {
CheckpointLoader
laneID string
checkpoint ExtractCheckpoint
}
func (l completedExtractLoader) Extract(laneID, _ string, _ []CheckpointFingerprint) (ExtractCheckpoint, CheckpointDecision) {
if laneID == l.laneID {
return l.checkpoint, CheckpointDecision{Reused: true, Reason: "coordinated extract result"}
}
return ExtractCheckpoint{}, CheckpointDecision{Reason: "extract result unavailable"}
}
func (r *Runner) runLanes(parent context.Context, input RunInput, checkpoints CheckpointRecorder, loader CheckpointLoader, doc *source.SourceDocument, sourceInput contracts.LLMInputMaterial, sessionID string, chunks []source.Chunk) (RunOutput, error) {
output := RunOutput{Manifest: manifestFromPipeline(input)}
states := make([]*laneExtractState, len(input.Prepared.lanes))
for i, prepared := range input.Prepared.lanes {
if prepared.typed == nil {
return output, fmt.Errorf("typed lane %q executor is not prepared", prepared.resolved.ID)
}
setTypedLaneManifestMetadata(&output, prepared.resolved.ID, prepared.typed.extractor, prepared.typed.merger, prepared.typed.normalizer)
state, err := prepareLaneExtract(input, loader, doc, chunks, i, prepared)
if err != nil {
return output, err
}
if !state.decision.Reused {
if err := checkpoints.ExtractRunning(prepared.resolved.ID, prepared.resolved.Extract.Module, state.deps); err != nil {
return output, fmt.Errorf("write extract checkpoint for lane %q: %w", prepared.resolved.ID, err)
}
}
states[i] = state
}
workerCount := input.ExtractWorkers
if workerCount < 1 {
workerCount = 1
}
ctx, cancel := context.WithCancel(parent)
defer cancel()
jobs := make(chan extractJob, workerCount)
results := make(chan extractJobResult, workerCount)
completions := make(chan laneCompletion, len(states))
var workers sync.WaitGroup
for i := 0; i < workerCount; i++ {
workers.Add(1)
go func() {
defer workers.Done()
for job := range jobs {
if ctx.Err() != nil {
continue
}
result := r.runExtractJob(ctx, input, doc, sourceInput, sessionID, job)
results <- result
}
}()
}
go func() {
defer close(jobs)
for chunkIndex := range chunks {
for laneIndex := range states {
state := states[laneIndex]
if state.decision.Reused {
continue
}
select {
case jobs <- extractJob{lane: state, chunk: chunks[chunkIndex]}:
case <-ctx.Done():
return
}
}
}
}()
go func() { workers.Wait(); close(results) }()
completedOutputs := make([]RunOutput, len(states))
var runErrors []orderedRunError
launched, completed := 0, 0
launch := func(state *laneExtractState) {
launched++
go func() {
laneOutput, err := r.continueLane(ctx, input, checkpoints, loader, doc, sourceInput, sessionID, chunks, state)
completions <- laneCompletion{index: state.index, output: laneOutput, err: err}
}()
}
for _, state := range states {
if state.decision.Reused {
launch(state)
}
}
resultChannel := results
for resultChannel != nil || completed < launched {
select {
case result, ok := <-resultChannel:
if !ok {
resultChannel = nil
continue
}
state := states[result.laneIndex]
state.remaining--
if result.err != nil {
state.failed = true
runErrors = append(runErrors, orderedRunError{stage: 0, lane: state.index, chunk: result.chunkIndex, err: result.err})
_ = checkpoints.ExtractFailed(state.prepared.resolved.ID, state.prepared.resolved.Extract.Module, state.deps, result.err)
cancel()
} else {
state.results[result.chunkIndex] = result
}
if state.remaining == 0 && !state.failed && ctx.Err() == nil {
if err := finalizeLaneExtract(checkpoints, state); err != nil {
state.failed = true
runErrors = append(runErrors, orderedRunError{stage: 0, lane: state.index, chunk: len(chunks), err: err})
cancel()
} else {
launch(state)
}
}
case completion := <-completions:
completed++
completedOutputs[completion.index] = completion.output
if completion.err != nil {
runErrors = append(runErrors, classifyLaneError(completion.index, len(chunks), completion.err))
cancel()
}
}
}
for i := range completedOutputs {
mergeLaneOutput(&output, completedOutputs[i])
}
if err := selectRunError(parent, runErrors); err != nil {
return output, err
}
return output, nil
}
func prepareLaneExtract(input RunInput, loader CheckpointLoader, doc *source.SourceDocument, chunks []source.Chunk, index int, prepared preparedLaneExecutor) (*laneExtractState, error) {
lane, typed := prepared.resolved, prepared.typed
digest, err := joinedChunkDigest(chunks)
if err != nil {
return nil, fmt.Errorf("digest chunks for lane %q: %w", lane.ID, err)
}
state := &laneExtractState{index: index, prepared: prepared, deps: digestFingerprints("chunks", digest), remaining: len(chunks), results: make(map[int]extractJobResult, len(chunks))}
cp, decision := loadExtract(loader, lane.ID, lane.Extract.Module, state.deps)
if decision.Reused {
for _, stored := range cp.Outputs {
if _, decodeErr := decodeCheckpointArtifact(typed.codec, stored); decodeErr != nil {
decision = CheckpointDecision{Reason: "extract artifact checkpoint codec is incompatible: " + decodeErr.Error()}
break
}
}
}
state.decision = decision
if decision.Reused {
state.remaining = 0
for _, stored := range cp.Outputs {
value, decodeErr := decodeCheckpointArtifact(typed.codec, stored)
if decodeErr != nil {
return nil, fmt.Errorf("decode extract checkpoint for lane %q: %w", lane.ID, decodeErr)
}
stored = hydrateCheckpointArtifact(typed.codec, stored, value)
artifact := erasedExtractArtifact{LaneID: lane.ID, ExtractorKey: lane.Extract.Module, SourceID: doc.ID, ChunkID: stored.ChunkID, ChunkIndex: stored.ChunkIndex, ChunkRef: stored.ChunkRef, Value: value}
if stored.ChunkIndex >= 0 && stored.ChunkIndex < len(chunks) && artifact.ChunkRef == (source.SourceRef{}) {
artifact.ChunkRef = chunks[stored.ChunkIndex].Ref
}
state.values = append(state.values, artifact)
state.serialized = append(state.serialized, cloneCheckpointArtifact(stored))
}
state.warnings, state.rejected = cloneWarnings(cp.Warnings), cloneRejectedOutputs(cp.Rejected)
}
return state, nil
}
func (r *Runner) runExtractJob(ctx context.Context, input RunInput, doc *source.SourceDocument, sourceInput contracts.LLMInputMaterial, sessionID string, job extractJob) extractJobResult {
state, chunk := job.lane, job.chunk
lane, typed := state.prepared.resolved, state.prepared.typed
result := extractJobResult{laneIndex: state.index, chunkIndex: chunk.Index}
var accepted erasedExtractArtifact
var serialized CheckpointArtifact
var acceptedWarnings []contracts.Warning
ok, rejection, err := runWithRetry(ctx, lane.Extract.Retries, func(attempt int) (bool, *contracts.RejectedOutput, error) {
started := time.Now().UTC()
attemptPath := path.Join("extract", debugPathComponent(lane.ID), fmt.Sprintf("chunk-%06d", chunk.Index+1), fmt.Sprintf("attempt-%02d", attempt))
attemptCtx, llmScope := withDebugLLMScope(ctx, attemptPath)
extracted, callErr := typed.extract(attemptCtx, typed.extractor, contracts.TypedExtractionRequest{Source: doc, Chunk: &chunk, SourceInput: chunkInputMaterial(sourceInput, chunk), SessionID: sessionID, References: CloneReferenceSet(lane.ExtractReferences.ReferenceSet), LLMProfile: lane.Extract.LLMProfile, Metadata: cloneMetadata(input.Metadata)})
if callErr != nil {
_ = writeDebugTimed(input.Debug, attemptPath+".json", debugEnvelopeWithLLMCalls(debugTimedEnvelope{Stage: string(StageExtract), LaneID: lane.ID, ModuleKey: lane.Extract.Module, Attempt: attempt, StartedAt: started, Error: callErr.Error()}, llmScope))
return false, nil, fmt.Errorf("extract lane %q chunk %q with extractor %q: %w", lane.ID, chunk.ID, lane.Extract.Module, callErr)
}
artifact := erasedExtractArtifact{LaneID: lane.ID, ExtractorKey: lane.Extract.Module, SourceID: doc.ID, ChunkID: chunk.ID, ChunkIndex: chunk.Index, ChunkRef: chunk.Ref, Value: extracted.Value}
warnings, rejected, validateErr := r.validateTypedArtifact(attemptCtx, typed.codec, typedValidationTarget{stage: StageExtract, laneID: lane.ID, moduleKey: lane.Extract.Module, source: doc, sourceID: doc.ID, sourceInput: chunkInputMaterial(sourceInput, chunk), sessionID: sessionID, references: lane.ExtractReferences.ReferenceSet, metadata: input.Metadata, chunk: &chunk, ref: chunk.Ref, value: extracted.Value}, state.prepared.extractValidators, attempt, input.Debug)
if validateErr != nil || rejected != nil {
return false, rejected, validateErr
}
stored, encodeErr := checkpointArtifact(typed.codec, artifact.LaneID, artifact.ExtractorKey, artifact.SourceID, artifact.Value)
if encodeErr != nil {
return false, nil, encodeErr
}
stored.ChunkID, stored.ChunkIndex, stored.ChunkRef = artifact.ChunkID, artifact.ChunkIndex, artifact.ChunkRef
accepted, serialized = artifact, stored
acceptedWarnings = append(cloneWarnings(extracted.Warnings), warnings...)
if debugErr := writeDebugTimed(input.Debug, attemptPath+".json", debugEnvelopeWithLLMCalls(debugTimedEnvelope{Stage: string(StageExtract), LaneID: lane.ID, ModuleKey: lane.Extract.Module, Attempt: attempt, StartedAt: started, Payload: map[string]any{"output": debugCheckpointArtifact(stored), "warnings": debugWarningEnvelopes(acceptedWarnings)}}, llmScope)); debugErr != nil {
return false, nil, debugErr
}
return true, nil, nil
})
result.err = err
if err == nil && !ok {
result.rejected = rejection
return result
}
result.value, result.serialized, result.warnings = accepted, serialized, acceptedWarnings
return result
}
func finalizeLaneExtract(checkpoints CheckpointRecorder, state *laneExtractState) error {
lane := state.prepared.resolved
indexes := make([]int, 0, len(state.results))
for index := range state.results {
indexes = append(indexes, index)
}
sort.Ints(indexes)
for _, index := range indexes {
result := state.results[index]
if result.rejected != nil {
state.rejected = append(state.rejected, *result.rejected)
continue
}
state.values = append(state.values, result.value)
state.serialized = append(state.serialized, result.serialized)
state.warnings = append(state.warnings, result.warnings...)
}
sort.SliceStable(state.values, func(i, j int) bool { return state.values[i].ChunkIndex < state.values[j].ChunkIndex })
sort.SliceStable(state.serialized, func(i, j int) bool { return state.serialized[i].ChunkIndex < state.serialized[j].ChunkIndex })
sort.SliceStable(state.rejected, func(i, j int) bool { return state.rejected[i].ChunkIndex < state.rejected[j].ChunkIndex })
if !state.decision.Reused {
if err := recordExtract(checkpoints, lane.ID, lane.Extract.Module, state.deps, state.serialized, state.rejected, state.warnings); err != nil {
return fmt.Errorf("write extract checkpoint for lane %q: %w", lane.ID, err)
}
}
return nil
}
func (r *Runner) continueLane(ctx context.Context, input RunInput, checkpoints CheckpointRecorder, loader CheckpointLoader, doc *source.SourceDocument, sourceInput contracts.LLMInputMaterial, sessionID string, chunks []source.Chunk, state *laneExtractState) (RunOutput, error) {
lane := state.prepared.resolved
local := RunOutput{Manifest: manifestFromPipeline(input)}
checkpoint := ExtractCheckpoint{Outputs: state.serialized, Rejected: state.rejected, Warnings: state.warnings}
coordinatedLoader := completedExtractLoader{CheckpointLoader: loader, laneID: lane.ID, checkpoint: checkpoint}
input.extractDecision = &state.decision
err := r.runTypedLane(ctx, input, checkpoints, coordinatedLoader, doc, sourceInput, sessionID, chunks, state.prepared, &local)
return local, err
}
func classifyLaneError(lane, sentinel int, err error) orderedRunError {
stage := 1
var laneErr *laneRunError
if errors.As(err, &laneErr) {
switch laneErr.stage {
case StageExtract:
stage = 0
case StageNormalize:
stage = 2
}
} else if strings.Contains(strings.ToLower(err.Error()), "normalize") {
stage = 2
}
return orderedRunError{stage: stage, lane: lane, chunk: sentinel, err: err}
}
func selectRunError(parent context.Context, values []orderedRunError) error {
if err := parent.Err(); err != nil {
return err
}
if len(values) == 0 {
return nil
}
hasReal := false
for _, value := range values {
if !errors.Is(value.err, context.Canceled) && !errors.Is(value.err, context.DeadlineExceeded) {
hasReal = true
break
}
}
filtered := values[:0]
for _, value := range values {
if hasReal && (errors.Is(value.err, context.Canceled) || errors.Is(value.err, context.DeadlineExceeded)) {
continue
}
filtered = append(filtered, value)
}
sort.SliceStable(filtered, func(i, j int) bool {
if filtered[i].stage != filtered[j].stage {
return filtered[i].stage < filtered[j].stage
}
if filtered[i].lane != filtered[j].lane {
return filtered[i].lane < filtered[j].lane
}
return filtered[i].chunk < filtered[j].chunk
})
return filtered[0].err
}
func mergeLaneOutput(dst *RunOutput, src RunOutput) {
if dst == nil {
return
}
dst.NormalizeOutputs = append(dst.NormalizeOutputs, cloneSerializedOutputs(src.NormalizeOutputs)...)
dst.Rejected = append(dst.Rejected, cloneRejectedOutputs(src.Rejected)...)
dst.Warnings = append(dst.Warnings, cloneWarnings(src.Warnings)...)
dst.CheckpointEvents = append(dst.CheckpointEvents, src.CheckpointEvents...)
for i := range dst.Manifest.ArtifactLanes {
for j := range src.Manifest.ArtifactLanes {
if dst.Manifest.ArtifactLanes[i].ID == src.Manifest.ArtifactLanes[j].ID && src.Manifest.ArtifactLanes[j].Metadata != nil {
dst.Manifest.ArtifactLanes[i].Metadata = cloneMetadata(src.Manifest.ArtifactLanes[j].Metadata)
}
}
}
}