diff --git a/internal/adapters/promptkit/adapter.go b/internal/adapters/promptkit/adapter.go new file mode 100644 index 0000000..1696a2f --- /dev/null +++ b/internal/adapters/promptkit/adapter.go @@ -0,0 +1,321 @@ +// Package promptkitadapter implements promptexec with Promptkit. +package promptkitadapter + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "time" + + promptkit "gitea.maximumdirect.net/eric/promptkit" + "gitea.maximumdirect.net/eric/weatherreporter/internal/promptassets" + "gitea.maximumdirect.net/eric/weatherreporter/internal/promptexec" +) + +// Config selects the Promptkit sources and optional local backend for one engine. +type Config struct { + ProfileDirectory string + ProfileFile string + LocalEndpoint string + LocalConcurrencyLimit int + Timeout time.Duration +} + +// Adapter owns one Promptkit engine and its opaque prepared execution handles. +type Adapter struct { + engine *promptkit.Engine +} + +var _ promptexec.Executor = (*Adapter)(nil) + +// New constructs a Promptkit-backed executor from Weatherreporter-owned settings. +func New(config Config) (*Adapter, error) { + return newAdapter(config) +} + +func newAdapter(config Config, additionalOptions ...promptkit.Option) (*Adapter, error) { + if config.ProfileDirectory != "" && config.ProfileFile != "" { + return nil, promptexec.NewError(promptexec.InvalidConfiguration, "profile directory and profile file cannot both be configured", nil) + } + if config.LocalEndpoint == "" && config.LocalConcurrencyLimit != 0 { + return nil, promptexec.NewError(promptexec.InvalidConfiguration, "local concurrency requires a local endpoint", nil) + } + + options := []promptkit.Option{ + promptkit.WithPromptFS(promptassets.PromptFS(), "."), + promptkit.WithSchemaFS(promptassets.SchemaFS(), "."), + } + if config.ProfileFile != "" { + options = append(options, promptkit.WithProfileFile(config.ProfileFile)) + } + if config.LocalEndpoint != "" { + options = append(options, promptkit.WithBackend(promptkit.LocalBackend(config.LocalEndpoint, config.LocalConcurrencyLimit))) + } + options = append(options, additionalOptions...) + engine, err := promptkit.NewEngine(promptkit.Config{ + ProfileDir: config.ProfileDirectory, + Timeout: config.Timeout, + }, options...) + if err != nil { + return nil, classifyConfigurationError(err) + } + return &Adapter{engine: engine}, nil +} + +func newAdapterForTest(config Config, client promptkit.LLMClient) (*Adapter, error) { + return newAdapter(config, promptkit.WithLLMClient(client)) +} + +// InspectPrompt maps an exact Promptkit prompt inspection into project-owned values. +func (adapter *Adapter) InspectPrompt(ctx context.Context, promptID string, promptVersion string) (promptexec.PromptInspection, error) { + if adapter == nil || adapter.engine == nil { + return promptexec.PromptInspection{}, promptexec.NewError(promptexec.InvalidConfiguration, "prompt executor is not configured", nil) + } + inspection, err := adapter.engine.InspectPrompt(ctx, promptID, promptVersion) + if err != nil { + return promptexec.PromptInspection{}, classifyError(err) + } + inputs := make([]promptexec.InputDefinition, len(inspection.Inputs)) + for index, input := range inspection.Inputs { + inputs[index] = promptexec.InputDefinition{ + Name: input.Name, + Required: input.Required, + ContentType: input.ContentType, + Description: input.Description, + } + } + return promptexec.PromptInspection{ + PromptID: inspection.PromptID, + PromptVersion: inspection.PromptVersion, + PromptHash: inspection.PromptHash, + DefaultProfileID: inspection.DefaultProfileID, + Inputs: inputs, + Output: outputContract(inspection.OutputContract), + }, nil +} + +// InspectProfile maps one explicit Promptkit profile inspection into safe values. +func (adapter *Adapter) InspectProfile(ctx context.Context, profileID string) (promptexec.ProfileInspection, error) { + if adapter == nil || adapter.engine == nil { + return promptexec.ProfileInspection{}, promptexec.NewError(promptexec.InvalidConfiguration, "prompt executor is not configured", nil) + } + inspection, err := adapter.engine.InspectProfile(ctx, profileID) + if err != nil { + return promptexec.ProfileInspection{}, classifyError(err) + } + return promptexec.ProfileInspection{ + ProfileID: inspection.ProfileID, + BackendID: inspection.EffectiveModelParams.BackendID, + ModelName: inspection.EffectiveModelParams.Model, + CredentialRequired: inspection.APIKeyRequired || inspection.EffectiveModelParams.APIKeyEnv != "", + }, nil +} + +// Execute prepares one exact inline data package, invokes prepared after a +// successful preparation, and then runs the same opaque prepared handle. +func (adapter *Adapter) Execute(ctx context.Context, request promptexec.ExecuteRequest, preparedCallback promptexec.PreparationCallback) (*promptexec.Execution, error) { + if adapter == nil || adapter.engine == nil { + return nil, promptexec.NewError(promptexec.InvalidConfiguration, "prompt executor is not configured", nil) + } + prepared, err := adapter.engine.PrepareExecution(ctx, promptkit.RunRequest{ + PromptID: request.PromptID, + PromptVersion: request.PromptVersion, + ProfileID: request.ProfileID, + Inputs: map[string]promptkit.ArtifactRef{ + "data_package": promptkit.InlineWithURI(request.DataPackagePath, string(append([]byte(nil), request.DataPackage...))), + }, + }) + if err != nil { + return nil, classifyError(err) + } + defer prepared.Discard() + + details := prepared.Details() + preparation, debug := preparationValues(details, request.DataPackagePath, request.CaptureDebug) + if preparedCallback != nil { + if err := preparedCallback(preparation, debug); err != nil { + return nil, err + } + } + + result, err := adapter.engine.RunPrepared(ctx, prepared) + if err != nil { + return nil, classifyError(err) + } + return executionValue(result, request.DataPackagePath, request.CaptureDebug), nil +} + +func outputContract(value promptkit.OutputContract) promptexec.OutputContract { + return promptexec.OutputContract{ + Format: string(value.Format), + ValidationMode: string(value.ValidationMode), + SchemaPath: value.SchemaPath, + } +} + +func preparationValues(value promptkit.PreparedRun, dataPackagePath string, captureDebug bool) (promptexec.Preparation, *promptexec.PreparationDebug) { + preparation := promptexec.Preparation{ + PromptID: value.PromptID, + PromptVersion: value.PromptVersion, + PromptHash: value.PromptHash, + RenderedPromptHash: value.RenderedPromptHash, + InputHashes: copyInputHashes(value.InputHashes), + ProfileID: value.SelectedProfileID, + BackendID: value.SelectedBackendID, + ModelName: value.EffectiveModelParams.Model, + Output: outputContract(value.OutputContract), + StartedAt: value.StartTime, + EndedAt: value.EndTime, + Duration: time.Duration(value.DurationMS) * time.Millisecond, + DataPackagePath: dataPackagePath, + } + if !captureDebug { + return preparation, nil + } + debug := &promptexec.PreparationDebug{ + RenderedMessages: renderedMessages(value.Messages), + Endpoint: value.EffectiveModelParams.Endpoint, + ParametersJSON: marshalDebugParameters(value.EffectiveModelParams), + } + if value.StructuredOutput != nil && value.StructuredOutput.JSONSchema != nil { + debug.StructuredSchema, _ = json.Marshal(value.StructuredOutput.JSONSchema.Schema) + } + return preparation, debug +} + +func executionValue(value *promptkit.RunResult, dataPackagePath string, captureDebug bool) *promptexec.Execution { + if value == nil { + return nil + } + validation := promptexec.NewValidation( + promptexec.ValidationStatus(value.Validation.Status), + string(value.Validation.Mode), + value.Validation.SchemaPath, + value.Validation.Errors, + ) + execution := &promptexec.Execution{ + RunID: value.RunID, + PromptID: value.PromptID, + PromptVersion: value.PromptVersion, + PromptHash: value.PromptHash, + RenderedPromptHash: value.RenderedPromptHash, + InputHashes: copyInputHashes(value.InputHashes), + ProfileID: value.SelectedProfileID, + BackendID: value.SelectedBackendID, + ModelName: value.ModelName, + GeneratedHash: value.Artifact.Hash, + Usage: promptexec.TokenUsage{ + PromptTokens: value.Usage.PromptTokens, + CompletionTokens: value.Usage.CompletionTokens, + TotalTokens: value.Usage.TotalTokens, + CachedTokens: value.Usage.CachedTokens, + CacheWriteTokens: value.Usage.CacheWriteTokens, + }, + StartedAt: value.StartTime, + EndedAt: value.EndTime, + Duration: value.Duration, + Validation: validation, + DataPackagePath: dataPackagePath, + RawOutput: []byte(value.RawOutput), + } + if captureDebug { + execution.Debug = &promptexec.ExecutionDebug{ + RawOutput: append([]byte(nil), value.RawOutput...), + ValidationDiagnostics: append([]string(nil), validation.Diagnostics...), + } + } + return execution +} + +func renderedMessages(values []promptkit.RenderedMessage) []promptexec.RenderedMessage { + messages := make([]promptexec.RenderedMessage, len(values)) + for index, value := range values { + messages[index] = promptexec.RenderedMessage{Role: value.Role, Content: value.Content} + } + return messages +} + +func copyInputHashes(values map[string]string) map[string]string { + if values == nil { + return nil + } + copy := make(map[string]string, len(values)) + for key, value := range values { + copy[key] = value + } + return copy +} + +func marshalDebugParameters(value promptkit.ExecutionTarget) []byte { + parameters := struct { + Temperature float64 `json:"temperature"` + MaxTokens int `json:"max_tokens"` + TopP float64 `json:"top_p"` + TimeoutSeconds int `json:"timeout_seconds"` + ServiceTier string `json:"service_tier"` + ReasoningEffort string `json:"reasoning_effort"` + ExtraParams map[string]any `json:"extra_params"` + }{ + Temperature: value.Temperature, + MaxTokens: value.MaxTokens, + TopP: value.TopP, + TimeoutSeconds: value.TimeoutSeconds, + ServiceTier: value.ServiceTier, + ReasoningEffort: value.ReasoningEffort, + ExtraParams: value.ExtraParams, + } + data, _ := json.Marshal(parameters) + return data +} + +func classifyConfigurationError(err error) error { + if err == nil { + return nil + } + return promptexec.NewError(promptexec.InvalidConfiguration, "prompt executor configuration is invalid", err) +} + +func classifyError(err error) error { + if err == nil { + return nil + } + if errors.Is(err, context.Canceled) { + return promptexec.NewError(promptexec.Canceled, "prompt operation was canceled", err) + } + if errors.Is(err, context.DeadlineExceeded) { + return promptexec.NewError(promptexec.DeadlineExceeded, "prompt operation exceeded its deadline", err) + } + var capacityError *promptkit.CapacityError + if errors.As(err, &capacityError) { + return promptexec.NewCapacityError(capacityError.BackendID, "prompt backend capacity is unavailable", err) + } + switch { + case errors.Is(err, promptkit.ErrInvalidConfig): + return promptexec.NewError(promptexec.InvalidConfiguration, "prompt executor configuration is invalid", err) + case errors.Is(err, promptkit.ErrPromptNotFound): + return promptexec.NewError(promptexec.PromptNotFound, "prompt definition was not found", err) + case errors.Is(err, promptkit.ErrPromptLoad): + return promptexec.NewError(promptexec.PromptLoad, "prompt definition could not be loaded", err) + case errors.Is(err, promptkit.ErrProfileNotFound): + return promptexec.NewError(promptexec.ProfileNotFound, "execution profile was not found", err) + case errors.Is(err, promptkit.ErrProfileLoad): + return promptexec.NewError(promptexec.ProfileLoad, "execution profile could not be loaded", err) + case errors.Is(err, promptkit.ErrAPIKeyEnvMissing): + return promptexec.NewError(promptexec.MissingCredential, "execution credential is unavailable", err) + case errors.Is(err, promptkit.ErrArtifactLoad): + return promptexec.NewError(promptexec.ArtifactLoad, "prompt input could not be loaded", err) + case errors.Is(err, promptkit.ErrPromptRender): + return promptexec.NewError(promptexec.PromptRender, "prompt could not be rendered", err) + case errors.Is(err, promptkit.ErrCapacityExceeded): + return promptexec.NewCapacityError("", "prompt backend capacity is unavailable", err) + case errors.Is(err, promptkit.ErrLLMGenerate): + return promptexec.NewError(promptexec.Generation, "prompt generation failed", err) + case errors.Is(err, promptkit.ErrValidation): + return promptexec.NewError(promptexec.OperationalValidation, "prompt output validation could not be completed", err) + case errors.Is(err, promptkit.ErrInvalidRequest), errors.Is(err, promptkit.ErrProfileRequired): + return promptexec.NewError(promptexec.InvalidRequest, "prompt execution request is invalid", err) + default: + return promptexec.NewError(promptexec.Generation, "prompt operation failed", fmt.Errorf("%w", err)) + } +} diff --git a/internal/adapters/promptkit/adapter_test.go b/internal/adapters/promptkit/adapter_test.go new file mode 100644 index 0000000..f6a79b6 --- /dev/null +++ b/internal/adapters/promptkit/adapter_test.go @@ -0,0 +1,380 @@ +package promptkitadapter + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + promptkit "gitea.maximumdirect.net/eric/promptkit" + "gitea.maximumdirect.net/eric/weatherreporter/internal/promptexec" +) + +type fakeClient struct { + mu sync.Mutex + response *promptkit.GenerateResponse + err error + calls int + requests []promptkit.GenerateRequest + block bool +} + +type recordingReader struct { + ref promptkit.ArtifactRef +} + +func (reader *recordingReader) Read(_ context.Context, ref promptkit.ArtifactRef) (*promptkit.Artifact, error) { + reader.ref = ref + return &promptkit.Artifact{ + Name: "data_package", + ContentType: "application/yaml", + Body: []byte(ref.Body), + URI: ref.URI, + Hash: "input-hash", + }, nil +} + +func (client *fakeClient) Generate(ctx context.Context, request promptkit.GenerateRequest) (*promptkit.GenerateResponse, error) { + client.mu.Lock() + client.calls++ + client.requests = append(client.requests, request) + block := client.block + response := client.response + err := client.err + client.mu.Unlock() + if block { + <-ctx.Done() + return nil, ctx.Err() + } + return response, err +} + +func (client *fakeClient) callCount() int { + client.mu.Lock() + defer client.mu.Unlock() + return client.calls +} + +func (client *fakeClient) request() promptkit.GenerateRequest { + client.mu.Lock() + defer client.mu.Unlock() + return client.requests[0] +} + +func TestInspectPromptAndProfile(t *testing.T) { + adapter := newTestAdapter(t, &fakeClient{}) + inspection, err := adapter.InspectPrompt(context.Background(), "weather.daily_generated_text", "1.0.0") + if err != nil { + t.Fatalf("InspectPrompt() error = %v", err) + } + if inspection.PromptID != "weather.daily_generated_text" || inspection.PromptVersion != "1.0.0" || inspection.DefaultProfileID != "gemini-flash-latest" { + t.Fatalf("inspection = %#v", inspection) + } + if len(inspection.Inputs) != 1 || inspection.Inputs[0].Name != "data_package" || !inspection.Inputs[0].Required || inspection.Inputs[0].ContentType != "application/yaml" { + t.Fatalf("inputs = %#v", inspection.Inputs) + } + if inspection.Output.Format != "json" || inspection.Output.ValidationMode != "json_schema" || inspection.Output.SchemaPath != "daily.generated_text.schema.json" { + t.Fatalf("output = %#v", inspection.Output) + } + + profile, err := adapter.InspectProfile(context.Background(), "test-profile") + if err != nil { + t.Fatalf("InspectProfile() error = %v", err) + } + if profile.ProfileID != "test-profile" || profile.BackendID != "" || profile.ModelName != "test-model" || profile.CredentialRequired { + t.Fatalf("profile = %#v", profile) + } + if strings.Contains(fmt.Sprintf("%#v", profile), "https://profile.example") { + t.Fatalf("profile leaks endpoint: %#v", profile) + } + + builtin, err := adapter.InspectProfile(context.Background(), "gemini-flash-latest") + if err != nil { + t.Fatalf("InspectProfile(builtin) error = %v", err) + } + if builtin.ProfileID != "gemini-flash-latest" || builtin.ModelName == "" { + t.Fatalf("builtin profile = %#v", builtin) + } +} + +func TestExecuteUsesPreparedInlineDataPackage(t *testing.T) { + client := &fakeClient{response: validResponse()} + adapter := newTestAdapter(t, client) + request := testExecuteRequest() + callbackCalls := 0 + result, err := adapter.Execute(context.Background(), request, func(preparation promptexec.Preparation, debug *promptexec.PreparationDebug) error { + callbackCalls++ + if preparation.PromptID != request.PromptID || preparation.PromptVersion != request.PromptVersion || preparation.DataPackagePath != request.DataPackagePath || preparation.ModelName != "test-model" { + t.Fatalf("preparation = %#v", preparation) + } + if debug != nil { + t.Fatalf("debug = %#v, want nil", debug) + } + if client.callCount() != 0 { + t.Fatal("provider called before preparation callback") + } + return nil + }) + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + if callbackCalls != 1 || client.callCount() != 1 { + t.Fatalf("callback/provider calls = %d/%d, want 1/1", callbackCalls, client.callCount()) + } + if result == nil || result.Validation.Status != promptexec.ValidationPassed || string(result.RawOutput) != client.response.Content || result.DataPackagePath != request.DataPackagePath { + t.Fatalf("result = %#v", result) + } + if result.Debug != nil { + t.Fatalf("debug = %#v, want nil", result.Debug) + } + providerRequest := client.request() + if providerRequest.Target.Model != "test-model" || providerRequest.Target.Endpoint != "https://profile.example/v1" { + t.Fatalf("provider target = %#v", providerRequest.Target) + } + if len(providerRequest.Prompt.Messages) == 0 || !strings.Contains(providerRequest.Prompt.Messages[2].Content, string(request.DataPackage)) { + t.Fatalf("rendered messages do not contain exact data package: %#v", providerRequest.Prompt.Messages) + } +} + +func TestExecuteUsesExactInlineDataPackageProvenance(t *testing.T) { + client := &fakeClient{response: validResponse()} + reader := &recordingReader{} + adapter := newTestAdapterWithOptions(t, client, promptkit.WithArtifactReader(reader)) + request := testExecuteRequest() + if _, err := adapter.Execute(context.Background(), request, nil); err != nil { + t.Fatalf("Execute() error = %v", err) + } + if reader.ref.Type != promptkit.ArtifactRefInline || reader.ref.URI != request.DataPackagePath || reader.ref.Body != string(request.DataPackage) { + t.Fatalf("artifact ref = %#v, want exact inline data package provenance", reader.ref) + } +} + +func TestExecuteCapturesSensitiveDebugOnlyWhenRequested(t *testing.T) { + client := &fakeClient{response: validResponse()} + adapter := newTestAdapter(t, client) + request := testExecuteRequest() + request.CaptureDebug = true + var preparationDebug *promptexec.PreparationDebug + result, err := adapter.Execute(context.Background(), request, func(preparation promptexec.Preparation, debug *promptexec.PreparationDebug) error { + preparationDebug = debug + if strings.Contains(fmt.Sprintf("%#v", preparation), "https://profile.example") || strings.Contains(fmt.Sprintf("%#v", preparation), string(request.DataPackage)) { + t.Fatalf("safe preparation leaks sensitive content: %#v", preparation) + } + return nil + }) + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + if preparationDebug == nil || preparationDebug.Endpoint != "https://profile.example/v1" || len(preparationDebug.RenderedMessages) == 0 || len(preparationDebug.StructuredSchema) == 0 || len(preparationDebug.ParametersJSON) == 0 { + t.Fatalf("preparation debug = %#v", preparationDebug) + } + if result.Debug == nil || string(result.Debug.RawOutput) != client.response.Content { + t.Fatalf("execution debug = %#v", result.Debug) + } +} + +func TestExecuteCallbackFailurePreventsGeneration(t *testing.T) { + client := &fakeClient{response: validResponse()} + adapter := newTestAdapter(t, client) + callbackError := errors.New("save preparation") + result, err := adapter.Execute(context.Background(), testExecuteRequest(), func(promptexec.Preparation, *promptexec.PreparationDebug) error { + return callbackError + }) + if result != nil || !errors.Is(err, callbackError) || client.callCount() != 0 { + t.Fatalf("result/error/provider calls = %#v/%v/%d", result, err, client.callCount()) + } +} + +func TestExecuteReturnsCompletedValidationRejection(t *testing.T) { + client := &fakeClient{response: &promptkit.GenerateResponse{Content: `{"summary":42}`, Usage: promptkit.TokenUsage{TotalTokens: 5}}} + adapter := newTestAdapter(t, client) + result, err := adapter.Execute(context.Background(), testExecuteRequest(), nil) + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + if result == nil || result.Validation.Status != promptexec.ValidationFailed || len(result.Validation.Diagnostics) == 0 || string(result.RawOutput) != client.response.Content { + t.Fatalf("result = %#v", result) + } +} + +func TestExecuteClassifiesOperationalFailures(t *testing.T) { + tests := []struct { + name string + client *fakeClient + context func() (context.Context, context.CancelFunc) + category promptexec.ErrorCategory + }{ + { + name: "generation", + client: &fakeClient{err: errors.New("provider response body")}, + context: func() (context.Context, context.CancelFunc) { + return context.WithCancel(context.Background()) + }, + category: promptexec.Generation, + }, + { + name: "canceled", + client: &fakeClient{block: true}, + context: func() (context.Context, context.CancelFunc) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + return ctx, func() {} + }, + category: promptexec.Canceled, + }, + { + name: "deadline", + client: &fakeClient{block: true}, + context: func() (context.Context, context.CancelFunc) { + return context.WithTimeout(context.Background(), time.Nanosecond) + }, + category: promptexec.DeadlineExceeded, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + adapter := newTestAdapter(t, test.client) + ctx, cancel := test.context() + defer cancel() + result, err := adapter.Execute(ctx, testExecuteRequest(), nil) + if result != nil || err == nil || promptexec.CategoryOf(err) != test.category { + t.Fatalf("result/error/category = %#v/%v/%q, want %q", result, err, promptexec.CategoryOf(err), test.category) + } + if strings.Contains(err.Error(), "provider response body") { + t.Fatalf("error leaks provider detail: %v", err) + } + }) + } +} + +func TestClassifyPromptkitErrors(t *testing.T) { + tests := []struct { + err error + category promptexec.ErrorCategory + }{ + {promptkit.ErrInvalidConfig, promptexec.InvalidConfiguration}, + {promptkit.ErrInvalidRequest, promptexec.InvalidRequest}, + {promptkit.ErrPromptNotFound, promptexec.PromptNotFound}, + {promptkit.ErrPromptLoad, promptexec.PromptLoad}, + {promptkit.ErrProfileNotFound, promptexec.ProfileNotFound}, + {promptkit.ErrProfileLoad, promptexec.ProfileLoad}, + {promptkit.ErrAPIKeyEnvMissing, promptexec.MissingCredential}, + {promptkit.ErrArtifactLoad, promptexec.ArtifactLoad}, + {promptkit.ErrPromptRender, promptexec.PromptRender}, + {promptkit.ErrLLMGenerate, promptexec.Generation}, + {promptkit.ErrValidation, promptexec.OperationalValidation}, + {&promptkit.CapacityError{BackendID: "local"}, promptexec.Capacity}, + } + for _, test := range tests { + t.Run(string(test.category), func(t *testing.T) { + got := classifyError(test.err) + if promptexec.CategoryOf(got) != test.category { + t.Fatalf("category = %q, want %q", promptexec.CategoryOf(got), test.category) + } + }) + } +} + +func TestNewValidatesConfiguration(t *testing.T) { + if _, err := New(Config{ProfileDirectory: "profiles", ProfileFile: "profile.yml"}); promptexec.CategoryOf(err) != promptexec.InvalidConfiguration { + t.Fatalf("profile source error = %v", err) + } + if _, err := New(Config{LocalConcurrencyLimit: 1}); promptexec.CategoryOf(err) != promptexec.InvalidConfiguration { + t.Fatalf("local concurrency error = %v", err) + } + if _, err := New(Config{LocalEndpoint: "not a URL"}); promptexec.CategoryOf(err) != promptexec.InvalidConfiguration { + t.Fatalf("local endpoint error = %v", err) + } +} + +func TestLocalBackendAndMissingCredentialBehavior(t *testing.T) { + profiles := testProfileDirectory(t, `id: local-profile +backend: local +model: local-model +`) + adapter, err := newAdapterForTest(Config{ + ProfileDirectory: profiles, + LocalEndpoint: "https://local.example/v1", + LocalConcurrencyLimit: 1, + }, &fakeClient{}) + if err != nil { + t.Fatalf("newAdapterForTest(local) error = %v", err) + } + profile, err := adapter.InspectProfile(context.Background(), "local-profile") + if err != nil || profile.BackendID != promptkit.BackendLocal || profile.ModelName != "local-model" { + t.Fatalf("local profile/error = %#v/%v", profile, err) + } + if got := classifyError(&promptkit.CapacityError{BackendID: promptkit.BackendLocal}); promptexec.CategoryOf(got) != promptexec.Capacity { + t.Fatalf("capacity classification = %v", got) + } + + credentialProfiles := testProfileDirectory(t, `id: credential-profile +endpoint: https://profile.example/v1 +model: test-model +api_key_env: WEATHERREPORTER_TEST_MISSING_KEY +`) + client := &fakeClient{response: validResponse()} + credentialAdapter, err := newAdapterForTest(Config{ProfileDirectory: credentialProfiles}, client) + if err != nil { + t.Fatalf("newAdapterForTest(credential) error = %v", err) + } + request := testExecuteRequest() + request.ProfileID = "credential-profile" + result, err := credentialAdapter.Execute(context.Background(), request, nil) + if result != nil || promptexec.CategoryOf(err) != promptexec.MissingCredential || client.callCount() != 0 { + t.Fatalf("credential result/category/calls = %#v/%q/%d", result, promptexec.CategoryOf(err), client.callCount()) + } +} + +func newTestAdapter(t *testing.T, client promptkit.LLMClient) *Adapter { + return newTestAdapterWithOptions(t, client) +} + +func newTestAdapterWithOptions(t *testing.T, client promptkit.LLMClient, options ...promptkit.Option) *Adapter { + t.Helper() + profiles := testProfileDirectory(t, `id: test-profile +endpoint: https://profile.example/v1 +model: test-model +temperature: 0.2 +max_tokens: 300 +top_p: 1 +timeout_seconds: 30 +`) + options = append(options, promptkit.WithLLMClient(client)) + adapter, err := newAdapter(Config{ProfileDirectory: profiles, Timeout: time.Second}, options...) + if err != nil { + t.Fatalf("newAdapter() error = %v", err) + } + return adapter +} + +func testProfileDirectory(t *testing.T, profile string) string { + t.Helper() + profiles := t.TempDir() + if err := os.WriteFile(filepath.Join(profiles, "profile.yml"), []byte(profile), 0o600); err != nil { + t.Fatalf("write profile: %v", err) + } + return profiles +} + +func testExecuteRequest() promptexec.ExecuteRequest { + return promptexec.ExecuteRequest{ + PromptID: "weather.daily_generated_text", + PromptVersion: "1.0.0", + ProfileID: "test-profile", + DataPackage: []byte("report:\n id: daily\nbriefing: {}\n"), + DataPackagePath: "data-packages/daily/data_package.yaml", + } +} + +func validResponse() *promptkit.GenerateResponse { + return &promptkit.GenerateResponse{ + Content: `{"summary":"A quiet day is expected.","forecast_discussion":["High pressure keeps conditions settled."],"confidence":"High."}`, + Usage: promptkit.TokenUsage{PromptTokens: 12, CompletionTokens: 8, TotalTokens: 20}, + } +} diff --git a/internal/promptexec/promptexec.go b/internal/promptexec/promptexec.go index 71b78ac..6a985f8 100644 --- a/internal/promptexec/promptexec.go +++ b/internal/promptexec/promptexec.go @@ -148,6 +148,16 @@ type Validation struct { Diagnostics []string } +// NewValidation returns a completed validation value with bounded diagnostics. +func NewValidation(status ValidationStatus, mode string, schemaPath string, diagnostics []string) Validation { + return Validation{ + Status: status, + Mode: mode, + SchemaPath: schemaPath, + Diagnostics: boundDiagnostics(diagnostics), + } +} + // ValidationStatus identifies the completed validation state. type ValidationStatus string diff --git a/internal/promptexec/promptexec_test.go b/internal/promptexec/promptexec_test.go index cc05139..523da8c 100644 --- a/internal/promptexec/promptexec_test.go +++ b/internal/promptexec/promptexec_test.go @@ -172,6 +172,10 @@ func TestBoundDiagnosticAndErrorText(t *testing.T) { if bounded[0] != "a�b" { t.Fatalf("invalid UTF-8 diagnostic = %q, want replacement", bounded[0]) } + validation := NewValidation(ValidationFailed, "json_schema", "daily.schema.json", values) + if len(validation.Diagnostics) != maxValidationDiagnostics || validation.Diagnostics[0] != "a�b" { + t.Fatalf("validation = %#v, want bounded diagnostics", validation) + } err := NewError(Generation, strings.Repeat("é", maxErrorMessageBytes), nil) if len(err.Error()) > maxErrorMessageBytes || !utf8.ValidString(err.Error()) {