Files
promptkit/generation_error_contract_test.go

96 lines
3.2 KiB
Go

package promptkit_test
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"strings"
"testing"
"gitea.maximumdirect.net/eric/promptkit"
)
func TestBuiltInGenerationError(t *testing.T) {
const (
codeMarker = "provider-code-marker"
typeMarker = "provider-type-marker"
messageMarker = "provider-message-marker"
)
engine := newBuiltInGenerationErrorEngine(t, http.StatusUnprocessableEntity,
`{"error":{"code":"`+codeMarker+`","type":"`+typeMarker+`","message":"`+messageMarker+`"}}`)
result, err := engine.Run(context.Background(), generationErrorRunRequest())
if result != nil {
t.Fatalf("Run result = %#v, want nil", result)
}
assertGenerationError(t, err, http.StatusUnprocessableEntity, codeMarker, typeMarker, messageMarker)
preparedEngine := newBuiltInGenerationErrorEngine(t, http.StatusServiceUnavailable, `{"error":{}}`)
prepared, err := preparedEngine.PrepareExecution(context.Background(), generationErrorRunRequest())
if err != nil {
t.Fatalf("PrepareExecution: %v", err)
}
result, err = preparedEngine.RunPrepared(context.Background(), prepared)
if result != nil {
t.Fatalf("RunPrepared result = %#v, want nil", result)
}
assertGenerationError(t, err, http.StatusServiceUnavailable, "", "", "")
}
func assertGenerationError(t *testing.T, err error, statusCode int, code, providerType, message string) {
t.Helper()
if !errors.Is(err, promptkit.ErrLLMGenerate) {
t.Fatalf("errors.Is(%v, ErrLLMGenerate) = false", err)
}
var generationErr *promptkit.GenerationError
if !errors.As(err, &generationErr) || generationErr == nil {
t.Fatalf("error = %T, want *GenerationError", err)
}
if generationErr.StatusCode() != statusCode || generationErr.ProviderCode() != code || generationErr.ProviderType() != providerType || generationErr.ProviderMessage() != message {
t.Fatalf("GenerationError = %#v", generationErr)
}
wantFormatted := fmt.Sprintf("failed to generate output: provider returned HTTP status %d", statusCode)
for _, rendered := range []string{fmt.Sprintf("%v", generationErr), fmt.Sprintf("%+v", generationErr), fmt.Sprintf("%#v", generationErr)} {
if rendered != wantFormatted {
t.Fatalf("formatted error = %q, want %q", rendered, wantFormatted)
}
for _, marker := range []string{code, providerType, message} {
if marker != "" && strings.Contains(rendered, marker) {
t.Fatalf("formatted error exposed provider marker %q: %q", marker, rendered)
}
}
}
}
func newBuiltInGenerationErrorEngine(t *testing.T, statusCode int, body string) *promptkit.Engine {
t.Helper()
config := contractConfig(frameworkSchemaDir)
config.HTTPClient = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: statusCode,
ContentLength: int64(len(body)),
Body: io.NopCloser(strings.NewReader(body)),
}, nil
})}
engine, err := promptkit.NewEngine(config)
if err != nil {
t.Fatalf("NewEngine: %v", err)
}
return engine
}
func generationErrorRunRequest() promptkit.RunRequest {
return promptkit.RunRequest{
PromptID: frameworkMarkdownSummaryPromptID,
Inputs: map[string]promptkit.ArtifactRef{
"transcript": promptkit.Inline("Rin opens the gate."),
"glossary": promptkit.Inline("gate: A guarded passage."),
},
}
}