Add feedback-aware correction contracts
This commit is contained in:
169
internal/framework/contracts/correction.go
Normal file
169
internal/framework/contracts/correction.go
Normal file
@@ -0,0 +1,169 @@
|
||||
package contracts
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
|
||||
)
|
||||
|
||||
// CorrectionProtocol identifies how a producer can represent the model output
|
||||
// that directly controlled a candidate.
|
||||
type CorrectionProtocol string
|
||||
|
||||
const (
|
||||
CorrectionProtocolSingleResponseV1 CorrectionProtocol = "single_response_v1"
|
||||
)
|
||||
|
||||
const (
|
||||
MaxValidationReasonCodeBytes = 128
|
||||
MaxValidationCorrectionGuidanceBytes = 4 * 1024
|
||||
MaxAssistantResponseBytes = 1 << 20
|
||||
MaxCorrectionGuidanceBytes = 64 * 1024
|
||||
MaxCorrectionContentBytes = MaxAssistantResponseBytes + MaxCorrectionGuidanceBytes
|
||||
)
|
||||
|
||||
func (protocol CorrectionProtocol) Validate() error {
|
||||
if protocol == CorrectionProtocolSingleResponseV1 {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("unsupported correction protocol %q", protocol)
|
||||
}
|
||||
|
||||
// SemanticCorrection carries the latest model response and application-owned
|
||||
// guidance for one fresh corrected request. Its content is sensitive and is
|
||||
// deliberately excluded from ordinary JSON serialization.
|
||||
type SemanticCorrection struct {
|
||||
AssistantResponse []byte `json:"-"`
|
||||
UserGuidance string `json:"-"`
|
||||
}
|
||||
|
||||
func NewSemanticCorrection(assistantResponse []byte, userGuidance string) (*SemanticCorrection, error) {
|
||||
correction := &SemanticCorrection{
|
||||
AssistantResponse: append([]byte(nil), assistantResponse...),
|
||||
UserGuidance: userGuidance,
|
||||
}
|
||||
if err := correction.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return correction, nil
|
||||
}
|
||||
|
||||
func CloneSemanticCorrection(correction *SemanticCorrection) (*SemanticCorrection, error) {
|
||||
if correction == nil {
|
||||
return nil, nil
|
||||
}
|
||||
return NewSemanticCorrection(correction.AssistantResponse, correction.UserGuidance)
|
||||
}
|
||||
|
||||
// CloneStructuredCompletionRequest returns a request whose mutable values are
|
||||
// owned by the caller. It is suitable for clients that retain requests after
|
||||
// CompleteStructured returns.
|
||||
func CloneStructuredCompletionRequest(request StructuredCompletionRequest) (StructuredCompletionRequest, error) {
|
||||
correction, err := CloneSemanticCorrection(request.Correction)
|
||||
if err != nil {
|
||||
return StructuredCompletionRequest{}, fmt.Errorf("clone correction: %w", err)
|
||||
}
|
||||
vars, err := source.CloneMetadata(request.Vars)
|
||||
if err != nil {
|
||||
return StructuredCompletionRequest{}, fmt.Errorf("clone variables: %w", err)
|
||||
}
|
||||
request.Inputs = request.Inputs.Clone()
|
||||
request.Vars = vars
|
||||
request.Correction = correction
|
||||
if request.StructuredOutputRepairAttempts != nil {
|
||||
attempts := *request.StructuredOutputRepairAttempts
|
||||
request.StructuredOutputRepairAttempts = &attempts
|
||||
}
|
||||
return request, nil
|
||||
}
|
||||
|
||||
func (correction SemanticCorrection) Validate() error {
|
||||
if err := validateAssistantResponse(correction.AssistantResponse); err != nil {
|
||||
return fmt.Errorf("semantic correction assistant response: %w", err)
|
||||
}
|
||||
if err := validateBoundedText(correction.UserGuidance, MaxCorrectionGuidanceBytes, "semantic correction user guidance", false); err != nil {
|
||||
return err
|
||||
}
|
||||
if len(correction.AssistantResponse)+len(correction.UserGuidance) > MaxCorrectionContentBytes {
|
||||
return errors.New("semantic correction content exceeds maximum length")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ModelCandidate preserves the exact single model response that directly
|
||||
// controlled a producer result. It is attempt-local and never durable data.
|
||||
type ModelCandidate struct {
|
||||
Response []byte `json:"-"`
|
||||
Protocol CorrectionProtocol `json:"-"`
|
||||
}
|
||||
|
||||
func NewModelCandidate(response []byte, protocol CorrectionProtocol) (*ModelCandidate, error) {
|
||||
candidate := &ModelCandidate{Response: append([]byte(nil), response...), Protocol: protocol}
|
||||
if err := candidate.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return candidate, nil
|
||||
}
|
||||
|
||||
func CloneModelCandidate(candidate *ModelCandidate) (*ModelCandidate, error) {
|
||||
if candidate == nil {
|
||||
return nil, nil
|
||||
}
|
||||
return NewModelCandidate(candidate.Response, candidate.Protocol)
|
||||
}
|
||||
|
||||
func (candidate ModelCandidate) Validate() error {
|
||||
if err := candidate.Protocol.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateAssistantResponse(candidate.Response); err != nil {
|
||||
return fmt.Errorf("model candidate response: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ValidateValidationResult(result ValidationResult) error {
|
||||
if result.ReasonCode != "" {
|
||||
if err := validateBoundedText(result.ReasonCode, MaxValidationReasonCodeBytes, "validation reason code", false); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if result.CorrectionGuidance != "" {
|
||||
if err := validateBoundedText(result.CorrectionGuidance, MaxValidationCorrectionGuidanceBytes, "validation correction guidance", false); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateAssistantResponse(response []byte) error {
|
||||
if len(response) > MaxAssistantResponseBytes {
|
||||
return errors.New("exceeds maximum length")
|
||||
}
|
||||
if !utf8.Valid(response) {
|
||||
return errors.New("must be valid UTF-8")
|
||||
}
|
||||
if len(strings.TrimSpace(string(response))) == 0 {
|
||||
return errors.New("must not be blank")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateBoundedText(value string, maximum int, name string, optional bool) error {
|
||||
if value == "" && optional {
|
||||
return nil
|
||||
}
|
||||
if !utf8.ValidString(value) {
|
||||
return fmt.Errorf("%s must be valid UTF-8", name)
|
||||
}
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fmt.Errorf("%s must not be blank", name)
|
||||
}
|
||||
if len(value) > maximum {
|
||||
return fmt.Errorf("%s exceeds maximum length", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user