Add feedback-aware correction contracts
This commit is contained in:
@@ -691,6 +691,7 @@ func debugResponseModel(response contracts.StructuredCompletionResponse) string
|
||||
|
||||
func debugValidationResultEnvelope(result contracts.ValidationResult) contracts.ValidationResult {
|
||||
result.Message = string(redactSecretBytes([]byte(result.Message)))
|
||||
result.CorrectionGuidance = ""
|
||||
result.DiagnosticArtifactPath = string(redactSecretBytes([]byte(result.DiagnosticArtifactPath)))
|
||||
for i := range result.Warnings {
|
||||
result.Warnings[i].Message = string(redactSecretBytes([]byte(result.Warnings[i].Message)))
|
||||
|
||||
@@ -58,11 +58,20 @@ func RegisterExtractorBuilder[T any](registry *ExtractorRegistry, spec ModuleSpe
|
||||
if !ok {
|
||||
return erasedTypedResult{}, fmt.Errorf("extractor %q has incompatible implementation %T", normalized.Key, implementation)
|
||||
}
|
||||
correction, err := contracts.CloneSemanticCorrection(request.Correction)
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, fmt.Errorf("clone extraction correction: %w", err)
|
||||
}
|
||||
request.Correction = correction
|
||||
result, err := extractor.Extract(ctx, request)
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, err
|
||||
}
|
||||
return erasedTypedResult{Value: result.Value, Warnings: result.Warnings}, nil
|
||||
candidate, err := contracts.CloneModelCandidate(result.ModelCandidate)
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, fmt.Errorf("clone extraction model candidate: %w", err)
|
||||
}
|
||||
return erasedTypedResult{Value: result.Value, Warnings: cloneWarnings(result.Warnings), ModelCandidate: candidate}, nil
|
||||
}}
|
||||
if registry.typedEntries == nil {
|
||||
registry.typedEntries = map[string]typedExtractorEntry{}
|
||||
|
||||
@@ -85,11 +85,19 @@ func RegisterMergerBuilder[T any](registry *MergerRegistry, spec ModuleSpec, val
|
||||
}
|
||||
outputs[i] = contracts.ExtractArtifact[T]{LaneID: output.LaneID, ExtractorKey: output.ExtractorKey, SourceID: output.SourceID, ChunkID: output.ChunkID, ChunkIndex: output.ChunkIndex, ChunkRef: output.ChunkRef, Value: value}
|
||||
}
|
||||
result, err := merger.Merge(ctx, contracts.TypedMergeRequest[T]{Source: request.Source, LaneID: request.LaneID, ExtractOutputs: outputs, SourceInput: request.SourceInput, SessionID: request.SessionID, References: request.References, LLMProfile: request.LLMProfile, StructuredOutputRepairAttempts: request.StructuredOutputRepairAttempts, Metadata: request.Metadata})
|
||||
correction, err := contracts.CloneSemanticCorrection(request.Correction)
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, fmt.Errorf("clone merge correction: %w", err)
|
||||
}
|
||||
result, err := merger.Merge(ctx, contracts.TypedMergeRequest[T]{Source: request.Source, LaneID: request.LaneID, ExtractOutputs: outputs, SourceInput: request.SourceInput, SessionID: request.SessionID, References: request.References, LLMProfile: request.LLMProfile, StructuredOutputRepairAttempts: request.StructuredOutputRepairAttempts, Correction: correction, Metadata: request.Metadata})
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, err
|
||||
}
|
||||
return erasedTypedResult{Value: result.Value, Warnings: result.Warnings}, nil
|
||||
candidate, err := contracts.CloneModelCandidate(result.ModelCandidate)
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, fmt.Errorf("clone merge model candidate: %w", err)
|
||||
}
|
||||
return erasedTypedResult{Value: result.Value, Warnings: cloneWarnings(result.Warnings), ModelCandidate: candidate}, nil
|
||||
},
|
||||
}
|
||||
return nil
|
||||
|
||||
@@ -76,11 +76,19 @@ func RegisterNormalizerBuilder[T any](registry *NormalizerRegistry, spec ModuleS
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, err
|
||||
}
|
||||
result, err := normalizer.Normalize(ctx, contracts.TypedNormalizeRequest[T]{Source: request.Source, LaneID: request.LaneID, MergeOutput: contracts.MergeArtifact[T]{LaneID: request.MergeOutput.LaneID, MergerKey: request.MergeOutput.MergerKey, SourceID: request.MergeOutput.SourceID, Value: value}, SourceInput: request.SourceInput, SessionID: request.SessionID, References: request.References, LLMProfile: request.LLMProfile, StructuredOutputRepairAttempts: request.StructuredOutputRepairAttempts, Metadata: request.Metadata})
|
||||
correction, err := contracts.CloneSemanticCorrection(request.Correction)
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, fmt.Errorf("clone normalize correction: %w", err)
|
||||
}
|
||||
result, err := normalizer.Normalize(ctx, contracts.TypedNormalizeRequest[T]{Source: request.Source, LaneID: request.LaneID, MergeOutput: contracts.MergeArtifact[T]{LaneID: request.MergeOutput.LaneID, MergerKey: request.MergeOutput.MergerKey, SourceID: request.MergeOutput.SourceID, Value: value}, SourceInput: request.SourceInput, SessionID: request.SessionID, References: request.References, LLMProfile: request.LLMProfile, StructuredOutputRepairAttempts: request.StructuredOutputRepairAttempts, Correction: correction, Metadata: request.Metadata})
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, err
|
||||
}
|
||||
return erasedTypedResult{Value: result.Value, Warnings: cloneWarnings(result.Warnings), Retry: cloneNormalizeRetry(result.Retry)}, nil
|
||||
candidate, err := contracts.CloneModelCandidate(result.ModelCandidate)
|
||||
if err != nil {
|
||||
return erasedTypedResult{}, fmt.Errorf("clone normalize model candidate: %w", err)
|
||||
}
|
||||
return erasedTypedResult{Value: result.Value, Warnings: cloneWarnings(result.Warnings), Retry: cloneNormalizeRetry(result.Retry), ModelCandidate: candidate}, nil
|
||||
},
|
||||
}
|
||||
return nil
|
||||
|
||||
@@ -499,6 +499,9 @@ func (r *Runner) validateChunks(ctx context.Context, doc *source.SourceDocument,
|
||||
default:
|
||||
return nil, nil, fmt.Errorf("validator %q is incompatible with chunk validation", binding.Module)
|
||||
}
|
||||
if err == nil {
|
||||
err = contracts.ValidateValidationResult(result)
|
||||
}
|
||||
debugContent := debugContentEnvelope(content, "application/json", nil, nil)
|
||||
debugContent.ContentDigest = debugContentDigest(content)
|
||||
debugCall := debugValidationCall{ValidatorName: binding.Module, Request: map[string]any{"stage": string(StageChunk), "module_key": moduleKey, "source_id": doc.ID, "schema": schema, "schema_digest": contracts.DigestArtifactSchema(schema), "content": debugContent, "metadata": redactSensitiveMap(metadata)}, Result: debugValidationResultEnvelope(result)}
|
||||
|
||||
@@ -523,6 +523,9 @@ func (r *Runner) validateTypedArtifact(ctx context.Context, codec artifactCodecE
|
||||
default:
|
||||
return nil, nil, fmt.Errorf("validator %q is incompatible with typed artifact validation", binding.Module)
|
||||
}
|
||||
if err == nil {
|
||||
err = contracts.ValidateValidationResult(result)
|
||||
}
|
||||
artifact, _ := validationCandidateArtifact(codec, target)
|
||||
debugCall := debugValidationCall{ValidatorName: binding.Module, Request: map[string]any{"stage": string(target.stage), "lane_id": target.laneID, "module_key": target.moduleKey, "source_id": target.sourceID, "artifact": debugCheckpointArtifact(artifact), "metadata": redactSensitiveMap(target.metadata)}, Result: debugValidationResultEnvelope(result)}
|
||||
if err != nil {
|
||||
|
||||
@@ -22,9 +22,10 @@ type erasedMergeArtifact struct {
|
||||
}
|
||||
|
||||
type erasedTypedResult struct {
|
||||
Value any
|
||||
Warnings []contracts.Warning
|
||||
Retry *contracts.NormalizeRetry
|
||||
Value any
|
||||
Warnings []contracts.Warning
|
||||
Retry *contracts.NormalizeRetry
|
||||
ModelCandidate *contracts.ModelCandidate
|
||||
}
|
||||
|
||||
type typedValidationTarget struct {
|
||||
|
||||
@@ -1,12 +1,70 @@
|
||||
package pipeline
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
|
||||
)
|
||||
|
||||
type correctionObservingNotesExtractor struct {
|
||||
correction *contracts.SemanticCorrection
|
||||
candidate *contracts.ModelCandidate
|
||||
}
|
||||
|
||||
func (*correctionObservingNotesExtractor) Key() string { return "test/correction-observing-extract" }
|
||||
func (*correctionObservingNotesExtractor) ReferenceSlots() []contracts.ReferenceSlot {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *correctionObservingNotesExtractor) Extract(_ context.Context, request contracts.TypedExtractionRequest) (contracts.TypedExtractionResult[codecNotes], error) {
|
||||
e.correction = request.Correction
|
||||
candidate, err := contracts.NewModelCandidate([]byte(`{"items":["extract"]}`), contracts.CorrectionProtocolSingleResponseV1)
|
||||
if err != nil {
|
||||
return contracts.TypedExtractionResult[codecNotes]{}, err
|
||||
}
|
||||
e.candidate = candidate
|
||||
return contracts.TypedExtractionResult[codecNotes]{Value: codecNotes{Items: []string{"extract"}}, ModelCandidate: candidate}, nil
|
||||
}
|
||||
|
||||
type correctionObservingNotesMerger struct {
|
||||
correction *contracts.SemanticCorrection
|
||||
candidate *contracts.ModelCandidate
|
||||
}
|
||||
|
||||
func (*correctionObservingNotesMerger) Key() string { return "test/correction-observing-merge" }
|
||||
|
||||
func (m *correctionObservingNotesMerger) Merge(_ context.Context, request contracts.TypedMergeRequest[codecNotes]) (contracts.TypedMergeResult[codecNotes], error) {
|
||||
m.correction = request.Correction
|
||||
candidate, err := contracts.NewModelCandidate([]byte(`{"items":["merge"]}`), contracts.CorrectionProtocolSingleResponseV1)
|
||||
if err != nil {
|
||||
return contracts.TypedMergeResult[codecNotes]{}, err
|
||||
}
|
||||
m.candidate = candidate
|
||||
return contracts.TypedMergeResult[codecNotes]{Value: codecNotes{Items: []string{"merge"}}, ModelCandidate: candidate}, nil
|
||||
}
|
||||
|
||||
type correctionObservingNotesNormalizer struct {
|
||||
correction *contracts.SemanticCorrection
|
||||
candidate *contracts.ModelCandidate
|
||||
}
|
||||
|
||||
func (*correctionObservingNotesNormalizer) Key() string { return "test/correction-observing-normalize" }
|
||||
func (*correctionObservingNotesNormalizer) ReferenceSlots() []contracts.ReferenceSlot {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (n *correctionObservingNotesNormalizer) Normalize(_ context.Context, request contracts.TypedNormalizeRequest[codecNotes]) (contracts.TypedNormalizeResult[codecNotes], error) {
|
||||
n.correction = request.Correction
|
||||
candidate, err := contracts.NewModelCandidate([]byte(`{"items":["normalize"]}`), contracts.CorrectionProtocolSingleResponseV1)
|
||||
if err != nil {
|
||||
return contracts.TypedNormalizeResult[codecNotes]{}, err
|
||||
}
|
||||
n.candidate = candidate
|
||||
return contracts.TypedNormalizeResult[codecNotes]{Value: codecNotes{Items: []string{"normalize"}}, ModelCandidate: candidate}, nil
|
||||
}
|
||||
|
||||
type repairObservingNotesMerger struct {
|
||||
attempts *int
|
||||
}
|
||||
@@ -93,3 +151,103 @@ func TestTypedRegistryErasurePreservesStructuredOutputRepairAttempts(t *testing.
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestTypedRegistryErasurePreservesCorrectionAndCandidateOwnership(t *testing.T) {
|
||||
const artifactKind contracts.ArtifactKind = "test/notes"
|
||||
|
||||
t.Run("extract", func(t *testing.T) {
|
||||
correction := newTestCorrection(t)
|
||||
implementation := &correctionObservingNotesExtractor{}
|
||||
registry := NewExtractorRegistry()
|
||||
if err := RegisterExtractor(registry, ModuleSpec{Key: implementation.Key(), Stage: StageExtract, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: artifactKind}, func() (contracts.Extractor[codecNotes], error) {
|
||||
return implementation, nil
|
||||
}); err != nil {
|
||||
t.Fatalf("RegisterExtractor() error = %v", err)
|
||||
}
|
||||
entry, ok := registry.typedEntry(implementation.Key())
|
||||
if !ok {
|
||||
t.Fatal("typed extractor entry missing")
|
||||
}
|
||||
built, err := entry.builder(BuildRequest{})
|
||||
if err != nil {
|
||||
t.Fatalf("builder() error = %v", err)
|
||||
}
|
||||
result, err := entry.extract(context.Background(), built, contracts.TypedExtractionRequest{Correction: correction})
|
||||
if err != nil {
|
||||
t.Fatalf("extract() error = %v", err)
|
||||
}
|
||||
assertCorrectionAndCandidateOwnership(t, correction, implementation.correction, implementation.candidate, result.ModelCandidate)
|
||||
})
|
||||
|
||||
t.Run("merge", func(t *testing.T) {
|
||||
correction := newTestCorrection(t)
|
||||
implementation := &correctionObservingNotesMerger{}
|
||||
registry := NewMergerRegistry()
|
||||
if err := RegisterMerger(registry, ModuleSpec{Key: implementation.Key(), Stage: StageMerge, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: artifactKind}, func() (contracts.Merger[codecNotes], error) {
|
||||
return implementation, nil
|
||||
}); err != nil {
|
||||
t.Fatalf("RegisterMerger() error = %v", err)
|
||||
}
|
||||
entry, ok := registry.typedEntry(implementation.Key(), artifactKind)
|
||||
if !ok {
|
||||
t.Fatal("typed merger entry missing")
|
||||
}
|
||||
built, err := entry.builder(BuildRequest{})
|
||||
if err != nil {
|
||||
t.Fatalf("builder() error = %v", err)
|
||||
}
|
||||
result, err := entry.merge(context.Background(), built, contracts.TypedMergeRequest[any]{Correction: correction, ExtractOutputs: []contracts.ExtractArtifact[any]{{Value: codecNotes{Items: []string{"extract"}}}}})
|
||||
if err != nil {
|
||||
t.Fatalf("merge() error = %v", err)
|
||||
}
|
||||
assertCorrectionAndCandidateOwnership(t, correction, implementation.correction, implementation.candidate, result.ModelCandidate)
|
||||
})
|
||||
|
||||
t.Run("normalize", func(t *testing.T) {
|
||||
correction := newTestCorrection(t)
|
||||
implementation := &correctionObservingNotesNormalizer{}
|
||||
registry := NewNormalizerRegistry()
|
||||
if err := RegisterNormalizer(registry, ModuleSpec{Key: implementation.Key(), Stage: StageNormalize, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: artifactKind}, func() (contracts.Normalizer[codecNotes], error) {
|
||||
return implementation, nil
|
||||
}); err != nil {
|
||||
t.Fatalf("RegisterNormalizer() error = %v", err)
|
||||
}
|
||||
entry, ok := registry.typedEntry(implementation.Key(), artifactKind)
|
||||
if !ok {
|
||||
t.Fatal("typed normalizer entry missing")
|
||||
}
|
||||
built, err := entry.builder(BuildRequest{})
|
||||
if err != nil {
|
||||
t.Fatalf("builder() error = %v", err)
|
||||
}
|
||||
result, err := entry.normalize(context.Background(), built, contracts.TypedNormalizeRequest[any]{Correction: correction, MergeOutput: contracts.MergeArtifact[any]{Value: codecNotes{Items: []string{"merge"}}}})
|
||||
if err != nil {
|
||||
t.Fatalf("normalize() error = %v", err)
|
||||
}
|
||||
assertCorrectionAndCandidateOwnership(t, correction, implementation.correction, implementation.candidate, result.ModelCandidate)
|
||||
})
|
||||
}
|
||||
|
||||
func newTestCorrection(t *testing.T) *contracts.SemanticCorrection {
|
||||
t.Helper()
|
||||
correction, err := contracts.NewSemanticCorrection([]byte(`{"items":["original"]}`), "Correct the response.")
|
||||
if err != nil {
|
||||
t.Fatalf("NewSemanticCorrection() error = %v", err)
|
||||
}
|
||||
return correction
|
||||
}
|
||||
|
||||
func assertCorrectionAndCandidateOwnership(t *testing.T, callerCorrection, observedCorrection *contracts.SemanticCorrection, producerCandidate, returnedCandidate *contracts.ModelCandidate) {
|
||||
t.Helper()
|
||||
if observedCorrection == nil || producerCandidate == nil || returnedCandidate == nil {
|
||||
t.Fatal("correction and candidates must be present")
|
||||
}
|
||||
callerCorrection.AssistantResponse[0] = '['
|
||||
if got := string(observedCorrection.AssistantResponse); got != `{"items":["original"]}` {
|
||||
t.Fatalf("observed correction = %q, want owned original content", got)
|
||||
}
|
||||
producerCandidate.Response[0] = '['
|
||||
if bytes.Equal(producerCandidate.Response, returnedCandidate.Response) {
|
||||
t.Fatalf("returned candidate aliases producer candidate: %q", returnedCandidate.Response)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user