96 lines
3.2 KiB
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."),
|
|
},
|
|
}
|
|
}
|