Enforce bounded output repair contracts
This commit is contained in:
@@ -178,6 +178,8 @@ Stage 1 is complete when every input path accepts only a coherent zero-to-three
|
|||||||
budget, invalid contracts fail before model work, and validators no longer
|
budget, invalid contracts fail before model work, and validators no longer
|
||||||
misreport a configured budget as completed repair work.
|
misreport a configured budget as completed repair work.
|
||||||
|
|
||||||
|
**Status:** Complete.
|
||||||
|
|
||||||
## Stage 2: Distinguish Explicit Empty Content From A Malformed Envelope
|
## Stage 2: Distinguish Explicit Empty Content From A Malformed Envelope
|
||||||
|
|
||||||
### Objective
|
### Objective
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const maxOutputRepairAttempts = 3
|
||||||
|
|
||||||
// ValidateOutputContract validates source-neutral output-contract invariants.
|
// ValidateOutputContract validates source-neutral output-contract invariants.
|
||||||
func ValidateOutputContract(contract OutputContract) error {
|
func ValidateOutputContract(contract OutputContract) error {
|
||||||
switch contract.Format {
|
switch contract.Format {
|
||||||
@@ -26,5 +28,11 @@ func ValidateOutputContract(contract OutputContract) error {
|
|||||||
if contract.RepairAttempts < 0 {
|
if contract.RepairAttempts < 0 {
|
||||||
return errors.New("repair_attempts must be greater than or equal to 0")
|
return errors.New("repair_attempts must be greater than or equal to 0")
|
||||||
}
|
}
|
||||||
|
if contract.RepairAttempts > maxOutputRepairAttempts {
|
||||||
|
return fmt.Errorf("repair_attempts must be less than or equal to %d", maxOutputRepairAttempts)
|
||||||
|
}
|
||||||
|
if contract.ValidationMode == ValidationNone && contract.RepairAttempts > 0 {
|
||||||
|
return errors.New("repair_attempts requires basic, json, or json_schema validation")
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -30,9 +30,28 @@ func TestValidateOutputContract(t *testing.T) {
|
|||||||
}},
|
}},
|
||||||
{name: "empty validation mode", change: func(c *OutputContract) { c.ValidationMode = "" }, wantErr: "validation mode"},
|
{name: "empty validation mode", change: func(c *OutputContract) { c.ValidationMode = "" }, wantErr: "validation mode"},
|
||||||
{name: "unsupported validation mode", change: func(c *OutputContract) { c.ValidationMode = ValidationMode("unknown") }, wantErr: "validation mode"},
|
{name: "unsupported validation mode", change: func(c *OutputContract) { c.ValidationMode = ValidationMode("unknown") }, wantErr: "validation mode"},
|
||||||
{name: "negative repair attempts", change: func(c *OutputContract) { c.RepairAttempts = -1 }, wantErr: "repair_attempts"},
|
{name: "negative repair attempts", change: func(c *OutputContract) {
|
||||||
|
c.ValidationMode = ValidationBasic
|
||||||
|
c.RepairAttempts = -1
|
||||||
|
}, wantErr: "repair_attempts"},
|
||||||
{name: "zero repair attempts", change: func(c *OutputContract) { c.RepairAttempts = 0 }},
|
{name: "zero repair attempts", change: func(c *OutputContract) { c.RepairAttempts = 0 }},
|
||||||
{name: "positive repair attempts", change: func(c *OutputContract) { c.RepairAttempts = 1 }},
|
{name: "one repair attempt", change: func(c *OutputContract) {
|
||||||
|
c.ValidationMode = ValidationBasic
|
||||||
|
c.RepairAttempts = 1
|
||||||
|
}},
|
||||||
|
{name: "maximum repair attempts", change: func(c *OutputContract) {
|
||||||
|
c.ValidationMode = ValidationJSON
|
||||||
|
c.RepairAttempts = 3
|
||||||
|
}},
|
||||||
|
{name: "too many repair attempts", change: func(c *OutputContract) {
|
||||||
|
c.ValidationMode = ValidationJSONSchema
|
||||||
|
c.SchemaPath = "schema.json"
|
||||||
|
c.RepairAttempts = 4
|
||||||
|
}, wantErr: "repair_attempts"},
|
||||||
|
{name: "none validation with repair attempts", change: func(c *OutputContract) {
|
||||||
|
c.ValidationMode = ValidationNone
|
||||||
|
c.RepairAttempts = 1
|
||||||
|
}, wantErr: "repair_attempts"},
|
||||||
{name: "json schema empty path", change: func(c *OutputContract) {
|
{name: "json schema empty path", change: func(c *OutputContract) {
|
||||||
c.ValidationMode = ValidationJSONSchema
|
c.ValidationMode = ValidationJSONSchema
|
||||||
c.SchemaPath = ""
|
c.SchemaPath = ""
|
||||||
|
|||||||
@@ -892,6 +892,38 @@ output:
|
|||||||
format: text
|
format: text
|
||||||
validation_mode: none
|
validation_mode: none
|
||||||
repair_attempts: -1
|
repair_attempts: -1
|
||||||
|
`,
|
||||||
|
wantErr: true,
|
||||||
|
wantDiagnostic: "repair_attempts",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "repair attempts above maximum",
|
||||||
|
definition: `
|
||||||
|
id: normalization-rule
|
||||||
|
version: "1"
|
||||||
|
messages:
|
||||||
|
- role: user
|
||||||
|
content: test
|
||||||
|
output:
|
||||||
|
format: text
|
||||||
|
validation_mode: basic
|
||||||
|
repair_attempts: 4
|
||||||
|
`,
|
||||||
|
wantErr: true,
|
||||||
|
wantDiagnostic: "repair_attempts",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "none validation with repair attempts",
|
||||||
|
definition: `
|
||||||
|
id: normalization-rule
|
||||||
|
version: "1"
|
||||||
|
messages:
|
||||||
|
- role: user
|
||||||
|
content: test
|
||||||
|
output:
|
||||||
|
format: text
|
||||||
|
validation_mode: none
|
||||||
|
repair_attempts: 1
|
||||||
`,
|
`,
|
||||||
wantErr: true,
|
wantErr: true,
|
||||||
wantDiagnostic: "repair_attempts",
|
wantDiagnostic: "repair_attempts",
|
||||||
|
|||||||
@@ -106,6 +106,8 @@ func TestRunnerPreparationRejectsInvalidOutputContractsBeforeCompletion(t *testi
|
|||||||
{name: "unsupported format", override: domain.OutputContract{Format: "binary", ValidationMode: domain.ValidationNone}},
|
{name: "unsupported format", override: domain.OutputContract{Format: "binary", ValidationMode: domain.ValidationNone}},
|
||||||
{name: "empty validation mode", override: domain.OutputContract{Format: domain.FormatText}},
|
{name: "empty validation mode", override: domain.OutputContract{Format: domain.FormatText}},
|
||||||
{name: "negative repair attempts", override: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone, RepairAttempts: -1}},
|
{name: "negative repair attempts", override: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone, RepairAttempts: -1}},
|
||||||
|
{name: "repair attempts above maximum", override: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationBasic, RepairAttempts: 4}},
|
||||||
|
{name: "none validation with repair attempts", override: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone, RepairAttempts: 1}},
|
||||||
{name: "json schema without path", override: domain.OutputContract{Format: domain.FormatJSON, ValidationMode: domain.ValidationJSONSchema}},
|
{name: "json schema without path", override: domain.OutputContract{Format: domain.FormatJSON, ValidationMode: domain.ValidationJSONSchema}},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -2202,7 +2202,7 @@ func TestRunnerRepairStateMachine(t *testing.T) {
|
|||||||
{
|
{
|
||||||
name: "successful repair stops below larger budget",
|
name: "successful repair stops below larger budget",
|
||||||
mode: domain.ValidationJSON,
|
mode: domain.ValidationJSON,
|
||||||
budget: 4,
|
budget: 3,
|
||||||
validationResults: []domain.ValidationResult{
|
validationResults: []domain.ValidationResult{
|
||||||
failed(domain.ValidationJSON, "candidate zero"),
|
failed(domain.ValidationJSON, "candidate zero"),
|
||||||
failed(domain.ValidationJSON, "candidate one"),
|
failed(domain.ValidationJSON, "candidate one"),
|
||||||
|
|||||||
@@ -174,9 +174,8 @@ func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract d
|
|||||||
}
|
}
|
||||||
|
|
||||||
res := domain.ValidationResult{
|
res := domain.ValidationResult{
|
||||||
Mode: contract.ValidationMode,
|
Mode: contract.ValidationMode,
|
||||||
SchemaPath: contract.SchemaPath,
|
SchemaPath: contract.SchemaPath,
|
||||||
RepairAttempts: contract.RepairAttempts,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if artifact == nil {
|
if artifact == nil {
|
||||||
|
|||||||
@@ -64,6 +64,43 @@ func TestStandardValidatorBasicFailureEmpty(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestStandardValidatorReportsZeroRepairAttempts(t *testing.T) {
|
||||||
|
v := NewStandardValidator("")
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
body string
|
||||||
|
contract domain.OutputContract
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "passed basic validation",
|
||||||
|
body: "answer",
|
||||||
|
contract: domain.OutputContract{
|
||||||
|
ValidationMode: domain.ValidationBasic,
|
||||||
|
RepairAttempts: 3,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "failed JSON validation",
|
||||||
|
body: `{"answer":`,
|
||||||
|
contract: domain.OutputContract{
|
||||||
|
ValidationMode: domain.ValidationJSON,
|
||||||
|
RepairAttempts: 3,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
res, err := v.Validate(context.Background(), &domain.Artifact{Body: []byte(tc.body)}, tc.contract)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("validate: %v", err)
|
||||||
|
}
|
||||||
|
if res.RepairAttempts != 0 {
|
||||||
|
t.Fatalf("repair attempts = %d, want 0", res.RepairAttempts)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestStandardValidatorJSONSuccess(t *testing.T) {
|
func TestStandardValidatorJSONSuccess(t *testing.T) {
|
||||||
v := NewStandardValidator("")
|
v := NewStandardValidator("")
|
||||||
|
|
||||||
|
|||||||
@@ -22,27 +22,52 @@ func TestPreparationRejectsInvalidOutputContractWithPublicError(t *testing.T) {
|
|||||||
t.Fatalf("construct engine: %v", err)
|
t.Fatalf("construct engine: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
req := promptkit.RunRequest{
|
for _, tc := range []struct {
|
||||||
PromptID: "prompt",
|
name string
|
||||||
Validation: &promptkit.OutputContract{
|
contract promptkit.OutputContract
|
||||||
Format: promptkit.OutputFormat("binary"),
|
}{
|
||||||
ValidationMode: promptkit.ValidationNone,
|
{
|
||||||
|
name: "unsupported format",
|
||||||
|
contract: promptkit.OutputContract{
|
||||||
|
Format: promptkit.OutputFormat("binary"),
|
||||||
|
ValidationMode: promptkit.ValidationNone,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
}
|
{
|
||||||
|
name: "repair attempts above maximum",
|
||||||
|
contract: promptkit.OutputContract{
|
||||||
|
Format: promptkit.FormatText,
|
||||||
|
ValidationMode: promptkit.ValidationBasic,
|
||||||
|
RepairAttempts: 4,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "none validation with repair attempts",
|
||||||
|
contract: promptkit.OutputContract{
|
||||||
|
Format: promptkit.FormatText,
|
||||||
|
ValidationMode: promptkit.ValidationNone,
|
||||||
|
RepairAttempts: 1,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
req := promptkit.RunRequest{PromptID: "prompt", Validation: &tc.contract}
|
||||||
|
|
||||||
prepared, err := engine.Prepare(context.Background(), req)
|
prepared, err := engine.Prepare(context.Background(), req)
|
||||||
if prepared != nil {
|
if prepared != nil {
|
||||||
t.Fatalf("expected no partial prepared run, got %+v", prepared)
|
t.Fatalf("expected no partial prepared run, got %+v", prepared)
|
||||||
}
|
}
|
||||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||||
t.Fatalf("prepare error = %v, want ErrInvalidRequest", err)
|
t.Fatalf("prepare error = %v, want ErrInvalidRequest", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
preparedExecution, err := engine.PrepareExecution(context.Background(), req)
|
preparedExecution, err := engine.PrepareExecution(context.Background(), req)
|
||||||
if preparedExecution != nil {
|
if preparedExecution != nil {
|
||||||
t.Fatalf("expected no partial prepared execution, got %+v", preparedExecution)
|
t.Fatalf("expected no partial prepared execution, got %+v", preparedExecution)
|
||||||
}
|
}
|
||||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||||
t.Fatalf("prepare execution error = %v, want ErrInvalidRequest", err)
|
t.Fatalf("prepare execution error = %v, want ErrInvalidRequest", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user