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."), }, } }