Files
promptkit/internal/usecase/output_contract_test.go

186 lines
6.1 KiB
Go

package usecase
import (
"context"
"errors"
"testing"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
)
type outputContractTestCollaborators struct {
artifacts *fakeArtifactReader
renderer *fakeRenderer
llm *fakeLLM
validator *recordingValidationPreparer
admitter *fakeRunAdmitter
}
func newOutputContractTestRunner() (*Runner, outputContractTestCollaborators) {
collaborators := outputContractTestCollaborators{
artifacts: defaultArtifactReader(),
renderer: defaultRenderer(),
llm: &fakeLLM{forbid: true},
validator: &recordingValidationPreparer{plan: &recordingPreparedValidation{}},
admitter: &fakeRunAdmitter{},
}
return NewRunner(
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
nil,
collaborators.artifacts,
collaborators.renderer,
collaborators.llm,
collaborators.validator,
collaborators.admitter,
), collaborators
}
func TestRunnerPreparationNormalizesOutputContractConsistently(t *testing.T) {
tests := []struct {
name string
override domain.OutputContract
want domain.OutputContract
}{
{
name: "empty replacement format defaults to text",
override: domain.OutputContract{ValidationMode: domain.ValidationNone},
want: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone},
},
{
name: "markdown basic replacement",
override: domain.OutputContract{Format: domain.FormatMarkdown, ValidationMode: domain.ValidationBasic},
want: domain.OutputContract{Format: domain.FormatMarkdown, ValidationMode: domain.ValidationBasic},
},
{
name: "json replacement preserves non-schema fields",
override: domain.OutputContract{
Format: domain.FormatJSON,
ValidationMode: domain.ValidationJSON,
SchemaPath: "ignored.json",
RepairAttempts: 2,
},
want: domain.OutputContract{
Format: domain.FormatJSON,
ValidationMode: domain.ValidationJSON,
SchemaPath: "ignored.json",
RepairAttempts: 2,
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
runner, _ := newOutputContractTestRunner()
req := domain.RunRequest{
PromptID: "p",
ProfileID: "exec",
Inputs: singleInputRef(),
Validation: &tt.override,
}
prepared, err := runner.Prepare(context.Background(), req)
if err != nil {
t.Fatalf("prepare: %v", err)
}
preparedExecution, err := runner.PrepareExecution(context.Background(), req)
if err != nil {
t.Fatalf("prepare execution: %v", err)
}
details := preparedExecution.Details()
if details == nil {
t.Fatal("prepared execution returned nil details")
}
if prepared.OutputContract != tt.want || details.OutputContract != tt.want {
t.Fatalf("output contracts = (%+v, %+v), want %+v", prepared.OutputContract, details.OutputContract, tt.want)
}
})
}
}
func TestRunnerPreparationRejectsInvalidOutputContractsBeforeCompletion(t *testing.T) {
tests := []struct {
name string
override domain.OutputContract
}{
{name: "unsupported format", override: domain.OutputContract{Format: "binary", ValidationMode: domain.ValidationNone}},
{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: "json schema without path", override: domain.OutputContract{Format: domain.FormatJSON, ValidationMode: domain.ValidationJSONSchema}},
}
for _, tt := range tests {
for _, operation := range []string{"Prepare", "PrepareExecution"} {
t.Run(tt.name+"/"+operation, func(t *testing.T) {
runner, collaborators := newOutputContractTestRunner()
req := domain.RunRequest{
PromptID: "p",
ProfileID: "exec",
Inputs: singleInputRef(),
Validation: &tt.override,
}
var err error
switch operation {
case "Prepare":
var prepared *domain.PreparedRun
prepared, err = runner.Prepare(context.Background(), req)
if prepared != nil {
t.Fatalf("expected no partial prepared run, got %+v", prepared)
}
case "PrepareExecution":
var prepared *PreparedExecution
prepared, err = runner.PrepareExecution(context.Background(), req)
if prepared != nil {
t.Fatalf("expected no partial prepared execution, got %+v", prepared)
}
default:
t.Fatalf("unknown operation %q", operation)
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
assertOutputContractCompletionSkipped(t, collaborators)
})
}
}
}
func TestRunnerRunRejectsInvalidOutputContractBeforeAdmission(t *testing.T) {
runner, collaborators := newOutputContractTestRunner()
result, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p",
ProfileID: "exec",
Inputs: singleInputRef(),
Validation: &domain.OutputContract{
Format: domain.OutputFormat("binary"),
ValidationMode: domain.ValidationNone,
},
})
if result != nil {
t.Fatalf("expected no partial result, got %+v", result)
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
assertOutputContractCompletionSkipped(t, collaborators)
}
func assertOutputContractCompletionSkipped(t *testing.T, collaborators outputContractTestCollaborators) {
t.Helper()
if collaborators.artifacts.calls != 0 || collaborators.renderer.calls != 0 ||
collaborators.validator.prepareCalls != 0 || collaborators.validator.directValidateCalls != 0 ||
len(collaborators.admitter.backendIDs) != 0 || collaborators.llm.calls != 0 {
t.Fatalf(
"invalid output contract reached downstream work: artifacts=%d renderer=%d prepare_validation=%d validation=%d admissions=%d generation=%d",
collaborators.artifacts.calls,
collaborators.renderer.calls,
collaborators.validator.prepareCalls,
collaborators.validator.directValidateCalls,
len(collaborators.admitter.backendIDs),
collaborators.llm.calls,
)
}
}