Complete Phase 10 LLM validators
This commit is contained in:
260
internal/framework/validators/llm_validators.go
Normal file
260
internal/framework/validators/llm_validators.go
Normal file
@@ -0,0 +1,260 @@
|
||||
package validators
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/audita/internal/core/chunking"
|
||||
"gitea.maximumdirect.net/eric/audita/internal/core/config"
|
||||
"gitea.maximumdirect.net/eric/audita/internal/core/schema"
|
||||
"gitea.maximumdirect.net/eric/audita/internal/framework/proposals"
|
||||
)
|
||||
|
||||
type LLMPromptBuilder func(validationPayload []LLMValidationItem) ([]LLMMessage, error)
|
||||
|
||||
type LLMBackedValidator struct {
|
||||
name string
|
||||
validatorType LLMValidatorType
|
||||
promptBuilder LLMPromptBuilder
|
||||
model string
|
||||
estimator chunking.TokenEstimator
|
||||
}
|
||||
|
||||
func (v *LLMBackedValidator) Name() string {
|
||||
return v.name
|
||||
}
|
||||
|
||||
// SetTokenEstimator allows deterministic test control over batching behavior.
|
||||
func (v *LLMBackedValidator) SetTokenEstimator(estimator chunking.TokenEstimator) {
|
||||
if v == nil || estimator == nil {
|
||||
return
|
||||
}
|
||||
v.estimator = estimator
|
||||
}
|
||||
|
||||
func NewLLMBackedValidator(name string, validatorType LLMValidatorType, model string) (*LLMBackedValidator, error) {
|
||||
builder, err := promptBuilderForType(validatorType)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.TrimSpace(name) == "" {
|
||||
return nil, fmt.Errorf("validator name must not be empty")
|
||||
}
|
||||
return &LLMBackedValidator{
|
||||
name: name,
|
||||
validatorType: validatorType,
|
||||
promptBuilder: builder,
|
||||
model: strings.TrimSpace(model),
|
||||
estimator: chunking.NewSimpleTokenEstimator(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (v *LLMBackedValidator) Validate(ctx context.Context, req Request) (Result, error) {
|
||||
if v == nil {
|
||||
return Result{}, fmt.Errorf("validator is nil")
|
||||
}
|
||||
if req.LLMClient == nil {
|
||||
return Result{}, fmt.Errorf("LLM-backed validator %q requires a structured LLM client", v.name)
|
||||
}
|
||||
if len(req.CandidateProposal) == 0 {
|
||||
return Result{ValidatorName: v.name, Decisions: nil}, nil
|
||||
}
|
||||
validationReq, immediate := BuildLLMValidationRequest(v.name, v.validatorType, req)
|
||||
if len(validationReq.Items) == 0 {
|
||||
all := append([]Decision(nil), immediate...)
|
||||
sort.SliceStable(all, func(i, j int) bool { return all[i].ProposalIndex < all[j].ProposalIndex })
|
||||
if err := EnforceDecisionCardinality(req.CandidateProposal, all); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
return Result{ValidatorName: v.name, Decisions: all}, nil
|
||||
}
|
||||
|
||||
maxTokens := config.DefaultValidationMaxPromptTokens
|
||||
if req.Config != nil && req.Config.ValidationMaxPromptTokens > 0 {
|
||||
maxTokens = req.Config.ValidationMaxPromptTokens
|
||||
}
|
||||
|
||||
batches, err := ChunkLLMValidationItems(validationReq.Items, maxTokens, v.estimator)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
llmDecisions := make([]Decision, 0)
|
||||
for _, batch := range batches {
|
||||
messages, err := v.promptBuilder(batch.Items)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
var response LLMValidationResponse
|
||||
call := func(callCtx context.Context) error {
|
||||
_, err = req.LLMClient.CompleteStructured(callCtx, StructuredCompletionRequest{
|
||||
StageName: fmt.Sprintf("%s:%s:batch-%04d", req.ModuleInstance, v.name, batch.BatchIndex),
|
||||
Messages: messages,
|
||||
Model: resolvedValidationModel(req.Config, v.model),
|
||||
}, &response)
|
||||
return err
|
||||
}
|
||||
if req.Scheduler != nil {
|
||||
err = req.Scheduler.Run(ctx, call)
|
||||
} else {
|
||||
err = call(ctx)
|
||||
}
|
||||
artifacts := InteractionArtifacts{}
|
||||
if req.DiagnosticsWriter != nil {
|
||||
stage := fmt.Sprintf("%s:%s:batch-%04d", req.ModuleInstance, v.name, batch.BatchIndex)
|
||||
artifacts, _ = req.DiagnosticsWriter.WriteInteraction(
|
||||
stage,
|
||||
map[string]any{"validator_name": v.name, "validator_type": v.validatorType, "batch_index": batch.BatchIndex},
|
||||
map[string]any{"messages": messages, "items": batch.Items},
|
||||
response,
|
||||
errPayload(err),
|
||||
)
|
||||
}
|
||||
if err != nil {
|
||||
return Result{}, fmt.Errorf("LLM validator %q completion failed: %w", v.name, err)
|
||||
}
|
||||
|
||||
batchDecisions, err := mapLLMResponseToDecisions(batch.Items, response)
|
||||
if err != nil {
|
||||
return Result{}, fmt.Errorf("LLM validator %q response invalid: %w", v.name, err)
|
||||
}
|
||||
for i := range batchDecisions {
|
||||
batchDecisions[i].DiagnosticArtifactPath = artifacts.ResponsePayloadPath
|
||||
}
|
||||
llmDecisions = append(llmDecisions, batchDecisions...)
|
||||
}
|
||||
|
||||
all := append([]Decision(nil), immediate...)
|
||||
all = append(all, llmDecisions...)
|
||||
if err := EnforceDecisionCardinality(req.CandidateProposal, all); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
sort.SliceStable(all, func(i, j int) bool { return all[i].ProposalIndex < all[j].ProposalIndex })
|
||||
return Result{ValidatorName: v.name, Decisions: all}, nil
|
||||
}
|
||||
|
||||
func errPayload(err error) any {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
return map[string]any{"error": err.Error()}
|
||||
}
|
||||
|
||||
func resolvedValidationModel(cfg *config.Config, override string) string {
|
||||
if strings.TrimSpace(override) != "" {
|
||||
return strings.TrimSpace(override)
|
||||
}
|
||||
if cfg == nil {
|
||||
return ""
|
||||
}
|
||||
return cfg.EffectiveValidationLLMConfig().Model
|
||||
}
|
||||
|
||||
func BuildLLMValidationRequest(validatorName string, validatorType LLMValidatorType, req Request) (LLMValidationRequest, []Decision) {
|
||||
items := make([]LLMValidationItem, 0, len(req.CandidateProposal))
|
||||
immediate := make([]Decision, 0)
|
||||
segments := make(map[int]schema.Segment)
|
||||
if req.WorkingTranscript != nil {
|
||||
segments = make(map[int]schema.Segment, len(req.WorkingTranscript.Segments))
|
||||
for _, seg := range req.WorkingTranscript.Segments {
|
||||
segments[seg.ID] = seg
|
||||
}
|
||||
}
|
||||
|
||||
for _, p := range req.CandidateProposal {
|
||||
seg, ok := segments[p.TargetSegmentID]
|
||||
if !ok {
|
||||
immediate = append(immediate, rejection(p.ProposalIndex, ReasonMissingTargetSegment, "target segment was not found"))
|
||||
continue
|
||||
}
|
||||
|
||||
preview := proposals.PreviewProposalForSegment(&seg, p.CorrectionProposal, req.ReplacementPolicy)
|
||||
if !preview.Applicable {
|
||||
immediate = append(immediate, rejection(p.ProposalIndex, string(preview.SkipReason), "proposal is not previewable for LLM validation"))
|
||||
continue
|
||||
}
|
||||
|
||||
items = append(items, LLMValidationItem{
|
||||
CorrectionIndex: p.ProposalIndex,
|
||||
SegmentID: p.TargetSegmentID,
|
||||
OriginalText: p.OriginalText,
|
||||
CorrectedText: p.CorrectedText,
|
||||
OriginalSegmentText: seg.Text,
|
||||
CorrectedSegmentText: preview.CorrectedSegmentText,
|
||||
Categories: append([]string(nil), seg.Categories...),
|
||||
})
|
||||
}
|
||||
|
||||
return LLMValidationRequest{
|
||||
ValidatorName: validatorName,
|
||||
ValidatorType: validatorType,
|
||||
ModuleKey: req.ModuleKey,
|
||||
ModuleInstance: req.ModuleInstance,
|
||||
ReplacementPolicy: string(req.ReplacementPolicy),
|
||||
Glossary: req.Glossary,
|
||||
Items: items,
|
||||
}, immediate
|
||||
}
|
||||
|
||||
func promptBuilderForType(validatorType LLMValidatorType) (LLMPromptBuilder, error) {
|
||||
switch validatorType {
|
||||
case LLMValidatorTypeSpokenFormPlausibility:
|
||||
return BuildSpokenFormPlausibilityMessages, nil
|
||||
case LLMValidatorTypeMeaningReversal:
|
||||
return BuildMeaningReversalMessages, nil
|
||||
case LLMValidatorTypeEditorialReview:
|
||||
return BuildEditorialMessages, nil
|
||||
case LLMValidatorTypeGrammarReview:
|
||||
return BuildGrammarReviewMessages, nil
|
||||
case LLMValidatorTypeSpokenWordReview:
|
||||
return BuildSpokenWordReviewMessages, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported LLM validator type %q", validatorType)
|
||||
}
|
||||
}
|
||||
|
||||
func mapLLMResponseToDecisions(items []LLMValidationItem, response LLMValidationResponse) ([]Decision, error) {
|
||||
expected := make(map[int]LLMValidationItem, len(items))
|
||||
for _, item := range items {
|
||||
expected[item.CorrectionIndex] = item
|
||||
}
|
||||
if len(response.Validations) == 0 {
|
||||
return nil, fmt.Errorf("missing validations in structured response")
|
||||
}
|
||||
|
||||
seen := make(map[int]LLMValidationDecision, len(response.Validations))
|
||||
for _, d := range response.Validations {
|
||||
if d.Confidence < 0.0 || d.Confidence > 1.0 {
|
||||
return nil, fmt.Errorf("confidence for correction_index %d must be between 0.0 and 1.0", d.CorrectionIndex)
|
||||
}
|
||||
if _, ok := expected[d.CorrectionIndex]; !ok {
|
||||
return nil, fmt.Errorf("unknown correction_index %d", d.CorrectionIndex)
|
||||
}
|
||||
if _, exists := seen[d.CorrectionIndex]; exists {
|
||||
return nil, fmt.Errorf("duplicate correction_index %d", d.CorrectionIndex)
|
||||
}
|
||||
seen[d.CorrectionIndex] = d
|
||||
}
|
||||
|
||||
decisions := make([]Decision, 0, len(items))
|
||||
for _, item := range items {
|
||||
d, ok := seen[item.CorrectionIndex]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("missing correction_index %d", item.CorrectionIndex)
|
||||
}
|
||||
reasonCode := ReasonApproved
|
||||
if !d.Approved {
|
||||
reasonCode = "llm_rejected"
|
||||
}
|
||||
decisions = append(decisions, Decision{
|
||||
ProposalIndex: item.CorrectionIndex,
|
||||
Approved: d.Approved,
|
||||
ReasonCode: reasonCode,
|
||||
Message: strings.TrimSpace(d.Reason),
|
||||
})
|
||||
}
|
||||
return decisions, nil
|
||||
}
|
||||
Reference in New Issue
Block a user