|
|
|
|
@@ -4,11 +4,15 @@ import (
|
|
|
|
|
"context"
|
|
|
|
|
"errors"
|
|
|
|
|
"fmt"
|
|
|
|
|
"net/http"
|
|
|
|
|
"net/http/httptest"
|
|
|
|
|
"os"
|
|
|
|
|
"path/filepath"
|
|
|
|
|
"reflect"
|
|
|
|
|
"strings"
|
|
|
|
|
"sync"
|
|
|
|
|
"testing"
|
|
|
|
|
"testing/fstest"
|
|
|
|
|
"time"
|
|
|
|
|
|
|
|
|
|
promptkit "gitea.maximumdirect.net/eric/promptkit"
|
|
|
|
|
@@ -20,12 +24,19 @@ type fakeClient struct {
|
|
|
|
|
mu sync.Mutex
|
|
|
|
|
response *promptkit.GenerateResponse
|
|
|
|
|
err error
|
|
|
|
|
outcomes []generationOutcome
|
|
|
|
|
next int
|
|
|
|
|
calls int
|
|
|
|
|
requests []promptkit.GenerateRequest
|
|
|
|
|
block bool
|
|
|
|
|
started chan struct{}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
type generationOutcome struct {
|
|
|
|
|
response *promptkit.GenerateResponse
|
|
|
|
|
err error
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
type recordingReader struct {
|
|
|
|
|
ref promptkit.ArtifactRef
|
|
|
|
|
}
|
|
|
|
|
@@ -49,6 +60,11 @@ func (client *fakeClient) Generate(ctx context.Context, request promptkit.Genera
|
|
|
|
|
started := client.started
|
|
|
|
|
response := client.response
|
|
|
|
|
err := client.err
|
|
|
|
|
if client.next < len(client.outcomes) {
|
|
|
|
|
outcome := client.outcomes[client.next]
|
|
|
|
|
client.next++
|
|
|
|
|
response, err = outcome.response, outcome.err
|
|
|
|
|
}
|
|
|
|
|
client.mu.Unlock()
|
|
|
|
|
if started != nil {
|
|
|
|
|
started <- struct{}{}
|
|
|
|
|
@@ -102,6 +118,12 @@ func (client *fakeClient) request() promptkit.GenerateRequest {
|
|
|
|
|
return client.requests[0]
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (client *fakeClient) allRequests() []promptkit.GenerateRequest {
|
|
|
|
|
client.mu.Lock()
|
|
|
|
|
defer client.mu.Unlock()
|
|
|
|
|
return append([]promptkit.GenerateRequest(nil), client.requests...)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestInspectPromptAndProfile(t *testing.T) {
|
|
|
|
|
adapter := newTestAdapter(t, &fakeClient{})
|
|
|
|
|
inspection, err := adapter.InspectPrompt(context.Background(), "weather.daily_generated_text", "2.0.0")
|
|
|
|
|
@@ -495,6 +517,114 @@ func TestExecuteReturnsCompletedValidationRejection(t *testing.T) {
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestExecuteMapsCorrectiveGenerationResults(t *testing.T) {
|
|
|
|
|
valid := `{"summary":"valid"}`
|
|
|
|
|
invalid := `{"summary":42}`
|
|
|
|
|
tests := []struct {
|
|
|
|
|
name string
|
|
|
|
|
outcomes []generationOutcome
|
|
|
|
|
wantStatus promptexec.ValidationStatus
|
|
|
|
|
wantRepairs int
|
|
|
|
|
wantCalls int
|
|
|
|
|
wantRaw string
|
|
|
|
|
wantUsage promptexec.TokenUsage
|
|
|
|
|
}{
|
|
|
|
|
{
|
|
|
|
|
name: "first pass valid",
|
|
|
|
|
outcomes: []generationOutcome{{response: generationResponse(valid, 2, 3, 5)}},
|
|
|
|
|
wantStatus: promptexec.ValidationPassed,
|
|
|
|
|
wantRepairs: 0,
|
|
|
|
|
wantCalls: 1,
|
|
|
|
|
wantRaw: valid,
|
|
|
|
|
wantUsage: promptexec.TokenUsage{PromptTokens: 2, CompletionTokens: 3, TotalTokens: 5},
|
|
|
|
|
},
|
|
|
|
|
{
|
|
|
|
|
name: "empty output repaired",
|
|
|
|
|
outcomes: []generationOutcome{{response: generationResponse("", 2, 3, 5)}, {response: generationResponse(valid, 7, 11, 18)}},
|
|
|
|
|
wantStatus: promptexec.ValidationPassed,
|
|
|
|
|
wantRepairs: 1,
|
|
|
|
|
wantCalls: 2,
|
|
|
|
|
wantRaw: valid,
|
|
|
|
|
wantUsage: promptexec.TokenUsage{PromptTokens: 9, CompletionTokens: 14, TotalTokens: 23},
|
|
|
|
|
},
|
|
|
|
|
{
|
|
|
|
|
name: "invalid output repaired",
|
|
|
|
|
outcomes: []generationOutcome{{response: generationResponse(invalid, 2, 3, 5)}, {response: generationResponse(valid, 7, 11, 18)}},
|
|
|
|
|
wantStatus: promptexec.ValidationPassed,
|
|
|
|
|
wantRepairs: 1,
|
|
|
|
|
wantCalls: 2,
|
|
|
|
|
wantRaw: valid,
|
|
|
|
|
wantUsage: promptexec.TokenUsage{PromptTokens: 9, CompletionTokens: 14, TotalTokens: 23},
|
|
|
|
|
},
|
|
|
|
|
{
|
|
|
|
|
name: "repair budget exhausted",
|
|
|
|
|
outcomes: []generationOutcome{{response: generationResponse(invalid, 2, 3, 5)}, {response: generationResponse(invalid, 7, 11, 18)}},
|
|
|
|
|
wantStatus: promptexec.ValidationFailed,
|
|
|
|
|
wantRepairs: 1,
|
|
|
|
|
wantCalls: 2,
|
|
|
|
|
wantRaw: invalid,
|
|
|
|
|
wantUsage: promptexec.TokenUsage{PromptTokens: 9, CompletionTokens: 14, TotalTokens: 23},
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
for _, test := range tests {
|
|
|
|
|
t.Run(test.name, func(t *testing.T) {
|
|
|
|
|
client := &fakeClient{outcomes: test.outcomes}
|
|
|
|
|
adapter := newRepairAdapter(t, client, "https://repair.example/v1")
|
|
|
|
|
var preparation promptexec.Preparation
|
|
|
|
|
result, err := adapter.Execute(context.Background(), repairExecuteRequest(), func(value promptexec.Preparation, _ *promptexec.PreparationDebug) error {
|
|
|
|
|
preparation = value
|
|
|
|
|
return nil
|
|
|
|
|
})
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("Execute() error = %v", err)
|
|
|
|
|
}
|
|
|
|
|
if preparation.Output.RepairAttempts != 1 || result == nil || result.Validation.Status != test.wantStatus || result.Validation.RepairAttempts != test.wantRepairs || string(result.RawOutput) != test.wantRaw || result.Usage != test.wantUsage {
|
|
|
|
|
t.Fatalf("preparation/result = %#v/%#v", preparation, result)
|
|
|
|
|
}
|
|
|
|
|
requests := client.allRequests()
|
|
|
|
|
if len(requests) != test.wantCalls {
|
|
|
|
|
t.Fatalf("provider requests = %d, want %d", len(requests), test.wantCalls)
|
|
|
|
|
}
|
|
|
|
|
if test.wantCalls == 2 && !reflect.DeepEqual(requests[0].Target, requests[1].Target) {
|
|
|
|
|
t.Fatalf("corrective target = %#v, want same prepared identity as %#v", requests[1].Target, requests[0].Target)
|
|
|
|
|
}
|
|
|
|
|
if result.ProfileID != preparation.ProfileID || result.BackendID != preparation.BackendID || result.ModelName != preparation.ModelName || result.PromptID != preparation.PromptID || result.PromptVersion != preparation.PromptVersion || result.PromptHash != preparation.PromptHash {
|
|
|
|
|
t.Fatalf("prepared/result identity = %#v/%#v", preparation, result)
|
|
|
|
|
}
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestExecuteMapsCorrectiveGenerationError(t *testing.T) {
|
|
|
|
|
const providerBody = `{"error":{"code":"repair-code","type":"repair-type","message":"repair-message"}}`
|
|
|
|
|
calls := 0
|
|
|
|
|
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
|
|
|
|
|
calls++
|
|
|
|
|
if calls == 1 {
|
|
|
|
|
writer.Header().Set("Content-Type", "application/json")
|
|
|
|
|
_, _ = fmt.Fprintf(writer, `{"choices":[{"message":{"content":%q}}],"usage":{"prompt_tokens":2,"completion_tokens":3,"total_tokens":5}}`, `{"summary":42}`)
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
writer.Header().Set("Content-Type", "application/json")
|
|
|
|
|
writer.WriteHeader(http.StatusUnprocessableEntity)
|
|
|
|
|
_, _ = writer.Write([]byte(providerBody))
|
|
|
|
|
}))
|
|
|
|
|
defer server.Close()
|
|
|
|
|
|
|
|
|
|
adapter := newRepairAdapter(t, nil, server.URL)
|
|
|
|
|
result, err := adapter.Execute(context.Background(), repairExecuteRequest(), nil)
|
|
|
|
|
if result != nil || err == nil || calls != 2 {
|
|
|
|
|
t.Fatalf("result/error/calls = %#v/%v/%d", result, err, calls)
|
|
|
|
|
}
|
|
|
|
|
var generationError *promptexec.GenerationError
|
|
|
|
|
if !errors.As(err, &generationError) || generationError.StatusCode() != http.StatusUnprocessableEntity || generationError.ProviderCode() != "repair-code" || generationError.ProviderType() != "repair-type" || generationError.ProviderMessage() != "repair-message" {
|
|
|
|
|
t.Fatalf("generation error = %#v", err)
|
|
|
|
|
}
|
|
|
|
|
if strings.Contains(err.Error(), "repair-message") || !errors.Is(err, promptkit.ErrLLMGenerate) {
|
|
|
|
|
t.Fatalf("generation error = %v", err)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestExecuteDropsOversizedGeneratedOutput(t *testing.T) {
|
|
|
|
|
client := &fakeClient{response: &promptkit.GenerateResponse{Content: strings.Repeat("x", generatedtext.MaxGeneratedTextBytes+1)}}
|
|
|
|
|
adapter := newTestAdapter(t, client)
|
|
|
|
|
@@ -715,6 +845,61 @@ func validResponse() *promptkit.GenerateResponse {
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func generationResponse(content string, promptTokens int, completionTokens int, totalTokens int) *promptkit.GenerateResponse {
|
|
|
|
|
return &promptkit.GenerateResponse{
|
|
|
|
|
Content: content,
|
|
|
|
|
Usage: promptkit.TokenUsage{
|
|
|
|
|
PromptTokens: promptTokens,
|
|
|
|
|
CompletionTokens: completionTokens,
|
|
|
|
|
TotalTokens: totalTokens,
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func newRepairAdapter(t *testing.T, client promptkit.LLMClient, endpoint string) *Adapter {
|
|
|
|
|
t.Helper()
|
|
|
|
|
profiles := testProfileDirectory(t, map[string]string{"profile.yml": "id: repair-profile\nendpoint: " + endpoint + "\nmodel: repair-model\n"})
|
|
|
|
|
options := []promptkit.Option{
|
|
|
|
|
promptkit.WithPromptFS(fstest.MapFS{
|
|
|
|
|
"repair.yml": &fstest.MapFile{Data: []byte(`id: weather.repair
|
|
|
|
|
version: "1.0.0"
|
|
|
|
|
default_profile: repair-profile
|
|
|
|
|
inputs:
|
|
|
|
|
- name: data_package
|
|
|
|
|
required: true
|
|
|
|
|
content_type: application/yaml
|
|
|
|
|
messages:
|
|
|
|
|
- role: user
|
|
|
|
|
content: "{{input \"data_package\"}}"
|
|
|
|
|
output:
|
|
|
|
|
format: json
|
|
|
|
|
validation_mode: json_schema
|
|
|
|
|
schema_path: repair.schema.json
|
|
|
|
|
repair_attempts: 1
|
|
|
|
|
`)}}, "."),
|
|
|
|
|
promptkit.WithSchemaFS(fstest.MapFS{
|
|
|
|
|
"repair.schema.json": &fstest.MapFile{Data: []byte(`{"type":"object","properties":{"summary":{"type":"string"}},"required":["summary"],"additionalProperties":false}`)},
|
|
|
|
|
}, "."),
|
|
|
|
|
}
|
|
|
|
|
if client != nil {
|
|
|
|
|
options = append(options, promptkit.WithLLMClient(client))
|
|
|
|
|
}
|
|
|
|
|
adapter, err := newAdapter(Config{ProfileDirectory: profiles}, options...)
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("newAdapter() error = %v", err)
|
|
|
|
|
}
|
|
|
|
|
return adapter
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func repairExecuteRequest() promptexec.ExecuteRequest {
|
|
|
|
|
return promptexec.ExecuteRequest{
|
|
|
|
|
PromptID: "weather.repair",
|
|
|
|
|
PromptVersion: "1.0.0",
|
|
|
|
|
ProfileID: "repair-profile",
|
|
|
|
|
DataPackage: []byte("report: repair\n"),
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func hourlyValidResponse() *promptkit.GenerateResponse {
|
|
|
|
|
return &promptkit.GenerateResponse{
|
|
|
|
|
Content: `{"summary":"A quiet hour is expected.","forecast_discussion":"Conditions remain settled.","precipitation_timing":""}`,
|
|
|
|
|
|