Compare commits
46 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| a4264bf4b4 | |||
| 81f41564a2 | |||
| 6b5a2497cc | |||
| 422cc6c978 | |||
| f48e042565 | |||
| d9442850ef | |||
| 2e828157b6 | |||
| 4189146536 | |||
| 59be507d2f | |||
| b45cea4084 | |||
| efe885c6b2 | |||
| 404a4c331d | |||
| 174516cb39 | |||
| b0999112cc | |||
| ed0f7527d5 | |||
| 5064cf833d | |||
| 8745d256bd | |||
| 2e76003fd5 | |||
| 36ce5a5099 | |||
| 465dc1389d | |||
| ae6f1a9865 | |||
| ee99dc9478 | |||
| 00ee5893e9 | |||
| 64d1cffd89 | |||
| f9e8afa2c3 | |||
| e827631d8c | |||
| 115fe8ba58 | |||
| c53250f023 | |||
| 3d99483219 | |||
| 67f788b1e2 | |||
| 764103a2e2 | |||
| a08dd83d1f | |||
| e8922d8ec5 | |||
| 2d44305a8a | |||
| a11c80291e | |||
| 3239567297 | |||
| c239304c2a | |||
| 159b02116f | |||
| e5b7adfb49 | |||
| fa2384e696 | |||
| af0bd3f31a | |||
| 5fff8cd623 | |||
| b2a6c47778 | |||
| a1805fe550 | |||
| 93af155254 | |||
| d783b687a5 |
13
README.md
13
README.md
@@ -33,10 +33,19 @@ boundary and constraints that framework work must preserve.
|
||||
|
||||
## Release Guidance
|
||||
|
||||
Consumers upgrading from `v0.5.0` to `v0.6.0` should read the
|
||||
[v0.6.0 changelog and migration guide](docs/releases/v0.6.0.md).
|
||||
Consumers upgrading from `v0.8.0` to `v0.9.0` should read the
|
||||
[v0.9.0 changelog and migration guide](docs/releases/v0.9.0.md).
|
||||
|
||||
Consumers upgrading from `v0.7.0` to `v0.8.0` can consult the
|
||||
[v0.8.0 changelog and migration guide](docs/releases/v0.8.0.md).
|
||||
|
||||
Earlier adopters can consult the
|
||||
[v0.7.0 changelog and migration guide](docs/releases/v0.7.0.md).
|
||||
|
||||
Consumers upgrading from `v0.5.0` to `v0.6.0` can consult the
|
||||
[v0.6.0 changelog and migration guide](docs/releases/v0.6.0.md).
|
||||
|
||||
Consumers upgrading from `v0.4.0` to `v0.5.0` can consult the
|
||||
[v0.5.0 changelog and migration guide](docs/releases/v0.5.0.md).
|
||||
|
||||
Consumers upgrading from `v0.3.0` to `v0.4.0` should read the
|
||||
|
||||
145
appended_messages_contract_test.go
Normal file
145
appended_messages_contract_test.go
Normal file
@@ -0,0 +1,145 @@
|
||||
package promptkit_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"reflect"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit"
|
||||
)
|
||||
|
||||
func TestAppendedMessageRoleConstantsAreStrings(t *testing.T) {
|
||||
var (
|
||||
developer string = promptkit.RoleDeveloper
|
||||
system string = promptkit.RoleSystem
|
||||
user string = promptkit.RoleUser
|
||||
assistant string = promptkit.RoleAssistant
|
||||
)
|
||||
if developer != "developer" || system != "system" || user != "user" || assistant != "assistant" {
|
||||
t.Fatalf("unexpected role constants: %q %q %q %q", developer, system, user, assistant)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppendedMessagesComposeFrozenEffectivePrompt(t *testing.T) {
|
||||
cacheControl := &promptkit.CacheControl{Type: promptkit.CacheControlEphemeral, TTL: "1h"}
|
||||
appended := []promptkit.RenderedMessage{
|
||||
{Role: " \tAsSiStAnT\n", Content: "previous response"},
|
||||
{Role: promptkit.RoleUser, Content: "corrective request", CacheControl: cacheControl},
|
||||
}
|
||||
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "output"}}
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "profile", "configured prompt"), "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{ID: "profile", Endpoint: "http://example.test/v1", Model: "model"}),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
request := promptkit.RunRequest{PromptID: "prompt", AppendedMessages: appended}
|
||||
|
||||
plain, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare plain request: %v", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), request)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare appended request: %v", err)
|
||||
}
|
||||
expected := []promptkit.RenderedMessage{
|
||||
{Role: promptkit.RoleUser, Content: "configured prompt"},
|
||||
{Role: promptkit.RoleAssistant, Content: "previous response"},
|
||||
{Role: promptkit.RoleUser, Content: "corrective request", CacheControl: &promptkit.CacheControl{Type: promptkit.CacheControlEphemeral, TTL: "1h"}},
|
||||
}
|
||||
if !reflect.DeepEqual(prepared.Messages, expected) || prepared.PromptHash != plain.PromptHash || prepared.RenderedPromptHash == plain.RenderedPromptHash {
|
||||
t.Fatalf("prepared effective prompt = %#v, plain = %#v", prepared, plain)
|
||||
}
|
||||
|
||||
equivalent, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{
|
||||
{Role: promptkit.RoleAssistant, Content: "previous response"},
|
||||
{Role: promptkit.RoleUser, Content: "corrective request", CacheControl: &promptkit.CacheControl{Type: promptkit.CacheControlEphemeral, TTL: "1h"}},
|
||||
}})
|
||||
if err != nil || equivalent.RenderedPromptHash != prepared.RenderedPromptHash {
|
||||
t.Fatalf("normalized-equivalent request = (%#v, %v), want matching rendered hash %q", equivalent, err, prepared.RenderedPromptHash)
|
||||
}
|
||||
|
||||
for _, variant := range []promptkit.RunRequest{
|
||||
{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{{Role: promptkit.RoleAssistant, Content: "changed"}, expected[2]}},
|
||||
{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{expected[2], expected[1]}},
|
||||
{PromptID: "prompt", AppendedMessages: []promptkit.RenderedMessage{{Role: promptkit.RoleAssistant, Content: "previous response"}, {Role: promptkit.RoleUser, Content: "corrective request"}}},
|
||||
} {
|
||||
variantPrepared, prepareErr := engine.Prepare(context.Background(), variant)
|
||||
if prepareErr != nil || variantPrepared.RenderedPromptHash == prepared.RenderedPromptHash || variantPrepared.PromptHash != prepared.PromptHash {
|
||||
t.Fatalf("variant preparation = (%#v, %v)", variantPrepared, prepareErr)
|
||||
}
|
||||
}
|
||||
|
||||
result, err := engine.Run(context.Background(), request)
|
||||
if err != nil || result.RenderedPromptHash != prepared.RenderedPromptHash || !reflect.DeepEqual(client.requests[0].Prompt.Messages, expected) {
|
||||
t.Fatalf("run result = (%#v, %v), request = %#v", result, err, client.requests)
|
||||
}
|
||||
|
||||
execution, err := engine.PrepareExecution(context.Background(), request)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
appended[0].Content = "changed caller content"
|
||||
cacheControl.TTL = ""
|
||||
firstDetails := execution.Details()
|
||||
firstDetails.Messages[1].Content = "changed details content"
|
||||
firstDetails.Messages[2].CacheControl.Type = "changed details cache"
|
||||
secondDetails := execution.Details()
|
||||
if !reflect.DeepEqual(secondDetails.Messages, expected) || secondDetails.RenderedPromptHash != prepared.RenderedPromptHash {
|
||||
t.Fatalf("prepared execution details = %#v, want frozen %#v", secondDetails, expected)
|
||||
}
|
||||
|
||||
result, err = engine.RunPrepared(context.Background(), execution)
|
||||
if err != nil || result.RenderedPromptHash != prepared.RenderedPromptHash || !reflect.DeepEqual(client.requests[1].Prompt.Messages, expected) {
|
||||
t.Fatalf("prepared result = (%#v, %v), requests = %#v", result, err, client.requests)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvalidAppendedMessagesFailBeforeSourceOrModelWork(t *testing.T) {
|
||||
promptSource := &inspectionCountingFS{}
|
||||
var modelCalls atomic.Int64
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(promptSource, "."),
|
||||
promptkit.WithLLMClient(countingLLMClient{calls: &modelCalls}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
request := promptkit.RunRequest{
|
||||
PromptID: "unreached",
|
||||
AppendedMessages: []promptkit.RenderedMessage{{
|
||||
Role: "unsupported-role",
|
||||
Content: "sensitive appended content",
|
||||
}},
|
||||
}
|
||||
|
||||
operations := []struct {
|
||||
name string
|
||||
run func() error
|
||||
}{
|
||||
{name: "Prepare", run: func() error { _, err := engine.Prepare(context.Background(), request); return err }},
|
||||
{name: "PrepareExecution", run: func() error { _, err := engine.PrepareExecution(context.Background(), request); return err }},
|
||||
{name: "Run", run: func() error { _, err := engine.Run(context.Background(), request); return err }},
|
||||
}
|
||||
for _, operation := range operations {
|
||||
t.Run(operation.name, func(t *testing.T) {
|
||||
err := operation.run()
|
||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("error = %v, want ErrInvalidRequest", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
if promptSource.opens.Load() != 0 {
|
||||
t.Fatalf("invalid appended message opened prompt sources %d times", promptSource.opens.Load())
|
||||
}
|
||||
if modelCalls.Load() != 0 {
|
||||
t.Fatalf("invalid appended message invoked the model %d times", modelCalls.Load())
|
||||
}
|
||||
}
|
||||
25
backends.go
25
backends.go
@@ -9,6 +9,10 @@ import (
|
||||
// backend.
|
||||
const BackendOpenRouter = backend.OpenRouterID
|
||||
|
||||
// BackendRakestrawHome is the reserved ID of Promptkit's built-in
|
||||
// Rakestrawhome backend.
|
||||
const BackendRakestrawHome = backend.RakestrawHomeID
|
||||
|
||||
// BackendLocal is the case-sensitive conventional ID used by [LocalBackend].
|
||||
// It is not a built-in or reserved backend and must be registered with
|
||||
// [WithBackend].
|
||||
@@ -20,15 +24,18 @@ const BackendLocal = "local"
|
||||
// to this configuration value do not break source compatibility.
|
||||
type Backend struct {
|
||||
// ID is the stable, case-sensitive registry key. NewEngine trims it and
|
||||
// requires a non-blank value. BackendOpenRouter is reserved.
|
||||
// requires a non-blank value. Built-in backend IDs are reserved.
|
||||
ID string
|
||||
// Endpoint is the OpenAI-compatible base endpoint. NewEngine trims it and
|
||||
// requires an absolute HTTP or HTTPS URL with a host and without user
|
||||
// information, a query string, or a fragment. Paths are allowed.
|
||||
Endpoint string
|
||||
// APIKeyEnv optionally names the environment variable containing the API
|
||||
// key. NewEngine trims it and requires the portable form
|
||||
// [A-Za-z_][A-Za-z0-9_]*. Store only the name, never a credential value.
|
||||
// APIKeyEnv optionally names an environment lookup source for an API key.
|
||||
// NewEngine trims it and requires the portable form [A-Za-z_][A-Za-z0-9_]*.
|
||||
// A direct RunRequest.APIKey takes precedence. When no usable credential is
|
||||
// available, the built-in client omits Authorization; injected clients own
|
||||
// their own credential-resolution behavior. Store only the name, never a
|
||||
// credential value.
|
||||
APIKeyEnv string
|
||||
// ExtraParams contains backend-wide request defaults. Values must be
|
||||
// JSON-compatible, finite, acyclic, and keyed by non-empty strings. Keys
|
||||
@@ -74,11 +81,11 @@ func LocalBackend(endpoint string, concurrencyLimit int) Backend {
|
||||
//
|
||||
// Registrations accumulate in option order. Every normalized ID must be unique
|
||||
// across consumer registrations and built-ins; a duplicate or invalid
|
||||
// definition makes NewEngine fail with ErrInvalidConfig. In particular,
|
||||
// BackendOpenRouter cannot be replaced. The immutable registration is scoped
|
||||
// to the resulting Engine and cannot be enumerated, replaced, removed, or
|
||||
// mutated after construction. WithBackend does not install package-global
|
||||
// state.
|
||||
// definition makes NewEngine fail with ErrInvalidConfig. Built-in IDs,
|
||||
// including [BackendOpenRouter] and [BackendRakestrawHome], cannot be
|
||||
// replaced. The immutable registration is scoped to the resulting Engine and
|
||||
// cannot be enumerated, replaced, removed, or mutated after construction.
|
||||
// WithBackend does not install package-global state.
|
||||
func WithBackend(backend Backend) Option {
|
||||
queueCapacity := 0
|
||||
queueCapacitySet := backend.QueueCapacity != nil
|
||||
|
||||
63
convert.go
63
convert.go
@@ -1,7 +1,9 @@
|
||||
package promptkit
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"unicode/utf8"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
@@ -12,19 +14,52 @@ func toDomainRunRequest(req RunRequest) (domain.RunRequest, error) {
|
||||
if err != nil {
|
||||
return domain.RunRequest{}, err
|
||||
}
|
||||
appendedMessages, err := toDomainAppendedMessages(req.AppendedMessages)
|
||||
if err != nil {
|
||||
return domain.RunRequest{}, err
|
||||
}
|
||||
return domain.RunRequest{
|
||||
PromptID: req.PromptID,
|
||||
PromptVersion: req.PromptVersion,
|
||||
ProfileID: req.ProfileID,
|
||||
SessionID: req.SessionID,
|
||||
APIKey: req.APIKey,
|
||||
Inputs: toDomainArtifactRefMap(req.Inputs),
|
||||
Vars: copyStringMap(req.Vars),
|
||||
Execution: execution,
|
||||
Validation: toDomainOutputContractPtr(req.Validation),
|
||||
PromptID: req.PromptID,
|
||||
PromptVersion: req.PromptVersion,
|
||||
ProfileID: req.ProfileID,
|
||||
SessionID: req.SessionID,
|
||||
APIKey: req.APIKey,
|
||||
Inputs: toDomainArtifactRefMap(req.Inputs),
|
||||
Vars: copyStringMap(req.Vars),
|
||||
Execution: execution,
|
||||
Validation: toDomainOutputContractPtr(req.Validation),
|
||||
AppendedMessages: appendedMessages,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func toDomainAppendedMessages(messages []RenderedMessage) ([]domain.RenderedMessage, error) {
|
||||
if messages == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
converted := make([]domain.RenderedMessage, len(messages))
|
||||
for index, message := range messages {
|
||||
if !utf8.ValidString(message.Content) {
|
||||
return nil, fmt.Errorf("appended message %d content must be valid UTF-8", index)
|
||||
}
|
||||
role, err := domain.NormalizeMessageRole(message.Role)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("appended message %d role: %w", index, err)
|
||||
}
|
||||
|
||||
cacheControl, err := domain.NormalizeCacheControl(toDomainCacheControl(message.CacheControl))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("appended message %d cache_control: %w", index, err)
|
||||
}
|
||||
converted[index] = domain.RenderedMessage{
|
||||
Role: role,
|
||||
Content: message.Content,
|
||||
CacheControl: cacheControl,
|
||||
}
|
||||
}
|
||||
return converted, nil
|
||||
}
|
||||
|
||||
func fromDomainPreparedRun(prepared *domain.PreparedRun) *PreparedRun {
|
||||
if prepared == nil {
|
||||
return nil
|
||||
@@ -285,6 +320,16 @@ func fromDomainRenderedMessages(messages []domain.RenderedMessage) []RenderedMes
|
||||
return out
|
||||
}
|
||||
|
||||
func toDomainCacheControl(cacheControl *CacheControl) *domain.CacheControl {
|
||||
if cacheControl == nil {
|
||||
return nil
|
||||
}
|
||||
return &domain.CacheControl{
|
||||
Type: domain.CacheControlType(cacheControl.Type),
|
||||
TTL: cacheControl.TTL,
|
||||
}
|
||||
}
|
||||
|
||||
func fromDomainCacheControl(cacheControl *domain.CacheControl) *CacheControl {
|
||||
if cacheControl == nil {
|
||||
return nil
|
||||
|
||||
108
convert_appended_messages_test.go
Normal file
108
convert_appended_messages_test.go
Normal file
@@ -0,0 +1,108 @@
|
||||
package promptkit
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
func TestToDomainAppendedMessages(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
messages []RenderedMessage
|
||||
want []domain.RenderedMessage
|
||||
wantErr string
|
||||
privateRole string
|
||||
privateText string
|
||||
}{
|
||||
{
|
||||
name: "normalizes supported roles",
|
||||
messages: []RenderedMessage{
|
||||
{Role: " DeVeLoPeR ", Content: "developer"},
|
||||
{Role: "SyStEm", Content: "system"},
|
||||
{Role: "\tUsEr\n", Content: "user"},
|
||||
{Role: "assistant", Content: "assistant"},
|
||||
},
|
||||
want: []domain.RenderedMessage{
|
||||
{Role: domain.RoleDeveloper, Content: "developer"},
|
||||
{Role: domain.RoleSystem, Content: "system"},
|
||||
{Role: domain.RoleUser, Content: "user"},
|
||||
{Role: domain.RoleAssistant, Content: "assistant"},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid role UTF-8",
|
||||
messages: []RenderedMessage{{Role: string([]byte{0xff}), Content: "private-content"}},
|
||||
wantErr: "appended message 0 role",
|
||||
},
|
||||
{
|
||||
name: "invalid content UTF-8",
|
||||
messages: []RenderedMessage{{Role: "private-role", Content: string([]byte{0xff})}},
|
||||
wantErr: "appended message 0 content",
|
||||
},
|
||||
{
|
||||
name: "unsupported role",
|
||||
messages: []RenderedMessage{{Role: "private-role", Content: "private-content"}},
|
||||
wantErr: "appended message 0 role",
|
||||
privateRole: "private-role",
|
||||
privateText: "private-content",
|
||||
},
|
||||
{
|
||||
name: "invalid cache control",
|
||||
messages: []RenderedMessage{{Role: RoleUser, Content: "private-content", CacheControl: &CacheControl{Type: "private-cache"}}},
|
||||
wantErr: "appended message 0 cache_control",
|
||||
privateText: "private-content",
|
||||
},
|
||||
{
|
||||
name: "empty and whitespace content",
|
||||
messages: []RenderedMessage{
|
||||
{Role: RoleUser, Content: ""},
|
||||
{Role: RoleAssistant, Content: " \t\n "},
|
||||
},
|
||||
want: []domain.RenderedMessage{
|
||||
{Role: domain.RoleUser, Content: ""},
|
||||
{Role: domain.RoleAssistant, Content: " \t\n "},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
got, err := toDomainAppendedMessages(test.messages)
|
||||
if test.wantErr != "" {
|
||||
if err == nil || !strings.Contains(err.Error(), test.wantErr) {
|
||||
t.Fatalf("error = %v, want %q", err, test.wantErr)
|
||||
}
|
||||
for _, privateValue := range []string{test.privateRole, test.privateText} {
|
||||
if privateValue != "" && strings.Contains(err.Error(), privateValue) {
|
||||
t.Fatalf("error exposed appended message data %q: %v", privateValue, err)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("toDomainAppendedMessages() error = %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(got, test.want) {
|
||||
t.Fatalf("toDomainAppendedMessages() = %#v, want %#v", got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestToDomainAppendedMessagesCopiesCacheControl(t *testing.T) {
|
||||
cacheControl := &CacheControl{Type: CacheControlEphemeral, TTL: "1h"}
|
||||
messages := []RenderedMessage{{Role: RoleUser, Content: "content", CacheControl: cacheControl}}
|
||||
request, err := toDomainRunRequest(RunRequest{AppendedMessages: messages})
|
||||
if err != nil {
|
||||
t.Fatalf("toDomainRunRequest() error = %v", err)
|
||||
}
|
||||
|
||||
messages[0].Content = "changed"
|
||||
cacheControl.TTL = ""
|
||||
if request.AppendedMessages[0].Content != "content" || request.AppendedMessages[0].CacheControl.TTL != "1h" {
|
||||
t.Fatalf("domain request did not retain an independent appended-message copy: %#v", request.AppendedMessages)
|
||||
}
|
||||
}
|
||||
16
doc.go
16
doc.go
@@ -23,8 +23,9 @@
|
||||
// InspectProfile return copied inspection values. Returned values and values
|
||||
// passed to extension interfaces are likewise isolated from engine state.
|
||||
// Callers own those copies and may mutate them after the call that supplied or
|
||||
// returned them. Returned structured errors are likewise caller-owned and may
|
||||
// be mutated without affecting engine state or another error.
|
||||
// returned them. [CapacityError] values are caller-owned and may be mutated
|
||||
// without affecting engine state or another error. Immutable [GenerationError]
|
||||
// values are also caller-owned and do not retain shared engine state.
|
||||
//
|
||||
// # Security and sensitive data
|
||||
//
|
||||
@@ -53,9 +54,14 @@
|
||||
// Construction, inspection, handle, and error values, including [Config],
|
||||
// [Backend], [RunRequest], [ArtifactRef], [ExecutionTargetOverride], [Profile],
|
||||
// [OpenAICompatibleProfileConfig], [ProfileInspection],
|
||||
// [PromptInputDefinition], [PromptInspection], [PreparedExecution], and
|
||||
// [CapacityError], do not have stable JSON representations. Direct API keys
|
||||
// are nevertheless excluded from JSON for every public value.
|
||||
// [PromptInputDefinition], [PromptInspection], [PreparedExecution],
|
||||
// [CapacityError], and [GenerationError], do not have stable JSON
|
||||
// representations. Direct API keys are nevertheless excluded from JSON for
|
||||
// every public value.
|
||||
|
||||
// Provider-derived [GenerationError] accessor values are untrusted and can
|
||||
// contain sensitive request or schema fragments. Applications must apply their
|
||||
// own disclosure policy before logging, displaying, or returning them.
|
||||
//
|
||||
// JSON timestamps use time.Time's RFC 3339 encoding and are omitted when zero.
|
||||
// PreparedRun and RunResult durations are encoded as integer milliseconds in
|
||||
|
||||
@@ -173,6 +173,49 @@ semantics. The
|
||||
[OpenAI-compatible integration contract](../integrations/openai-compatible-chat.md)
|
||||
owns the built-in client's outbound HTTP behavior.
|
||||
|
||||
### Repair A Structured Result
|
||||
|
||||
Set a small additional-call budget when a structurally invalid result can be
|
||||
corrected automatically:
|
||||
|
||||
```go
|
||||
request.Validation = &promptkit.OutputContract{
|
||||
Format: promptkit.FormatJSON,
|
||||
ValidationMode: promptkit.ValidationJSONSchema,
|
||||
SchemaPath: "events.schema.json",
|
||||
RepairAttempts: 1,
|
||||
}
|
||||
```
|
||||
|
||||
Each repair attempt is another model call, so it can increase latency and
|
||||
usage; `RunResult.Usage` is cumulative and `Validation.RepairAttempts` reports
|
||||
calls actually started. Exhaustion still returns the final failed validation
|
||||
result. `basic` validation can also repair an empty candidate, but structural
|
||||
validity is not evidence of factual or domain correctness. See the
|
||||
[output-contract format reference](../formats.md#output-contract) and
|
||||
[`OutputContract` GoDoc](../../types.go) for the exact budget and eligibility
|
||||
rules.
|
||||
|
||||
### Append Already-Rendered Messages
|
||||
|
||||
An application can include an earlier assistant response and its own corrective
|
||||
instruction in a fresh request without changing the configured prompt:
|
||||
|
||||
```go
|
||||
request.AppendedMessages = []promptkit.RenderedMessage{
|
||||
{Role: promptkit.RoleAssistant, Content: previousResponse},
|
||||
{Role: promptkit.RoleUser, Content: correction},
|
||||
}
|
||||
result, err := engine.Run(ctx, request)
|
||||
```
|
||||
|
||||
These messages are already rendered: Promptkit does not template or resolve
|
||||
files in them, and they can contain sensitive model output or application
|
||||
feedback. Promptkit remains stateless; every `Run` call re-resolves its current
|
||||
sources and the application owns any semantic retry budget. When a
|
||||
pre-execution equality check is required, use `PrepareExecution` and compare
|
||||
its opaque rendered-prompt hash before invoking `RunPrepared`.
|
||||
|
||||
## Inputs, Profiles, And Overrides
|
||||
|
||||
Use `File`, `Inline`, or `InlineWithURI` to construct request inputs. A request
|
||||
@@ -189,6 +232,53 @@ For programmatic profiles,
|
||||
[`OpenAICompatibleProfile`](../../profiles.go) converts ordinary
|
||||
OpenAI-compatible settings into a value accepted by `WithProfiles`.
|
||||
|
||||
### Alias A Built-In Profile
|
||||
|
||||
Give an application-owned profile ID a built-in base when prompts should select
|
||||
the application ID while inheriting the built-in target. The child can override
|
||||
only the setting it owns:
|
||||
|
||||
```go
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "weather-light",
|
||||
BaseProfileID: "deepseek-4-flash",
|
||||
ReasoningEffort: "high",
|
||||
})
|
||||
```
|
||||
|
||||
Select `weather-light` in a prompt or `RunRequest.ProfileID`; it remains the
|
||||
reported selected profile. See the [profile inheritance format
|
||||
reference](../formats.md#profile-inheritance) and the
|
||||
[`Profile` GoDoc](../../types.go) for exact lookup, merging, and validation
|
||||
behavior.
|
||||
|
||||
### Use The Rakestrawhome Built-In Profile
|
||||
|
||||
Set `RAKESTRAWHOME_INFERENCE_API_KEY` in the application environment, then
|
||||
select `rakestrawhome-gemma-4-31b` as an ordinary profile ID. For example, a
|
||||
prepared result identifies the selected built-in through
|
||||
`BackendRakestrawHome`:
|
||||
|
||||
```go
|
||||
prepared, err := engine.Prepare(ctx, promptkit.RunRequest{
|
||||
PromptID: "meeting.summary",
|
||||
ProfileID: "rakestrawhome-gemma-4-31b",
|
||||
Inputs: inputs,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if prepared.SelectedBackendID != promptkit.BackendRakestrawHome {
|
||||
return fmt.Errorf("unexpected backend %q", prepared.SelectedBackendID)
|
||||
}
|
||||
```
|
||||
|
||||
Do not register `rakestrawhome` manually. When adopting this built-in, remove
|
||||
an existing `WithBackend` registration with that exact ID; retaining it causes
|
||||
the intentional duplicate-ID configuration error. Direct request credentials
|
||||
and runtime endpoint overrides remain supported under their ordinary GoDoc and
|
||||
format contracts.
|
||||
|
||||
### Inspect A Profile Before Prompt Work
|
||||
|
||||
Use [`Engine.InspectProfile`](../../engine.go) to validate one configured
|
||||
@@ -204,7 +294,7 @@ if err != nil {
|
||||
|
||||
target := inspection.EffectiveModelParams
|
||||
if target.APIKeyEnv != "" {
|
||||
// Apply application policy for the named environment variable.
|
||||
// This is a configured optional environment lookup source.
|
||||
} else if inspection.APIKeyRequired {
|
||||
// Arrange a direct credential before later execution.
|
||||
}
|
||||
@@ -213,9 +303,11 @@ if target.APIKeyEnv != "" {
|
||||
Use this configuration-time boundary when only the profile and its target need
|
||||
checking. Use `Prepare` when the application also needs prompt, input, schema,
|
||||
or rendering work; use prepared execution when that work must remain tied to a
|
||||
later execution. Inspection reports credential requirements but leaves the
|
||||
timing of credential enforcement to the application. The method's
|
||||
[GoDoc](../../engine.go) owns its exact result and error contract.
|
||||
later execution. A reported `APIKeyEnv` is a configured optional source, while
|
||||
`APIKeyRequired` is the explicit local requirement. The
|
||||
[credential format reference](../formats.md#credentials) and the method's
|
||||
[GoDoc](../../engine.go) own the exact precedence, timing, result, and error
|
||||
contracts.
|
||||
|
||||
### Set A Per-Run Session And Reasoning
|
||||
|
||||
@@ -423,6 +515,24 @@ status; those choices remain with the consuming application. The
|
||||
contract, while the [`Engine.Run` and error GoDoc](../../engine.go) owns broad
|
||||
error and cancellation identities.
|
||||
|
||||
For a non-2xx response from the built-in OpenAI-compatible client, inspect the
|
||||
status and deliberately selected provider diagnostic when useful:
|
||||
|
||||
```go
|
||||
var generationErr *promptkit.GenerationError
|
||||
if errors.As(err, &generationErr) {
|
||||
status := generationErr.StatusCode()
|
||||
message := generationErr.ProviderMessage()
|
||||
_, _ = status, message // Apply application retry and presentation policy.
|
||||
}
|
||||
```
|
||||
|
||||
All provider fields are untrusted and can contain sensitive request or schema
|
||||
fragments. Do not log, display, or return them without an application-specific
|
||||
disclosure policy. Promptkit does not assign retry or presentation behavior.
|
||||
The [`GenerationError` GoDoc](../../generation_error.go) owns the exact typed
|
||||
error contract.
|
||||
|
||||
## Application Boundary
|
||||
|
||||
Promptkit is an importable library. It does not own a command, inbound HTTP
|
||||
|
||||
@@ -34,6 +34,7 @@ Start with:
|
||||
| Root public API | The [architecture policy](policy/architecture.md), [consumer guide](consumers/pkg-promptkit.md), [testing policy](policy/testing.md), and existing GoDoc. |
|
||||
| Prompt, profile, or schema formats | The [framework format reference](formats.md), owning parser or validator package, and [documentation policy](policy/documentation.md). |
|
||||
| Source loading or validation | The [framework format reference](formats.md), [internal source document](internal/sources.md), and owning package tests. |
|
||||
| Maintained external catalogs or built-ins | The [internal source document](internal/sources.md), [framework format reference](formats.md), [testing policy](policy/testing.md), both catalog module repositories, and each catalog's release procedure. |
|
||||
| Model-client behavior | The [OpenAI-compatible integration contract](integrations/openai-compatible-chat.md), [internal model-client document](internal/llm.md), and owning package tests. |
|
||||
| Internal package implementation | The [architecture policy](policy/architecture.md), [internal component overview](internal/overview.md), and focused internal document listed for that package. |
|
||||
| Tests or test fixtures | The [testing policy](policy/testing.md), owning package, and focused internal document listed by the component overview. |
|
||||
|
||||
162
docs/formats.md
162
docs/formats.md
@@ -82,7 +82,13 @@ allowed.
|
||||
|
||||
### Messages And Templates
|
||||
|
||||
Each message has a non-empty `role` and exactly one of:
|
||||
Each message has a `role` that Promptkit trims and lowercases. It must then be
|
||||
exactly one of `developer`, `system`, `user`, or `assistant`; blank, custom,
|
||||
`tool`, and `function` roles are invalid. This intentionally tightens the
|
||||
previous nonblank-string rule. Consumers migrating to the next minor release
|
||||
must update any nonstandard prompt-definition roles before upgrading.
|
||||
|
||||
Each message also has exactly one of:
|
||||
|
||||
- `content`, containing an inline Go template; or
|
||||
- `content_file`, naming a file whose contents are the Go template.
|
||||
@@ -126,7 +132,7 @@ outbound integration determines its wire representation.
|
||||
| `format` | yes | `text`, `markdown`, or `json`. |
|
||||
| `validation_mode` | yes | `none`, `basic`, `json`, or `json_schema`. |
|
||||
| `schema_path` | for `json_schema` | Path to a schema in the configured schema source. |
|
||||
| `repair_attempts` | no | Integer zero or greater; omitted means zero. |
|
||||
| `repair_attempts` | no | Integer from zero through three; omitted means zero. A positive value requires `basic`, `json`, or `json_schema` validation. |
|
||||
|
||||
The validation modes behave as follows:
|
||||
|
||||
@@ -136,14 +142,31 @@ The validation modes behave as follows:
|
||||
- `json_schema` requires valid JSON that satisfies the selected schema.
|
||||
|
||||
`format` controls output artifact metadata. JSON Schema mode also supplies the
|
||||
schema to compatible model clients as structured-output metadata. The public
|
||||
engine does not install an output repairer, so its validation is single-pass
|
||||
even when a positive `repair_attempts` value is present.
|
||||
schema to compatible model clients as structured-output metadata. Plain `json`
|
||||
validation accepts every valid JSON value and does not request a provider-native
|
||||
JSON-object constraint.
|
||||
|
||||
`repair_attempts` counts additional generation calls after a failed validation.
|
||||
Zero is single-pass. With a positive eligible budget, Promptkit stops at the
|
||||
first valid candidate. If the budget is exhausted, it returns the final
|
||||
candidate and its complete failed validation result; generation and operational
|
||||
validation failures remain errors. `none` never permits repair.
|
||||
|
||||
A request-level `OutputContract` replaces the complete prompt output contract.
|
||||
It does not merge individual fields. If its format is empty, Promptkit uses
|
||||
`text`.
|
||||
|
||||
## Built-In Backends
|
||||
|
||||
Every engine provides these reserved OpenAI-compatible backend IDs. Consumers
|
||||
must not register either ID with `WithBackend`; exact registration and
|
||||
reservation behavior belongs to the [`Backend` GoDoc](../backends.go).
|
||||
|
||||
| ID | Base endpoint | API-key environment variable | Active generation limit | Default queue capacity |
|
||||
| --- | --- | --- | ---: | ---: |
|
||||
| `openrouter` | `https://openrouter.ai/api/v1` | `OPENROUTER_API_KEY` | 16 | 1024 |
|
||||
| `rakestrawhome` | `https://inference.ai.rakestrawhome.com/v1` | `RAKESTRAWHOME_INFERENCE_API_KEY` | 4 | 1024 |
|
||||
|
||||
## Profile Definitions
|
||||
|
||||
A profile supplies model execution settings:
|
||||
@@ -162,9 +185,19 @@ extra_params:
|
||||
provider_option: enabled
|
||||
```
|
||||
|
||||
A derived profile can use a named base and override only the settings it owns:
|
||||
|
||||
```yaml
|
||||
id: local-summary-fast
|
||||
base_profile: local-summary
|
||||
timeout_seconds: 30
|
||||
reasoning_effort: low
|
||||
```
|
||||
|
||||
| Field | Required | Meaning |
|
||||
| --- | --- | --- |
|
||||
| `id` | yes | Profile identifier, trimmed before selection and publication. It must be non-empty after trimming and unique within one source after normalization. |
|
||||
| `base_profile` | no | One optional parent profile ID. A derived profile may inherit target fields from it. |
|
||||
| `backend` | unless `endpoint` is present | Backend registry ID. It is trimmed and registry membership is checked when the profile is prepared or inspected. |
|
||||
| `endpoint` | unless `backend` is present | OpenAI-compatible base URL, including an API version path when required. A nonempty value is trimmed and must be absolute HTTP or HTTPS with a host and without user information, a query, or a fragment. When both connection fields are present, this overrides the backend endpoint without changing backend identity. |
|
||||
| `model` | yes | Non-empty provider model name. |
|
||||
@@ -174,16 +207,21 @@ extra_params:
|
||||
| `timeout_seconds` | no | Per-generation deadline in whole seconds; integer zero or greater. |
|
||||
| `service_tier` | no | Provider-specific request tier. |
|
||||
| `reasoning_effort` | no | Provider-specific reasoning setting. |
|
||||
| `api_key_env` | no | Name of an environment variable containing the API key. |
|
||||
| `api_key_env` | no | Optional environment-variable lookup source for an API key. |
|
||||
| `extra_params` | no | JSON-compatible provider-specific outbound fields. |
|
||||
|
||||
Raw `api_key` is prohibited in profile YAML. Store only an environment
|
||||
variable name in `api_key_env`.
|
||||
|
||||
A standalone profile must provide a model and at least one of `backend` or
|
||||
`endpoint`. A derived profile may omit those target fields because its selected
|
||||
base chain can provide them. Local parsing still validates a derived profile's
|
||||
own ID, supplied endpoint, execution-setting bounds, and `extra_params`.
|
||||
|
||||
Promptkit does not infer a backend from a model or endpoint. Endpoint-only
|
||||
profiles remain supported and have no effective backend ID.
|
||||
The engine always provides the built-in `openrouter` ID. Consumers can add
|
||||
engine-scoped IDs with
|
||||
The engine always provides the built-in `openrouter` and `rakestrawhome` IDs.
|
||||
Consumers can add engine-scoped IDs with
|
||||
[`WithBackend`](../backends.go); exact registration validation belongs to its
|
||||
GoDoc.
|
||||
|
||||
@@ -245,7 +283,7 @@ Profile sources resolve matching IDs in this order:
|
||||
2. the ordinary configured source selected by a profile file, `fs.FS`, or
|
||||
configured profile directory;
|
||||
3. application fallback profiles supplied with `WithFallbackProfileFS`; and
|
||||
4. embedded built-in profiles.
|
||||
4. maintained external catalog profiles.
|
||||
|
||||
A profile source supplies a complete definition; definitions and their fields
|
||||
are not merged across sources. A higher-precedence source falls back only when
|
||||
@@ -255,41 +293,66 @@ profiles. They use `APIKeyRequired` for request-scoped credentials instead of
|
||||
`api_key_env`. Preparation and exact profile inspection use this same source
|
||||
precedence.
|
||||
|
||||
When a selected definition names `base_profile`, every profile ID in that
|
||||
chain is looked up through this same precedence order. A higher-precedence
|
||||
definition therefore shadows a lower-precedence definition of the same base
|
||||
ID, including a built-in. References are not source-qualified.
|
||||
|
||||
### Profile Inheritance
|
||||
|
||||
Promptkit resolves one linear base chain of at most 32 profiles, including the
|
||||
selected profile. It merges settings from the root base to the selected leaf.
|
||||
The leaf's `id` remains the selected profile identity. Nonblank string fields
|
||||
(`backend`, `endpoint`, `model`, `service_tier`, `reasoning_effort`, and
|
||||
`api_key_env`) and nonzero numeric fields replace inherited values. A nonempty
|
||||
`extra_params` map replaces the complete inherited map rather than merging
|
||||
keys, and `APIKeyRequired: true` remains true through the chain. Backend and
|
||||
endpoint are independent: replacing one does not clear the other.
|
||||
|
||||
There is no profile-level clearing syntax. Blank strings, zero numbers, false,
|
||||
and empty maps remain unspecified and inherit from a base. Use existing
|
||||
presence-aware request overrides where an execution needs an explicit zero or
|
||||
empty reasoning setting.
|
||||
|
||||
An absent directly selected profile reports the ordinary not-found error. Once
|
||||
the selected profile exists, a missing base, cycle, overlong chain, or
|
||||
incomplete resolved target is a profile-load failure. Ordinary operations
|
||||
resolve chains afresh; prepared execution retains the fully resolved target.
|
||||
|
||||
## Built-In Profile Catalog
|
||||
|
||||
Every built-in selects the `openrouter` backend. The engine's built-in backend
|
||||
registry supplies `https://openrouter.ai/api/v1` and the environment-variable
|
||||
name `OPENROUTER_API_KEY`, so individual profiles contain only model and
|
||||
generation settings. Built-in profile files do not repeat those connection
|
||||
values. A configured, application fallback, or in-memory profile with the same
|
||||
profile ID takes precedence.
|
||||
Every built-in profile selects one maintained built-in backend and inherits
|
||||
that backend's connection and credential metadata. Profile files do not repeat
|
||||
those values. A configured, application fallback, or in-memory profile with
|
||||
the same profile ID takes precedence.
|
||||
|
||||
| Provider | ID | Model |
|
||||
| --- | --- | --- |
|
||||
| aion-labs | `aion-2` | `aion-labs/aion-2.0` |
|
||||
| anthropic | `claude-fable-latest` | `~anthropic/claude-fable-latest` |
|
||||
| anthropic | `claude-haiku-latest` | `~anthropic/claude-haiku-latest` |
|
||||
| anthropic | `claude-opus-latest` | `~anthropic/claude-opus-latest` |
|
||||
| anthropic | `claude-sonnet-latest` | `~anthropic/claude-sonnet-latest` |
|
||||
| deepseek | `deepseek-3-2` | `deepseek/deepseek-v3.2` |
|
||||
| deepseek | `deepseek-4-flash` | `deepseek/deepseek-v4-flash` |
|
||||
| deepseek | `deepseek-4-pro` | `deepseek/deepseek-v4-pro` |
|
||||
| google | `gemini-2-flash` | `google/gemini-2.5-flash` |
|
||||
| google | `gemini-2-flash-lite` | `google/gemini-2.5-flash-lite` |
|
||||
| google | `gemini-2-pro` | `google/gemini-2.5-pro` |
|
||||
| google | `gemini-3-flash-lite` | `google/gemini-3.1-flash-lite` |
|
||||
| google | `gemini-flash-latest` | `~google/gemini-flash-latest` |
|
||||
| google | `gemini-pro-latest` | `~google/gemini-pro-latest` |
|
||||
| google | `gemma-4-31b` | `google/gemma-4-31b-it:exacto` |
|
||||
| minimax | `minimax-m2` | `minimax/minimax-m2.5` |
|
||||
| minimax | `minimax-m3` | `minimax/minimax-m3` |
|
||||
| mistral | `mistral-large-2512` | `mistralai/mistral-large-2512` |
|
||||
| mistral | `mistral-medium-3-5` | `mistralai/mistral-medium-3-5` |
|
||||
| mistral | `mistral-small-3` | `mistralai/mistral-small-3.2-24b-instruct` |
|
||||
| mistral | `mistral-small-4` | `mistralai/mistral-small-2603` |
|
||||
| nvidia | `nemotron-3-ultra` | `nvidia/nemotron-3-ultra-550b-a55b` |
|
||||
| openai | `gpt-5-mini` | `openai/gpt-5.4-mini` |
|
||||
| openai | `gpt-5-nano` | `openai/gpt-5.4-nano` |
|
||||
| Provider | ID | Backend | Model |
|
||||
| --- | --- | --- | --- |
|
||||
| aion-labs | `aion-2` | `openrouter` | `aion-labs/aion-2.0` |
|
||||
| anthropic | `claude-fable-latest` | `openrouter` | `~anthropic/claude-fable-latest` |
|
||||
| anthropic | `claude-haiku-latest` | `openrouter` | `~anthropic/claude-haiku-latest` |
|
||||
| anthropic | `claude-opus-latest` | `openrouter` | `~anthropic/claude-opus-latest` |
|
||||
| anthropic | `claude-sonnet-latest` | `openrouter` | `~anthropic/claude-sonnet-latest` |
|
||||
| deepseek | `deepseek-3-2` | `openrouter` | `deepseek/deepseek-v3.2` |
|
||||
| deepseek | `deepseek-4-flash` | `openrouter` | `deepseek/deepseek-v4-flash` |
|
||||
| deepseek | `deepseek-4-pro` | `openrouter` | `deepseek/deepseek-v4-pro` |
|
||||
| google | `gemini-2-flash` | `openrouter` | `google/gemini-2.5-flash` |
|
||||
| google | `gemini-2-flash-lite` | `openrouter` | `google/gemini-2.5-flash-lite` |
|
||||
| google | `gemini-2-pro` | `openrouter` | `google/gemini-2.5-pro` |
|
||||
| google | `gemini-3-flash-lite` | `openrouter` | `google/gemini-3.1-flash-lite` |
|
||||
| google | `gemini-flash-latest` | `openrouter` | `~google/gemini-flash-latest` |
|
||||
| google | `gemini-pro-latest` | `openrouter` | `~google/gemini-pro-latest` |
|
||||
| google | `gemma-4-31b` | `openrouter` | `google/gemma-4-31b-it:exacto` |
|
||||
| google | `rakestrawhome-gemma-4-31b` | `rakestrawhome` | `google/gemma-4-31b-it` |
|
||||
| minimax | `minimax-m2` | `openrouter` | `minimax/minimax-m2.5` |
|
||||
| minimax | `minimax-m3` | `openrouter` | `minimax/minimax-m3` |
|
||||
| mistral | `mistral-large-2512` | `openrouter` | `mistralai/mistral-large-2512` |
|
||||
| mistral | `mistral-medium-3-5` | `openrouter` | `mistralai/mistral-medium-3-5` |
|
||||
| mistral | `mistral-small-3` | `openrouter` | `mistralai/mistral-small-3.2-24b-instruct` |
|
||||
| mistral | `mistral-small-4` | `openrouter` | `mistralai/mistral-small-2603` |
|
||||
| nvidia | `nemotron-3-ultra` | `openrouter` | `nvidia/nemotron-3-ultra-550b-a55b` |
|
||||
| openai | `gpt-5-mini` | `openrouter` | `openai/gpt-5.4-mini` |
|
||||
| openai | `gpt-5-nano` | `openrouter` | `openai/gpt-5.4-nano` |
|
||||
|
||||
## Schemas
|
||||
|
||||
@@ -308,16 +371,25 @@ schema produces a failed validation result.
|
||||
Credential values belong at the request or environment boundary, never in
|
||||
prompt, profile, schema, or example files:
|
||||
|
||||
- a file profile names an environment variable with `api_key_env`;
|
||||
- an in-memory profile may set `APIKeyRequired`;
|
||||
- a request can provide a direct `APIKey` or override `APIKeyEnv`; and
|
||||
- a backend or file profile can name an optional environment lookup source
|
||||
with `APIKeyEnv` or `api_key_env`;
|
||||
- an in-memory profile may set `APIKeyRequired` as an explicit local
|
||||
requirement;
|
||||
- a request can provide a direct `APIKey` or override the optional `APIKeyEnv`
|
||||
source; and
|
||||
- a direct request key takes precedence over environment lookup.
|
||||
|
||||
After a direct request key, the credential-source precedence is request
|
||||
`APIKeyEnv`, profile `api_key_env`, then the backend default. An in-memory
|
||||
profile with `APIKeyRequired` clears an inherited backend environment name and
|
||||
requires a direct key unless the request explicitly supplies `APIKeyEnv`.
|
||||
Promptkit validates required credential availability during preparation.
|
||||
Named environment sources are optional: when the selected source is absent,
|
||||
empty, or whitespace-only, the built-in client omits the `Authorization`
|
||||
header and handles the provider response normally. `APIKeyRequired` is the
|
||||
only explicit local availability requirement. Promptkit validates required
|
||||
credential availability during preparation and rechecks it when a prepared
|
||||
execution runs. Injected clients receive resolved source metadata but define
|
||||
their own credential-resolution behavior.
|
||||
Direct keys are excluded from JSON results and redacted by public string
|
||||
formatters. Environment-variable names may appear in prepared metadata, but
|
||||
their values do not.
|
||||
|
||||
@@ -31,11 +31,13 @@ does not serialize it in the provider request.
|
||||
|
||||
## Authentication
|
||||
|
||||
A non-empty API key supplied directly on the execution target takes
|
||||
precedence. Otherwise, when an API-key environment-variable name is supplied,
|
||||
the client reads that variable and requires a non-empty value. The selected
|
||||
key is sent as `Authorization: Bearer <key>`. No authorization header is sent
|
||||
when neither mechanism is configured.
|
||||
A usable API key supplied directly on the execution target takes precedence.
|
||||
Otherwise, when an API-key environment-variable name is supplied, the client
|
||||
reads and trims that variable. A bearer header is sent only when the resolved
|
||||
direct or environment credential is non-empty. When neither source is usable,
|
||||
the client omits `Authorization` and handles the provider response normally.
|
||||
An explicitly required target with no usable source is rejected before
|
||||
transport.
|
||||
|
||||
The target contains the already resolved environment-variable name: an
|
||||
explicit request override takes precedence over profile metadata, which takes
|
||||
@@ -53,6 +55,12 @@ Each ordinary message contains its `role` and string `content`. A
|
||||
cache-controlled message instead uses a text content block containing `type`,
|
||||
`text`, and `cache_control`; an empty cache-control TTL is omitted.
|
||||
|
||||
Promptkit sends only `developer`, `system`, `user`, and `assistant` roles and
|
||||
does so without provider-specific translation. Tool and deprecated function
|
||||
payloads are outside this text-message contract. A backend or model that
|
||||
rejects an otherwise supported role or context returns its ordinary provider
|
||||
error, which follows the normal generation-error path.
|
||||
|
||||
The effective direct or prompt-rendered session ID is trimmed, limited to 256
|
||||
Unicode code points, and sent when nonempty as top-level `session_id`. It is
|
||||
never also sent as a session header.
|
||||
@@ -65,7 +73,8 @@ The client conditionally includes:
|
||||
- non-empty `service_tier` and effective `reasoning_effort`; an explicitly
|
||||
disabled reasoning setting is empty and therefore omitted; and
|
||||
- `response_format` for JSON Schema structured output, including its name,
|
||||
strict flag, and schema document.
|
||||
strict flag, and schema document. Plain JSON validation does not add an
|
||||
object-only response constraint.
|
||||
|
||||
The engine resolves backend, profile, and request extra-parameter maps by
|
||||
whole-map replacement rather than key merging. The resulting effective map is
|
||||
@@ -97,15 +106,33 @@ drained.
|
||||
|
||||
The bounded body must contain exactly one OpenAI-compatible JSON response
|
||||
object followed only by JSON whitespace and EOF. The client returns the first
|
||||
choice's non-empty message content and maps prompt, completion, total, cached,
|
||||
and cache-write token counts. Invalid or truncated JSON, trailing non-whitespace
|
||||
data, a second JSON value, absent choices, empty first-choice content, and size
|
||||
overflow are malformed responses and return no partial result.
|
||||
choice's explicitly present string message content, including an empty or
|
||||
whitespace-only string, and maps prompt, completion, total, cached, and
|
||||
cache-write token counts. Invalid or truncated JSON, trailing non-whitespace
|
||||
data, a second JSON value, absent choices, missing content, `null` content,
|
||||
non-string content, and size overflow are malformed responses and return no
|
||||
partial result.
|
||||
|
||||
For a non-2xx status, the error includes the status code but never the provider
|
||||
response body. Promptkit does not yet parse provider error envelopes; bounded
|
||||
non-success parsing belongs to the
|
||||
[structured-generation-error roadmap](../roadmap/structured-generation-errors.md).
|
||||
For a non-2xx status, Promptkit recognizes one JSON document with a top-level
|
||||
object-valued `error` member. Its optional `message` and `type` fields must be
|
||||
strings, and `code` may be a string or JSON number. Valid supported fields are
|
||||
handled independently, numeric codes retain their JSON number text, and
|
||||
unknown fields are ignored. Missing, invalid, malformed, or multiply framed
|
||||
envelopes contribute no provider detail.
|
||||
|
||||
Non-success bodies have a 65,536-byte limit. A larger declared
|
||||
`Content-Length` is not read; otherwise the client reads at most one additional
|
||||
byte to detect streamed or underreported overflow. Empty, unreadable,
|
||||
oversized, malformed, and unrecognized bodies retain only the received status.
|
||||
The body is always closed and no oversized stream is drained beyond that probe.
|
||||
|
||||
Extracted strings are made valid UTF-8, trimmed, and converted to one line by
|
||||
collapsing Unicode whitespace, control, and format-character runs. Blank
|
||||
values are omitted. Codes and types longer than 256 Unicode code points are
|
||||
omitted; messages longer than 4,096 code points are truncated at a code-point
|
||||
boundary with an ellipsis inside the limit. Promptkit never exposes raw bodies,
|
||||
headers, endpoints, credentials, request data, schemas, generated content, or
|
||||
unsupported provider metadata through this handling.
|
||||
|
||||
An outbound `http.Client.Do` failure retains both Promptkit's request-failure
|
||||
identity and the exact transport error for `errors.Is` and `errors.As` checks.
|
||||
|
||||
@@ -21,8 +21,9 @@ uses internal domain values for rendered prompts, execution targets,
|
||||
structured output, responses, and token usage.
|
||||
|
||||
The runner supplies a fully resolved target after applying backend, profile,
|
||||
and request precedence. The client uses its endpoint, credential metadata,
|
||||
generation fields, and extra parameters. `BackendID` remains routing metadata
|
||||
and request precedence, plus canonical provider-bound text messages. The
|
||||
client uses its endpoint, credential metadata, generation fields, and extra
|
||||
parameters. `BackendID` remains routing metadata
|
||||
for the generation boundary and is not mapped into the provider payload.
|
||||
|
||||
Construction trims and validates a nonempty configured base URL and clones any
|
||||
@@ -37,8 +38,9 @@ resolved request target may supply the endpoint. Generation then:
|
||||
4. composes `/chat/completions` through parsed URL path operations;
|
||||
5. resolves authentication;
|
||||
6. performs the outbound request under the applicable deadlines; and
|
||||
7. decodes one strictly framed, size-bounded response object and maps its first
|
||||
choice and token usage.
|
||||
7. decodes one strictly framed, size-bounded successful response object and
|
||||
maps its first choice and token usage, or decodes bounded structured
|
||||
non-success detail.
|
||||
|
||||
`internal/llm` owns the set of reserved OpenAI-compatible request fields used
|
||||
when validating extra parameters. Backend registration consumes the same rule
|
||||
@@ -54,12 +56,14 @@ the target, rendered messages, and structured-output constraint retained by
|
||||
executable preparation. Execution does not reopen or rerender consumer
|
||||
sources.
|
||||
|
||||
Before backend admission, the runner rechecks that the frozen credential
|
||||
environment-variable name is available. The handle does not retain the
|
||||
environment value; the model client resolves the value visible when generation
|
||||
begins. A direct request key remains in private execution state only until the
|
||||
claimed execution finishes or an unclaimed handle is discarded. Exact public
|
||||
ownership and redaction semantics belong to the
|
||||
Before backend admission, the runner rechecks a frozen credential
|
||||
environment-variable name only when the target explicitly requires a
|
||||
credential. The handle does not retain the environment value; the model client
|
||||
resolves the value visible when generation begins. For optional sources with no
|
||||
usable value, the built-in client omits `Authorization` and continues to the
|
||||
provider. A direct request key remains in private execution state only until
|
||||
the claimed execution finishes or an unclaimed handle is discarded. Exact
|
||||
public ownership and redaction semantics belong to the
|
||||
[`PreparedExecution` GoDoc](../../prepared_execution.go).
|
||||
|
||||
## Failure Categories
|
||||
@@ -67,21 +71,38 @@ ownership and redaction semantics belong to the
|
||||
The package preserves distinct error identities for invalid client
|
||||
configuration, invalid generation requests, request execution failures,
|
||||
non-success provider statuses, and malformed successful responses. Provider
|
||||
response bodies are not included in non-success errors.
|
||||
response bodies are never exposed in raw form through non-success errors.
|
||||
|
||||
Invalid nonempty configured endpoints are configuration failures. A missing or
|
||||
invalid final selected endpoint is an invalid generation request and is
|
||||
rejected before transport.
|
||||
|
||||
Authentication resolves a trimmed direct key before a trimmed configured
|
||||
environment value. Optional missing, empty, or whitespace-only sources do not
|
||||
block transport and produce no `Authorization` header. An explicitly required
|
||||
target with no usable source is rejected before transport with the existing
|
||||
invalid-request diagnostics.
|
||||
|
||||
Successful response bodies have a fixed 16 MiB limit enforced by declared
|
||||
length and by reading at most one byte beyond the boundary. The decoder accepts
|
||||
exactly one JSON object plus trailing whitespace and EOF. Size overflow,
|
||||
truncation, malformed JSON, trailing data, and a second value are malformed
|
||||
responses with no partial result or provider content in the error. Every body
|
||||
is closed, and an unbounded oversized stream is not drained. Non-success
|
||||
responses remain status-only; bounded provider error-envelope parsing belongs
|
||||
to the
|
||||
[structured-generation-error roadmap](../roadmap/structured-generation-errors.md).
|
||||
is closed, and an unbounded oversized stream is not drained.
|
||||
|
||||
After framing succeeds, the first choice must contain an explicitly present
|
||||
string `message.content`. The string is returned exactly, including empty or
|
||||
whitespace-only content. Missing choices, missing or `null` content, and
|
||||
non-string content are malformed responses. Output validation and correction
|
||||
eligibility remain outside this package.
|
||||
|
||||
For a non-success response, `ProviderHTTPError` retains the HTTP status and
|
||||
only normalized detail from the bounded recognized envelope. It retains
|
||||
`ErrUnexpectedStatus` through unwrapping. The client owns response closure;
|
||||
its bounded reader and parser never close or drain a body themselves. The root
|
||||
facade converts this concrete internal error into the public
|
||||
[`GenerationError`](../../generation_error.go), while arbitrary injected-client
|
||||
errors continue through the ordinary generation-error mapping unchanged.
|
||||
|
||||
An `http.Client.Do` failure is represented by a redacting multi-cause error:
|
||||
the package request-failure sentinel and the exact returned transport error are
|
||||
@@ -99,10 +120,14 @@ The
|
||||
own configuration, client cloning, deterministic deadline precedence,
|
||||
authentication, request and response mapping, malformed data, error identity,
|
||||
cancellation, endpoint selection and composition, pre-transport rejection, and
|
||||
bounded single-document response framing, closure, and response-body
|
||||
suppression. The root
|
||||
transport contract tests also verify that resolved backend settings reach this
|
||||
client without serializing backend identity and that ordinary-run cancellation
|
||||
retains its public generation and context identities. All use local test
|
||||
servers or controlled test transports; the default suite makes no live or paid
|
||||
provider requests.
|
||||
bounded single-document successful-response framing, closure, and
|
||||
response-body suppression. The focused
|
||||
[provider HTTP error tests](../../internal/llm/provider_http_error_test.go)
|
||||
own envelope parsing, normalization, and bounded-reader cases; their
|
||||
[transport tests](../../internal/llm/provider_http_error_transport_test.go)
|
||||
own non-success response closure and integration. Root transport contract tests
|
||||
own public `GenerationError` conversion, while also verifying that resolved
|
||||
backend settings reach this client without serializing backend identity and
|
||||
that ordinary-run cancellation retains its public generation and context
|
||||
identities. All use local test servers or controlled test transports; the
|
||||
default suite makes no live or paid provider requests.
|
||||
|
||||
@@ -11,23 +11,23 @@ contributor workflow and validation.
|
||||
|
||||
| Component | Implemented responsibility | References |
|
||||
| --- | --- | --- |
|
||||
| Root `promptkit` package | Provides the supported engine facade, source, backend-registration, and injection options, public request, result, prompt-inspection, and profile-inspection values, opaque prepared-execution handles, profile construction, extension interfaces, value conversion, redacted formatting, typed capacity errors, public error mapping, and engine-local profile-source assembly including application fallbacks. | [Package GoDoc](../../doc.go), [prepared execution](../../prepared_execution.go), [backend API](../../backends.go), [engine assembly](../../engine.go) |
|
||||
| Root `promptkit` package | Provides the supported engine facade, source, backend-registration, and injection options, public request, result, prompt-inspection, and profile-inspection values, opaque prepared-execution handles, profile construction, extension interfaces, value conversion, redacted formatting, typed capacity and generation error mapping, engine-local profile-source assembly including application fallbacks, and bounded output-repair assembly. | [Package GoDoc](../../doc.go), [prepared execution](../../prepared_execution.go), [backend API](../../backends.go), [engine assembly](../../engine.go) |
|
||||
| `examples/go-library/prepare` | Demonstrates an offline downstream consumer using a prompt file, in-memory profile, inline input, and `Prepare`. It is not a public library package. | [Example program](../../examples/go-library/prepare/main.go) |
|
||||
| `examples/go-library/run` | Demonstrates an offline downstream consumer using a prompt file, in-memory profile, inline input, an injected deterministic model client, and `Run`. It is not a public library package. | [Example program](../../examples/go-library/run/main.go) |
|
||||
| `internal/backend` | Constructs each engine's immutable registry from the built-in OpenRouter definition and consumer additions, validates and defensively copies definitions through the shared JSON-value package, and consumes the LLM-owned OpenAI-compatible reserved request-field rule. | [Backend registry](../../internal/backend/registry.go) |
|
||||
| `internal/backend` | Constructs each engine's immutable registry from maintained definitions and consumer additions, validates and defensively copies definitions through the shared JSON-value package, and consumes the LLM-owned OpenAI-compatible reserved request-field rule. | [Backend registry](../../internal/backend/registry.go) |
|
||||
| `internal/catalog` | Strictly validates imported immutable maintained backend and profile catalog assets before engine assembly uses them. | [Catalog adapter](../../internal/catalog/catalog.go), [internal sources](sources.md#profiles-and-built-ins) |
|
||||
| `internal/capacity` | Owns engine-local bounded execution admission and FIFO model-generation permits for limited backend IDs, including cancellation-safe waiter removal and client wrapping. | [Internal capacity management](capacity.md) |
|
||||
| `internal/domain` | Defines internal framework values for requests, artifacts, prompt definitions, profiles, execution targets, rendering, generation, and validation, and owns source-neutral invariants for shared execution settings, OpenAI-compatible base endpoints, session identifiers, and output contracts. Source parsing, required fields, other source-specific normalization, defaulting, and boundary-specific error classification remain with their callers. | [Domain declarations](../../internal/domain/domain.go), [endpoint invariant](../../internal/domain/endpoint.go) |
|
||||
| `internal/defaults` | Defines application-neutral framework constants and constructs the default execution target. It contains no CLI, server, or inbound HTTP limits. | [Framework defaults](../../internal/defaults/defaults.go) |
|
||||
| `internal/filecatalog` | Provides deterministic YAML discovery and path helpers for operating-system filesystems and `fs.FS` sources. | [File catalog](../../internal/filecatalog/catalog.go) |
|
||||
| `internal/jsonvalue` | Validates and deeply copies bounded JSON-compatible extra-parameter and prepared-schema trees while preserving supported concrete value types and rejecting cycles or excessive depth and work. | [JSON values](../../internal/jsonvalue/jsonvalue.go) |
|
||||
| `internal/promptdef` | Loads strictly decoded, validated prompt definitions from filesystem and `fs.FS` sources, including version selection and contained file-backed message content. | [Framework formats](../formats.md), [prompt-definition repository](../../internal/promptdef/filesystem_repository.go) |
|
||||
| `internal/profile` | Loads strictly decoded, validated execution profiles, including backend selection, from filesystem and `fs.FS` sources and composes repositories with error-preserving fallback. | [Framework formats](../formats.md), [profile repositories](../../internal/profile/filesystem_repository.go) |
|
||||
| `internal/profile/builtin` | Embeds the built-in profile catalog, whose entries select OpenRouter. | [Built-in catalog](../formats.md#built-in-profile-catalog), [repository](../../internal/profile/builtin/repository.go) |
|
||||
| `internal/profile` | Loads strictly decoded, locally validated execution profiles from filesystem and `fs.FS` sources, overlays raw sources with error-preserving fallback, and resolves inherited profiles. | [Framework formats](../formats.md), [profile repositories](../../internal/profile/filesystem_repository.go), [internal sources](sources.md#profiles-and-built-ins) |
|
||||
| `internal/prompt` | Renders prompt messages from Go templates with artifact, variable, session, and cache-control data. | [Go-template renderer](../../internal/prompt/go_renderer.go) |
|
||||
| `internal/artifact` | Resolves ordinary inline and unrestricted caller-selected file references into copied artifacts with metadata and hashes. | [Internal sources and validation](sources.md) |
|
||||
| `internal/validate` | Validates basic, JSON, and JSON Schema output using operating-system filesystem or `fs.FS` schema sources and creates operation-local validation plans with canonical contained schema resources. | [Framework formats](../formats.md#schemas), [internal sources and validation](sources.md) |
|
||||
| `internal/llm` | Defines the internal generation boundary and implements outbound OpenAI-compatible chat requests from resolved execution targets, including response decoding, authentication, deadline handling, and ownership of the OpenAI-compatible reserved request-field policy. | [Internal model client](llm.md) |
|
||||
| `internal/usecase` | Resolves prompt definitions and hashes, profiles, backends, and targets for exact inspection and request settings for preparation, and coordinates ordinary execution and one-attempt prepared execution across internal sources, rendering, artifact loading, operation-local validation plans, generation, capacity, and optional repair. | [Internal runner](runner.md), [prepared-execution implementation](../../internal/usecase/prepared_execution.go) |
|
||||
| `internal/llm` | Defines the internal generation boundary and implements outbound OpenAI-compatible chat requests from resolved execution targets, including bounded structured non-success response decoding, successful-response decoding, authentication, deadline handling, and ownership of the OpenAI-compatible reserved request-field policy. | [Internal model client](llm.md) |
|
||||
| `internal/usecase` | Resolves prompt definitions and hashes, profiles, backends, and targets for exact inspection and request settings for preparation, and coordinates ordinary execution and one-attempt prepared execution across internal sources, rendering, artifact loading, operation-local validation plans, generation, capacity, and bounded repair. | [Internal runner](runner.md), [prepared-execution implementation](../../internal/usecase/prepared_execution.go) |
|
||||
|
||||
The root package assembles these internal components without exposing their
|
||||
representations. Consumers depend only on the root facade.
|
||||
|
||||
@@ -23,8 +23,10 @@ built-in backend and validated consumer additions, one engine-local run
|
||||
admitter, and a model client wrapped by the same capacity manager. Validation
|
||||
plans and provider-facing schema metadata come from the validator's preparation
|
||||
interface.
|
||||
An output repairer can be injected internally, but the ordinary runner
|
||||
constructor does not enable one.
|
||||
The root engine supplies one default output repairer through the explicit
|
||||
runner constructor, using the same capacity-wrapped client as initial
|
||||
generation. The no-repair runner constructor remains available for focused
|
||||
internal callers and tests.
|
||||
|
||||
Each invocation carries its state in request, prepared-run, and result values.
|
||||
The runner has no durable run or session store.
|
||||
@@ -82,7 +84,8 @@ profile, or backend:
|
||||
schema metadata from it when required;
|
||||
2. load and hash input artifacts;
|
||||
3. render messages and the prompt-defined session;
|
||||
4. apply any direct session ID;
|
||||
4. apply any direct session ID, then append already-normalized request messages
|
||||
after the rendered definition messages;
|
||||
5. hash the effective rendered prompt; and
|
||||
6. construct the prepared value and preparation timing.
|
||||
|
||||
@@ -111,7 +114,10 @@ the prompt session template, and is applied after ordinary message rendering.
|
||||
A blank direct value retains prompt-template behavior. The runner clears the
|
||||
template only on a value copy of the definition, so the definition hash always
|
||||
describes the original source while the rendered-prompt hash includes the
|
||||
effective direct or rendered session.
|
||||
effective direct or rendered session and the complete effective message
|
||||
sequence. The rendered-prompt hash uses a versioned, length-framed SHA-256
|
||||
encoding of that session, every role and content value, and cache-control
|
||||
presence and values; its hexadecimal value is opaque.
|
||||
|
||||
The registry is read-only after engine construction. Concurrent `Prepare` and
|
||||
`Run` calls resolve independent defensive backend values and keep all
|
||||
@@ -125,8 +131,8 @@ admitter is an internal unlimited fallback. After successful admission, `Run`
|
||||
immediately defers the returned release function, performs the completion
|
||||
phase, makes one initial generation call, builds the named output artifact,
|
||||
and validates that artifact with the plan compiled during completion. Invalid
|
||||
generated content remains a validation result; an inability to perform
|
||||
validation is an operational error.
|
||||
generated content remains a validation result; an inability to generate or
|
||||
perform validation is an operational error.
|
||||
|
||||
Validation preparation and execution honor cancellation at every
|
||||
Promptkit-controlled boundary and do not publish a partial plan or result.
|
||||
@@ -142,17 +148,26 @@ serializing preparation or validation behind the active-generation limit.
|
||||
The wrapped model client separately acquires a FIFO active permit only around
|
||||
each actual generation call.
|
||||
|
||||
When an internal repairer is present, a JSON or JSON Schema content failure can
|
||||
trigger bounded repair attempts. Repair receives the effective execution
|
||||
target, explicit numeric-presence bits, credential, backend identity, session
|
||||
ID, validation errors, prior output, and structured-output specification. One
|
||||
request constructor supplies those common fields to initial and repair
|
||||
generation while their rendered prompts remain intentionally distinct. The
|
||||
default repairer uses the same wrapped client as initial generation, so each
|
||||
repair reacquires the selected backend's active permit while remaining inside
|
||||
its original admission lease. Repair never performs a second bounded
|
||||
admission, and repaired outputs use the operation's existing validation plan.
|
||||
This capability remains internal and is not a public option.
|
||||
After a failed `basic`, JSON, or JSON Schema validation with a positive frozen
|
||||
budget, the installed repairer can make a bounded corrective call. Each request
|
||||
starts with a fresh copy of the complete effective message sequence (the
|
||||
configured prefix followed by the request suffix), includes only the latest
|
||||
nonempty candidate as an assistant message, and appends one corrective user
|
||||
message. Empty candidates omit that assistant message. The
|
||||
correction carries validation diagnostics as JSON data bounded to 64 KiB; the
|
||||
full diagnostics remain in the validation result.
|
||||
|
||||
Repair receives the effective execution target, explicit numeric-presence bits,
|
||||
credential, backend identity, session ID, and structured-output specification.
|
||||
The same request constructor supplies those common fields to initial and repair
|
||||
generation. The default repairer uses the same wrapped client as initial
|
||||
generation, so each repair reacquires the selected backend's active permit
|
||||
while remaining inside its original admission lease. Repair never performs a
|
||||
second bounded admission, and repaired outputs use the operation's existing
|
||||
validation plan. The runner stops at the first valid candidate, sums completed
|
||||
generation usage, reports calls actually started, and returns the final failed
|
||||
validation result on exhaustion. A repair generation failure follows the
|
||||
ordinary generation-error category rather than becoming a validation error.
|
||||
|
||||
A successful result includes the output artifact and raw output, validation
|
||||
state, effective session ID, prompt and rendered-prompt hashes, selected
|
||||
|
||||
@@ -21,6 +21,12 @@ paths, content opening, and root containment. Each lookup remains a
|
||||
point-in-time scan: definitions and catalogs are not cached, and file-backed
|
||||
message content is opened only for the exact selected candidate.
|
||||
|
||||
Message roles are normalized through the shared domain owner by trimming
|
||||
Unicode whitespace and lowercasing. Only `developer`, `system`, `user`, and
|
||||
`assistant` are published; invalid roles remain selected prompt-definition
|
||||
failures rather than becoming request errors. Cache-control metadata uses the
|
||||
same shared domain normalization and defensive-copy rule.
|
||||
|
||||
Operating-system sources enforce containment against canonical roots and
|
||||
targets so symlinks cannot escape. Injected `fs.FS` sources enforce containment
|
||||
in their clean relative path namespace. A single-file source uses the selected
|
||||
@@ -39,17 +45,30 @@ duplicate detection, and source containment:
|
||||
|
||||
## Profiles And Built-Ins
|
||||
|
||||
`internal/profile` loads and validates execution profiles from an
|
||||
operating-system filesystem or an `fs.FS`. A file contains exactly one YAML
|
||||
document and its trimmed YAML `id` is its only selection identity; filenames do
|
||||
not confer authority. Each point lookup reads discovered files once for their
|
||||
metadata and reuses the selected file's bytes for strict decoding; unrelated
|
||||
profiles are not fully decoded. Strict selected decoding recognizes the
|
||||
optional `backend` field, trims its value, and requires a model plus at least
|
||||
one non-blank backend or endpoint. File-backed `extra_params` values are
|
||||
validated and defensively copied through the shared bounded JSON-value owner
|
||||
before a profile is published. OpenAI-compatible reserved-field policy remains
|
||||
with the model-client and backend-registry owners.
|
||||
`internal/profile` loads, locally validates, overlays, and resolves execution
|
||||
profiles from an operating-system filesystem or an `fs.FS`. A file contains
|
||||
exactly one YAML document and its trimmed YAML `id` is its only selection
|
||||
identity; filenames do not confer authority. Each point lookup reads discovered
|
||||
files once for their metadata and reuses the selected file's bytes for strict
|
||||
decoding; unrelated profiles are not fully decoded. Strict selected decoding
|
||||
recognizes `base_profile` and the optional `backend` field, trims their values,
|
||||
and permits inherited target fields only when a base is named. File-backed
|
||||
`extra_params` values are validated and defensively copied through the shared
|
||||
bounded JSON-value owner before a profile is published. OpenAI-compatible
|
||||
reserved-field policy remains with the model-client and backend-registry owners.
|
||||
|
||||
`LoadFSRepository` is the eager immutable loading boundary for internal
|
||||
catalog consumers. It discovers and strictly validates every raw profile once,
|
||||
preserves safe source metadata including explicitly present YAML fields, and
|
||||
publishes independently copied values from memory. It does not resolve profile
|
||||
inheritance. Configured consumer sources continue to use the lazy point lookup
|
||||
repositories described above.
|
||||
|
||||
`internal/catalog` validates the imported immutable OpenRouter and
|
||||
Rakestrawhome asset modules as one private adapter boundary. It enforces their
|
||||
manifest, layout, profile ownership, inheritance, and secret-safety rules
|
||||
before returning raw catalog sources. Root assembly uses those validated
|
||||
catalogs as the maintained lowest-precedence profile source.
|
||||
|
||||
The overlay repository consults the next repository only when the
|
||||
higher-precedence repository reports that a profile is absent. A reliably
|
||||
@@ -59,23 +78,31 @@ backend registry membership because the available registry belongs to the
|
||||
assembled engine; the runner checks membership during preparation and exact
|
||||
profile inspection.
|
||||
|
||||
The root engine assembles profile repositories in precedence order: in-memory
|
||||
profiles, one ordinary configured source, an application fallback source, then
|
||||
the embedded built-in catalog. An explicit file or `fs.FS` profile source
|
||||
replaces `Config.ProfileDir` within the ordinary configured-source category.
|
||||
The root engine assembles one raw composite catalog in precedence order:
|
||||
in-memory profiles, one ordinary configured source, an application fallback
|
||||
source, then the maintained external catalog. An explicit file or `fs.FS` profile
|
||||
source replaces `Config.ProfileDir` within the ordinary configured-source
|
||||
category. One outer resolving repository wraps that complete raw catalog, so
|
||||
each base lookup observes the same precedence and shadowing rules.
|
||||
|
||||
Exact profile inspection performs one point-in-time lookup through those
|
||||
profile sources and checks the resolved target without reading prompt, input,
|
||||
The resolving repository traverses every selected chain afresh, retains no
|
||||
cache, detects cycles, limits a chain to 32 profiles, merges root-to-leaf into a
|
||||
new caller-owned value, and validates the final target before publishing it. It
|
||||
does not check backend registry membership. Exact `base_profile` syntax, merge
|
||||
rules, and consumer-visible failure behavior belong to the [framework format
|
||||
reference](../formats.md#profile-inheritance).
|
||||
|
||||
Exact profile inspection performs one point-in-time resolved lookup through
|
||||
those profile sources and checks the final target without reading prompt, input,
|
||||
or schema sources. It does not retain that lookup for a later execution.
|
||||
Prepared execution instead freezes the fully resolved target; a later ordinary
|
||||
operation performs a fresh traversal.
|
||||
|
||||
`internal/profile/builtin` embeds the maintained built-in profile catalog.
|
||||
Every embedded profile selects `openrouter` and inherits its endpoint and
|
||||
credential environment-variable name from the built-in backend registry rather
|
||||
than repeating those values. Profile loading and overlay behavior are owned by
|
||||
the [profile repository tests](../../internal/profile/repository_test.go),
|
||||
while catalog completeness, the backend-selection invariant, and duplicate IDs
|
||||
are owned by the
|
||||
[built-in repository tests](../../internal/profile/builtin/repository_test.go).
|
||||
The maintained external catalog provides every built-in profile and its
|
||||
matching backend definition. Profile loading and overlay behavior are owned by the
|
||||
[profile repository tests](../../internal/profile/repository_test.go), while
|
||||
catalog completeness, backend selection, and duplicate IDs are owned by the
|
||||
[catalog adapter tests](../../internal/catalog/catalog_test.go).
|
||||
|
||||
## Ordinary Artifacts
|
||||
|
||||
@@ -110,8 +137,9 @@ Session and message parsing and execution remain synchronous. The renderer
|
||||
checks cancellation before and after each parse and execution boundary,
|
||||
between artifact conversion chunks, around each message, and before publishing
|
||||
the complete prompt. It cannot interrupt template work already in progress and
|
||||
never publishes a partial prompt after observing cancellation. It carries
|
||||
message roles, session IDs, and cache control into the rendered prompt. The
|
||||
never publishes a partial prompt after observing cancellation. It validates and
|
||||
canonicalizes message roles before carrying roles, session IDs, and cache
|
||||
control into the rendered prompt. The
|
||||
[renderer tests](../../internal/prompt/renderer_test.go) own rendering behavior.
|
||||
|
||||
## Schemas And Output Validation
|
||||
|
||||
@@ -22,21 +22,21 @@ The implemented internal components consist of:
|
||||
- `internal/domain`, which owns framework data values and source-neutral
|
||||
invariants shared by later internal components;
|
||||
- `internal/backend`, which owns validated immutable OpenAI-compatible backend
|
||||
definitions and the built-in OpenRouter definition;
|
||||
definitions;
|
||||
- `internal/capacity`, which owns engine-local bounded run admission and
|
||||
model-generation scheduling for limited backends;
|
||||
- `internal/defaults`, which owns application-neutral framework defaults and
|
||||
constructs the default execution target;
|
||||
- `internal/filecatalog`, which discovers YAML files and provides source-path
|
||||
helpers for filesystem and `fs.FS` consumers;
|
||||
- `internal/catalog`, which validates immutable external maintained-catalog
|
||||
assets before they are eligible for engine assembly;
|
||||
- `internal/jsonvalue`, which validates and defensively copies JSON-compatible
|
||||
extra-parameter trees;
|
||||
- `internal/promptdef`, which loads and validates prompt definitions from
|
||||
filesystem and `fs.FS` sources;
|
||||
- `internal/profile`, which loads, validates, and overlays execution profiles
|
||||
from filesystem and `fs.FS` sources;
|
||||
- `internal/profile/builtin`, which embeds the built-in execution profile
|
||||
catalog;
|
||||
- `internal/prompt`, which renders prompt messages from Go templates;
|
||||
- `internal/artifact`, which resolves ordinary inline and unrestricted
|
||||
caller-selected file references;
|
||||
@@ -54,13 +54,15 @@ packages or participate in internal assembly.
|
||||
The root facade assembles one immutable backend registry, one capacity manager,
|
||||
the internal repositories, renderer, validator, outbound client, and use-case
|
||||
runner while translating public values and errors at the library boundary. The
|
||||
registry contains built-ins plus validated engine-scoped consumer additions.
|
||||
registry contains validated maintained definitions plus engine-scoped consumer
|
||||
additions.
|
||||
The facade constructs the capacity manager from the registry's immutable
|
||||
policy snapshot, wraps the selected built-in or injected model client, and
|
||||
supplies bounded admission to the runner. The defaults and renderer depend on
|
||||
the domain model. Prompt-definition and profile repositories use the domain
|
||||
model, file catalog, and YAML decoder. The built-in profile repository supplies
|
||||
an embedded `fs.FS` to the profile package. Artifact reading uses the domain
|
||||
model, file catalog, and YAML decoder. The catalog adapter validates imported
|
||||
immutable maintained data before root assembly supplies it to the profile
|
||||
package. Artifact reading uses the domain
|
||||
model and application-neutral defaults. Validation uses the domain model, file
|
||||
catalog, and JSON Schema implementation. The model client uses the domain
|
||||
model, application-neutral defaults, and an injected or standard-library HTTP
|
||||
|
||||
@@ -146,6 +146,13 @@ grep -F 'Consumer action:' "$RELEASE_NOTES_FILE"
|
||||
Inspect the complete message and confirm that it accurately records the
|
||||
compatibility impact, public API changes, and required consumer action.
|
||||
|
||||
When a release introduces or updates the maintained external catalogs, state
|
||||
that it selects two independently versioned data dependencies, preserves the
|
||||
public API and configuration, requires no consumer migration, and guarantees
|
||||
only the catalog versions selected and tested by that Promptkit release. Link
|
||||
to the current architecture and format documentation rather than duplicating
|
||||
catalog contracts in the release note.
|
||||
|
||||
## Create And Inspect The Tag
|
||||
|
||||
Run the candidate guard again immediately before tag creation. This ensures
|
||||
|
||||
155
docs/releases/v0.7.0.md
Normal file
155
docs/releases/v0.7.0.md
Normal file
@@ -0,0 +1,155 @@
|
||||
# Promptkit v0.7.0
|
||||
|
||||
This supplemental changelog and migration guide summarizes the consumer-facing
|
||||
changes from `v0.6.0` to `v0.7.0`. The annotated `v0.7.0` tag is the
|
||||
authoritative release record. Exact current contracts belong to the linked
|
||||
GoDoc and durable documentation.
|
||||
|
||||
## Summary
|
||||
|
||||
`v0.7.0` expands provider integration and profile composition while making
|
||||
credential and generation-failure handling more flexible:
|
||||
|
||||
- Promptkit now includes the `rakestrawhome` backend and its Gemma profile;
|
||||
- built-in generation failures expose bounded structured provider details;
|
||||
- an unavailable optional API-key environment source no longer prevents a
|
||||
request from reaching an upstream that permits unauthenticated access; and
|
||||
- profiles can inherit from and selectively refine another profile.
|
||||
|
||||
## Compatibility
|
||||
|
||||
This release adds public declarations and fields but removes none. Existing
|
||||
keyed configuration literals and ordinary `errors.Is` handling continue to
|
||||
work.
|
||||
|
||||
Adding `BaseProfileID` to `Profile` and `OpenAICompatibleProfileConfig` changes
|
||||
their struct shape. Consumers using positional composite literals for either
|
||||
type must convert them to keyed literals. Existing keyed literals require no
|
||||
change.
|
||||
|
||||
The `rakestrawhome` backend ID is now built in and reserved. A consumer that
|
||||
previously registered that exact ID with `WithBackend` must remove its manual
|
||||
registration before upgrading. Other custom backend registrations are
|
||||
unchanged.
|
||||
|
||||
When an optional backend, profile, or request `APIKeyEnv` is unset, empty, or
|
||||
whitespace-only, the built-in client now omits `Authorization` and sends the
|
||||
request. Previously this condition could fail before transport. Set
|
||||
`Profile.APIKeyRequired` when missing credentials must remain a local
|
||||
preflight error.
|
||||
|
||||
Provider non-success responses continue to match `ErrLLMGenerate`. Their
|
||||
rendered wording is not a compatibility contract; consumers can now use
|
||||
`errors.As` with `*GenerationError` when structured status information is
|
||||
needed.
|
||||
|
||||
## Upgrade
|
||||
|
||||
Update the module dependency with:
|
||||
|
||||
```sh
|
||||
go get gitea.maximumdirect.net/eric/promptkit@v0.7.0
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Remove any manual `rakestrawhome` backend registration, convert positional
|
||||
profile literals to keyed literals, and run the consuming project's ordinary
|
||||
and race-enabled tests.
|
||||
|
||||
## Rakestrawhome Built-In Backend And Profile
|
||||
|
||||
Every engine now includes the reserved `rakestrawhome` backend, identified by
|
||||
`BackendRakestrawHome`. The built-in `rakestrawhome-gemma-4-31b` profile
|
||||
selects that backend. Consumers can use the maintained endpoint, credential,
|
||||
capacity, and model defaults without registering either definition themselves.
|
||||
|
||||
See the [built-in backend and profile catalogs](../formats.md#built-in-backends)
|
||||
and the [consumer adoption example](../consumers/pkg-promptkit.md#use-the-rakestrawhome-built-in-profile)
|
||||
for the current contracts.
|
||||
|
||||
## Structured Generation Errors
|
||||
|
||||
Non-2xx responses from the built-in OpenAI-compatible client now return an
|
||||
immutable `*GenerationError`. Consumers can inspect the HTTP status and any
|
||||
safely extracted provider code, type, or message while retaining the ordinary
|
||||
generation-error category:
|
||||
|
||||
```go
|
||||
var generationErr *promptkit.GenerationError
|
||||
if errors.As(err, &generationErr) {
|
||||
status := generationErr.StatusCode()
|
||||
_ = status
|
||||
}
|
||||
```
|
||||
|
||||
Provider fields are bounded and normalized but remain untrusted and may
|
||||
contain sensitive request or schema details. Default and Go-syntax formatting
|
||||
omit those fields. Applications must apply their own disclosure policy before
|
||||
logging or presenting accessor values.
|
||||
|
||||
See the [`GenerationError` GoDoc](../../generation_error.go), the
|
||||
[consumer error-handling guide](../consumers/pkg-promptkit.md#handle-errors),
|
||||
and the [OpenAI-compatible response contract](../integrations/openai-compatible-chat.md#response-handling).
|
||||
|
||||
## Optional Credential Sources
|
||||
|
||||
`APIKeyEnv` names an optional environment lookup source unless the selected
|
||||
profile explicitly sets `APIKeyRequired`. When neither a direct request key nor
|
||||
a usable environment value exists, the built-in client omits the bearer header
|
||||
and handles the upstream response normally. This supports local and other
|
||||
OpenAI-compatible providers that permit unauthenticated requests without
|
||||
hiding an authentication error returned by a provider that requires one.
|
||||
|
||||
The [credential format reference](../formats.md#credentials), the
|
||||
[`Backend` GoDoc](../../backends.go), the
|
||||
[`ExecutionTargetOverride` GoDoc](../../types.go), and the
|
||||
[authentication integration contract](../integrations/openai-compatible-chat.md#authentication)
|
||||
define the current precedence and availability rules.
|
||||
|
||||
## Profile Inheritance
|
||||
|
||||
YAML profiles can name one parent with `base_profile`; in-memory profiles use
|
||||
`Profile.BaseProfileID`, and `OpenAICompatibleProfileConfig` forwards the same
|
||||
field. A profile can act as an application-owned alias of a built-in or refine
|
||||
selected inherited settings:
|
||||
|
||||
```go
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "weather-light",
|
||||
BaseProfileID: "deepseek-4-flash",
|
||||
ReasoningEffort: "high",
|
||||
})
|
||||
```
|
||||
|
||||
Base lookup observes the existing source precedence. Chains are linear,
|
||||
cycle-safe, and resolved afresh for ordinary operations. Prepared execution
|
||||
freezes the fully resolved target. The selected leaf ID remains public while
|
||||
effective execution settings reflect the resolved chain.
|
||||
|
||||
See the [profile inheritance format reference](../formats.md#profile-inheritance),
|
||||
the [consumer alias example](../consumers/pkg-promptkit.md#alias-a-built-in-profile),
|
||||
and the [`Profile` GoDoc](../../types.go) for exact merge and validation
|
||||
behavior.
|
||||
|
||||
## Public API Changes
|
||||
|
||||
The release adds:
|
||||
|
||||
- `BackendRakestrawHome`;
|
||||
- `GenerationError`, including `StatusCode`, `ProviderCode`, `ProviderType`,
|
||||
`ProviderMessage`, `Error`, `GoString`, and `Unwrap`;
|
||||
- `Profile.BaseProfileID`; and
|
||||
- `OpenAICompatibleProfileConfig.BaseProfileID`.
|
||||
|
||||
No public declaration was removed.
|
||||
|
||||
## Consumer Action
|
||||
|
||||
- Remove a manual backend registration whose ID is exactly `rakestrawhome`.
|
||||
- Convert positional `Profile` or `OpenAICompatibleProfileConfig` literals to
|
||||
keyed literals.
|
||||
- Set `Profile.APIKeyRequired` where a missing credential must fail locally
|
||||
instead of reaching the provider unauthenticated.
|
||||
- Treat `GenerationError` provider fields as untrusted and potentially
|
||||
sensitive when adopting the new accessors.
|
||||
- Run consumer ordinary and race-enabled tests after updating the module.
|
||||
109
docs/releases/v0.8.0.md
Normal file
109
docs/releases/v0.8.0.md
Normal file
@@ -0,0 +1,109 @@
|
||||
# Promptkit v0.8.0
|
||||
|
||||
This supplemental changelog and migration guide summarizes the consumer-facing
|
||||
changes from `v0.7.0` to `v0.8.0`. The annotated `v0.8.0` tag is the
|
||||
authoritative release record. Exact current contracts belong to the linked
|
||||
GoDoc and durable documentation.
|
||||
|
||||
## Summary
|
||||
|
||||
`v0.8.0` activates Promptkit's bounded output-repair workflow:
|
||||
|
||||
- failed nonempty-text, JSON, and JSON Schema validation can make a limited
|
||||
number of corrective model calls;
|
||||
- corrective calls preserve the original rendered conversation, effective
|
||||
target, session, structured-output contract, and backend capacity policy;
|
||||
- results report cumulative usage and the number of corrective calls actually
|
||||
made; and
|
||||
- explicitly empty OpenAI-compatible response content now reaches output
|
||||
validation instead of being classified as a malformed provider envelope.
|
||||
|
||||
## Compatibility
|
||||
|
||||
This release adds no public declarations or fields and removes none. Existing
|
||||
source code remains source-compatible.
|
||||
|
||||
The behavior of the existing `OutputContract.RepairAttempts` field and prompt
|
||||
YAML `repair_attempts` field has changed. A positive value now authorizes real
|
||||
additional model calls after eligible validation failures; earlier releases
|
||||
accepted the field but the public engine remained single-pass. Consumers that
|
||||
set a positive value should expect additional latency, token usage, and
|
||||
provider cost when repair is needed.
|
||||
|
||||
Repair budgets must now be between zero and three. A positive budget requires
|
||||
`basic`, `json`, or `json_schema` validation. Values above three and a positive
|
||||
budget paired with `none` are invalid contracts rather than ignored settings.
|
||||
|
||||
An explicitly present empty or whitespace-only string returned by the built-in
|
||||
OpenAI-compatible client is now a completed generation candidate. `none`
|
||||
validation permits it, while `basic`, `json`, and `json_schema` classify it
|
||||
under their ordinary validation rules and may repair it when configured.
|
||||
Missing, `null`, or non-string content remains a malformed provider response.
|
||||
|
||||
## Upgrade
|
||||
|
||||
Update the module dependency with:
|
||||
|
||||
```sh
|
||||
go get gitea.maximumdirect.net/eric/promptkit@v0.8.0
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Review every prompt definition and request override that sets a positive repair
|
||||
budget. Use zero or omit the field to retain single-pass execution. Ensure each
|
||||
positive budget is no greater than three and uses an eligible validation mode,
|
||||
then run the consuming project's ordinary and race-enabled tests.
|
||||
|
||||
## Bounded Output Repair
|
||||
|
||||
`repair_attempts` counts corrective calls in addition to the initial model
|
||||
call. Promptkit validates each completed candidate, stops at the first valid
|
||||
one, and never exceeds the configured bound. If every candidate remains
|
||||
invalid, the run completes successfully with the final candidate and its
|
||||
failed validation result rather than returning an operational error.
|
||||
|
||||
Each correction starts from the original rendered messages and includes only
|
||||
the latest invalid candidate and latest validation diagnostics. JSON Schema
|
||||
mode retains the provider-native structured-output request as its first line of
|
||||
defense. Promptkit performs only deterministic structural validation; a valid
|
||||
response is not necessarily factual or correct for an application's domain.
|
||||
|
||||
Usage in the final result is cumulative across the initial response and every
|
||||
completed corrective response. `ValidationResult.RepairAttempts` reports the
|
||||
number of corrective calls actually made. Corrective generation failures use
|
||||
the same public generation-error categories and structured provider details as
|
||||
an initial generation failure.
|
||||
|
||||
See the [output-contract format reference](../formats.md#output-contract), the
|
||||
[consumer repair example](../consumers/pkg-promptkit.md#repair-a-structured-result),
|
||||
and the [`OutputContract` and `ValidationResult` GoDoc](../../types.go) for the
|
||||
current contracts.
|
||||
|
||||
## Explicit Empty Content
|
||||
|
||||
The built-in OpenAI-compatible client now distinguishes an explicitly present
|
||||
empty string from a missing or malformed `content` field. This aligns built-in
|
||||
and injected clients by letting the selected output contract decide whether an
|
||||
empty candidate is acceptable, invalid, or eligible for repair.
|
||||
|
||||
See the
|
||||
[OpenAI-compatible response contract](../integrations/openai-compatible-chat.md#response-handling)
|
||||
for the exact envelope behavior.
|
||||
|
||||
## Public API Changes
|
||||
|
||||
None. This release activates and tightens the documented behavior of existing
|
||||
fields.
|
||||
|
||||
## Consumer Action
|
||||
|
||||
- Remove or set `repair_attempts` to zero where execution must remain
|
||||
single-pass.
|
||||
- Keep every positive repair budget at three or fewer and pair it with
|
||||
`basic`, `json`, or `json_schema` validation.
|
||||
- Account for additional latency, usage, and provider cost when enabling
|
||||
repair.
|
||||
- Continue checking the returned validation status because bounded repair can
|
||||
exhaust without producing a valid candidate.
|
||||
- Review workflows that previously treated explicit empty provider content as
|
||||
a generation error.
|
||||
125
docs/releases/v0.9.0.md
Normal file
125
docs/releases/v0.9.0.md
Normal file
@@ -0,0 +1,125 @@
|
||||
# Promptkit v0.9.0
|
||||
|
||||
This supplemental changelog and migration guide summarizes the consumer-facing
|
||||
changes from `v0.8.0` to `v0.9.0`. The annotated `v0.9.0` tag is the
|
||||
authoritative release record. Exact current contracts belong to the linked
|
||||
GoDoc and durable documentation.
|
||||
|
||||
## Summary
|
||||
|
||||
`v0.9.0` adds stateless request-message composition and separates maintained
|
||||
provider data from Promptkit's core implementation:
|
||||
|
||||
- callers can append already-rendered messages to a configured prompt for
|
||||
application-owned conversations and semantic correction workflows;
|
||||
- the public package now publishes constants for the four supported text-chat
|
||||
roles, and prompt definitions use the same normalized role vocabulary; and
|
||||
- the OpenRouter and Rakestrawhome backend/profile catalogs now come from two
|
||||
independently versioned Go module dependencies.
|
||||
|
||||
## Compatibility
|
||||
|
||||
This release adds one field and four constants to the public API and removes no
|
||||
public declarations. Existing keyed `RunRequest` literals that omit
|
||||
`AppendedMessages` retain their behavior. Adding the field changes the struct
|
||||
shape, so consumers using positional `RunRequest` literals must convert them
|
||||
to keyed literals.
|
||||
|
||||
Message roles in prompt definitions are now trimmed, lowercased, and required
|
||||
to be `developer`, `system`, `user`, or `assistant`. Definitions using another
|
||||
role that earlier releases accepted as an arbitrary nonblank string now fail
|
||||
prompt loading. In particular, the text-only message contract does not support
|
||||
`tool` or the deprecated `function` role. Otherwise valid roles with different
|
||||
case or surrounding whitespace are normalized rather than rejected.
|
||||
|
||||
The catalog extraction preserves Promptkit's public API, built-in backend and
|
||||
profile IDs, configuration, precedence, credential handling, and capacity
|
||||
behavior. Consumers do not import or register either catalog themselves, and
|
||||
no configuration migration is required. Promptkit now selects two
|
||||
independently versioned data dependencies and guarantees only the catalog
|
||||
versions selected and tested by this Promptkit release.
|
||||
|
||||
## Upgrade
|
||||
|
||||
Update the module dependency with:
|
||||
|
||||
```sh
|
||||
go get gitea.maximumdirect.net/eric/promptkit@v0.9.0
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Convert any positional `RunRequest` literals to keyed literals. Review prompt
|
||||
definitions for unsupported roles, then run the consuming project's ordinary
|
||||
and race-enabled tests.
|
||||
|
||||
## Appended Request Messages
|
||||
|
||||
`RunRequest.AppendedMessages` accepts caller-owned `RenderedMessage` values
|
||||
that Promptkit validates, copies, and appends after every rendered
|
||||
prompt-definition message in caller order. Promptkit does not template this
|
||||
content, retain conversation state between calls, impose a retry policy, or
|
||||
apply a message-count, byte-size, token, or context-window limit. Upstream
|
||||
rejections continue through the ordinary generation-error boundary.
|
||||
|
||||
This primitive supports application-owned conversation continuations and
|
||||
domain-aware correction loops while preserving Promptkit's existing
|
||||
preparation, hashing, prepared-execution, structural repair, backend-capacity,
|
||||
credential, and cancellation behavior. Appended content can include sensitive
|
||||
model output or application feedback; prepared values expose the complete
|
||||
effective messages by design, while default request formatting reports only
|
||||
the appended-message count.
|
||||
|
||||
See the
|
||||
[consumer appended-message example](../consumers/pkg-promptkit.md#append-already-rendered-messages),
|
||||
the [`RunRequest`, `RenderedMessage`, and `CacheControl` GoDoc](../../types.go),
|
||||
and the [OpenAI-compatible request contract](../integrations/openai-compatible-chat.md#request-body)
|
||||
for current behavior.
|
||||
|
||||
## Supported Message Roles
|
||||
|
||||
The new `RoleDeveloper`, `RoleSystem`, `RoleUser`, and `RoleAssistant`
|
||||
constants identify the complete role vocabulary accepted by Promptkit's
|
||||
text-chat message model. The same validation and normalization now apply to
|
||||
prompt-definition messages and request-supplied appended messages. Promptkit
|
||||
does not translate between roles; provider- or model-specific rejection of an
|
||||
otherwise supported role remains an upstream generation error.
|
||||
|
||||
See the [message format reference](../formats.md#messages-and-templates) for
|
||||
the canonical prompt-definition contract.
|
||||
|
||||
## Independently Versioned Backend Catalogs
|
||||
|
||||
Promptkit imports immutable catalog data for its maintained OpenRouter and
|
||||
Rakestrawhome backends and profiles. Engine construction validates and
|
||||
assembles both catalogs behind the existing built-in registry and profile
|
||||
precedence rules. Promptkit no longer keeps duplicate embedded profile assets
|
||||
or hard-coded definitions for those maintained backends.
|
||||
|
||||
The module versions in Promptkit's `go.mod` identify the catalog releases
|
||||
tested with this release. The [built-in backend and profile format
|
||||
reference](../formats.md#built-in-backends) remains the canonical consumer
|
||||
contract, while the [internal source documentation](../internal/sources.md#profiles-and-built-ins)
|
||||
describes the dependency boundary.
|
||||
|
||||
## Public API Changes
|
||||
|
||||
The release adds:
|
||||
|
||||
- `RunRequest.AppendedMessages`;
|
||||
- `RoleDeveloper`;
|
||||
- `RoleSystem`;
|
||||
- `RoleUser`; and
|
||||
- `RoleAssistant`.
|
||||
|
||||
No public declaration was removed.
|
||||
|
||||
## Consumer Action
|
||||
|
||||
- Convert positional `RunRequest` literals to keyed literals.
|
||||
- Replace unsupported prompt-definition roles with an appropriate supported
|
||||
role, or keep richer tool-call protocols in an application-owned client.
|
||||
- Treat appended messages and prepared effective messages according to the
|
||||
application's sensitive-data policy.
|
||||
- Do not add direct catalog imports or registration calls; existing Promptkit
|
||||
construction and configuration remain correct.
|
||||
- Run consumer ordinary and race-enabled tests after updating the module.
|
||||
@@ -38,33 +38,8 @@ consumers.
|
||||
|
||||
## Ideas
|
||||
|
||||
### Public bounded output repair
|
||||
|
||||
After the codebase-audit remediations are complete, Promptkit should make its
|
||||
bounded output-repair capability available through the public engine. A
|
||||
consumer should be able to request a limited number of corrective generation
|
||||
attempts when JSON or JSON Schema output fails content validation, without
|
||||
having to reproduce Promptkit's generation, validation, capacity, and result-
|
||||
accounting orchestration.
|
||||
|
||||
- Repair is validation recovery, not a general provider retry, failover, or
|
||||
backoff policy. Transport failures, cancellation, and operational schema or
|
||||
validation errors must retain their ordinary error behavior.
|
||||
- Repair must stop after the first valid result or the configured attempt
|
||||
bound. Exhausting the bound should preserve the final invalid result and its
|
||||
validation diagnostics rather than inventing success.
|
||||
- Initial generation and every repair attempt must use the same resolved
|
||||
backend, effective execution settings and presence semantics, session,
|
||||
credential boundary, structured-output contract, and backend-capacity
|
||||
policy.
|
||||
- Results should report the number of repair attempts and cumulative usage for
|
||||
every model call made by the run.
|
||||
- Ordinary and prepared execution should expose coherent behavior, including
|
||||
cancellation, frozen prepared state, error identity, and capacity lifetime.
|
||||
|
||||
Select this work only after the accepted audit findings affecting shared
|
||||
execution invariants, validation, orchestration, transport, and repair
|
||||
internals have been remediated.
|
||||
No ideas are currently awaiting selection. Active feature work belongs in its
|
||||
focused roadmap rather than this catalog.
|
||||
|
||||
## Entry Format
|
||||
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
# Structured Generation Errors
|
||||
|
||||
## Purpose
|
||||
|
||||
Promptkit should give downstream applications actionable, machine-readable
|
||||
details when the built-in OpenAI-compatible client receives a non-success HTTP
|
||||
response. Today the client reports only the status code and discards the
|
||||
provider response body. This makes ordinary configuration failures—such as an
|
||||
unsupported strict JSON Schema keyword—unnecessarily difficult to diagnose.
|
||||
|
||||
## Target End State
|
||||
|
||||
Failures from the built-in transport are available through a public typed error
|
||||
that works with `errors.As` while continuing to match `ErrLLMGenerate` through
|
||||
`errors.Is`. The error should expose:
|
||||
|
||||
- the HTTP status code;
|
||||
- a normalized provider error code or type when supplied; and
|
||||
- a bounded provider message extracted from a recognized OpenAI-compatible
|
||||
JSON error envelope.
|
||||
|
||||
The ordinary `Error()` string should remain safe and concise: it should include
|
||||
the status and provider code or type, but not automatically include the
|
||||
provider message. Consumers that deliberately want the provider's diagnostic
|
||||
text can retrieve it from the typed error and apply their own disclosure and
|
||||
logging policy.
|
||||
|
||||
This contract should be available for both ordinary and prepared execution.
|
||||
Errors returned by injected model clients must continue to preserve their own
|
||||
identity and should not be converted into fabricated HTTP details.
|
||||
|
||||
## Safety And Compatibility Boundaries
|
||||
|
||||
- Never expose the raw response body, response headers, endpoint, credentials,
|
||||
request messages, schema document, or generated content through this API.
|
||||
- Read only a small fixed maximum response body, reject malformed or
|
||||
unrecognized envelopes, normalize invalid UTF-8 and control characters, and
|
||||
cap every retained diagnostic field independently.
|
||||
- Treat the extracted provider message as untrusted and potentially sensitive:
|
||||
its GoDoc must tell consumers not to log or display it without applying their
|
||||
own policy.
|
||||
- Preserve the existing generic behavior when a response is empty, non-JSON,
|
||||
oversized, or does not match a recognized error envelope.
|
||||
- Do not assign retryability from an HTTP status. Promptkit supplies facts;
|
||||
downstream applications retain retry and presentation policy.
|
||||
|
||||
## Recommended API Direction
|
||||
|
||||
Prefer one immutable public `GenerationError` value, constructed internally and
|
||||
carrying accessors for HTTP status, provider code or type, and provider message.
|
||||
This keeps the exact representation evolvable while giving consumers an
|
||||
idiomatic `errors.As` contract. Public Go declarations and GoDoc should own the
|
||||
final exact names and semantics.
|
||||
|
||||
The internal OpenAI-compatible client should parse only the conventional
|
||||
top-level `error` envelope and pass normalized details through the use-case and
|
||||
public error-mapping layers. The integration documentation should continue to
|
||||
own wire behavior; the public declarations should own the consumer contract.
|
||||
|
||||
## Acceptance Criteria
|
||||
|
||||
- A downstream consumer can distinguish a provider HTTP 400 from other
|
||||
generation failures and obtain a bounded provider explanation when present.
|
||||
- The typed error still satisfies `errors.Is(err, ErrLLMGenerate)`.
|
||||
- Existing cancellation, capacity, validation, and injected-client error
|
||||
identities remain unchanged.
|
||||
- Tests cover recognized string and numeric provider codes, absent and malformed
|
||||
envelopes, oversized bodies and fields, control characters, and error-chain
|
||||
behavior without making live provider requests.
|
||||
- Current-state GoDoc and the OpenAI-compatible integration and internal-client
|
||||
documents are updated only when the implementation lands.
|
||||
91
engine.go
91
engine.go
@@ -11,14 +11,16 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
openrouter "gitea.maximumdirect.net/eric/promptkit-backend-openrouter"
|
||||
rakestrawhome "gitea.maximumdirect.net/eric/promptkit-backend-rakestrawhome"
|
||||
artifactadapter "gitea.maximumdirect.net/eric/promptkit/internal/artifact"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/backend"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/capacity"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/catalog"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/defaults"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile/builtin"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/prompt"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/promptdef"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/usecase"
|
||||
@@ -53,9 +55,9 @@ var (
|
||||
// an execution profile or resolve its backend, except for the profile
|
||||
// not-found case represented by ErrProfileNotFound.
|
||||
ErrProfileLoad = errors.New("failed to load execution profile")
|
||||
// ErrAPIKeyEnvMissing identifies an APIKeyEnv whose environment variable is
|
||||
// unset or empty when no direct RunRequest.APIKey takes precedence. Such an
|
||||
// error also matches ErrInvalidRequest.
|
||||
// ErrAPIKeyEnvMissing identifies an explicitly required APIKeyEnv whose
|
||||
// environment variable is unset or empty after direct RunRequest.APIKey
|
||||
// precedence is applied. Such an error also matches ErrInvalidRequest.
|
||||
ErrAPIKeyEnvMissing = errors.New("api_key_env points to an unset environment variable")
|
||||
// ErrArtifactLoad identifies a failure to resolve an input artifact. Errors
|
||||
// returned by an injected ArtifactReader remain available through errors.Is.
|
||||
@@ -69,8 +71,9 @@ var (
|
||||
// request, an LLM or provider rate-limit response, or ErrLLMGenerate.
|
||||
ErrCapacityExceeded = errors.New("backend capacity exceeded")
|
||||
// ErrLLMGenerate identifies a model-client failure or a nil successful
|
||||
// response. Errors returned by an injected LLMClient remain available
|
||||
// through errors.Is.
|
||||
// response. A built-in OpenAI-compatible non-2xx response is available as a
|
||||
// [GenerationError]. Errors returned by an injected LLMClient remain
|
||||
// available through errors.Is.
|
||||
ErrLLMGenerate = errors.New("failed to generate output")
|
||||
// ErrValidation identifies an operational failure to load or compile a
|
||||
// schema or validate output. A completed validation whose Status is
|
||||
@@ -99,7 +102,7 @@ type Config struct {
|
||||
// prompt source.
|
||||
PromptDir string
|
||||
// ProfileDir is an optional ordinary configured source whose profiles take
|
||||
// precedence over application fallback and embedded built-in profiles. An
|
||||
// precedence over application fallback and maintained catalog profiles. An
|
||||
// empty value selects the lower-precedence sources unless a profile-source
|
||||
// option supplies the ordinary source.
|
||||
ProfileDir string
|
||||
@@ -272,7 +275,7 @@ func WithProfileFile(path string) Option {
|
||||
//
|
||||
// Profile lookup checks, in order, profiles supplied by WithProfiles; the
|
||||
// ordinary configured source selected by WithProfileFile, WithProfileFS, or
|
||||
// Config.ProfileDir; this fallback source; and Promptkit's embedded built-in
|
||||
// Config.ProfileDir; this fallback source; and Promptkit's maintained catalog
|
||||
// profiles. Each source supplies a complete profile definition; profile fields
|
||||
// are not merged between sources. Only an absent profile ID proceeds to the
|
||||
// next source. A matching read, parse, duplicate, validation, or credential
|
||||
@@ -303,10 +306,12 @@ func WithFallbackProfileFS(fsys fs.FS, root string) Option {
|
||||
// WithProfiles configures in-memory profiles that take precedence over
|
||||
// ordinary configured, application fallback, and built-in profiles.
|
||||
//
|
||||
// NewEngine validates and copies every profile. IDs must be unique within one
|
||||
// call. An invalid profile, duplicate ID, or unsupported ExtraParams value
|
||||
// makes construction fail with ErrInvalidConfig. Repeating WithProfiles
|
||||
// replaces the complete earlier in-memory set rather than merging it.
|
||||
// NewEngine locally validates and copies every profile. IDs must be unique
|
||||
// within one call. An invalid local definition, duplicate ID, or unsupported
|
||||
// ExtraParams value makes construction fail with ErrInvalidConfig. A derived
|
||||
// profile's base reference and resolved target completeness are checked when it
|
||||
// is selected or inspected. Repeating WithProfiles replaces the complete
|
||||
// earlier in-memory set rather than merging it.
|
||||
func WithProfiles(profiles ...Profile) Option {
|
||||
return optionFunc(func(options *engineOptions) error {
|
||||
repo, err := newMemoryProfileRepository(profiles)
|
||||
@@ -387,9 +392,17 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
|
||||
promptDefs = promptdef.NewFilesystemRepository(cfg.PromptDir)
|
||||
}
|
||||
|
||||
profiles := newProfileRepository(cfg.ProfileDir, options)
|
||||
maintainedCatalogs, err := catalog.Load(
|
||||
catalog.Source{Name: "OpenRouter", ExpectedBackendID: backend.OpenRouterID, FS: openrouter.FS(), Root: openrouter.Root},
|
||||
catalog.Source{Name: "Rakestrawhome", ExpectedBackendID: backend.RakestrawHomeID, FS: rakestrawhome.FS(), Root: rakestrawhome.Root},
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to load maintained catalogs: %v", ErrInvalidConfig, err)
|
||||
}
|
||||
|
||||
backendRegistry, err := backend.NewRegistry(options.backends)
|
||||
profiles := newProfileRepository(cfg.ProfileDir, options, maintainedCatalogs.Profiles)
|
||||
|
||||
backendRegistry, err := backend.NewRegistry(maintainedCatalogs.Backends, options.backends)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to construct backend registry: %v", ErrInvalidConfig, err)
|
||||
}
|
||||
@@ -427,7 +440,7 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
|
||||
}
|
||||
|
||||
return &Engine{
|
||||
runner: usecase.NewRunner(
|
||||
runner: usecase.NewRunnerWithRepairer(
|
||||
promptDefs,
|
||||
profiles,
|
||||
backendRegistry,
|
||||
@@ -435,13 +448,14 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
|
||||
prompt.NewGoRenderer(),
|
||||
llmClient,
|
||||
validator,
|
||||
usecase.NewDefaultOutputRepairer(llmClient),
|
||||
capacityManager,
|
||||
),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newProfileRepository(profileDir string, options engineOptions) profile.Repository {
|
||||
repository := builtin.NewRepository()
|
||||
func newProfileRepository(profileDir string, options engineOptions, maintained profile.Repository) profile.Repository {
|
||||
repository := maintained
|
||||
|
||||
if options.fallbackProfileSource {
|
||||
repository = profile.NewOverlayRepository(options.fallbackProfiles, repository)
|
||||
@@ -457,7 +471,7 @@ func newProfileRepository(profileDir string, options engineOptions) profile.Repo
|
||||
repository = profile.NewOverlayRepository(options.memoryProfiles, repository)
|
||||
}
|
||||
|
||||
return repository
|
||||
return profile.NewResolvingRepository(repository)
|
||||
}
|
||||
|
||||
func fileSource(name string) (fs.FS, string, error) {
|
||||
@@ -567,6 +581,8 @@ func (e *Engine) InspectProfile(ctx context.Context, profileID string) (*Profile
|
||||
}
|
||||
|
||||
// Prepare resolves and renders a prompt request without calling an LLM.
|
||||
// It appends any validated RunRequest.AppendedMessages after rendered
|
||||
// definition messages in the returned caller-owned snapshot.
|
||||
//
|
||||
// Prepare selects the prompt and profile, resolves any selected backend and
|
||||
// effective execution settings, resolves the output contract, loads and hashes
|
||||
@@ -601,7 +617,8 @@ func (e *Engine) Prepare(ctx context.Context, req RunRequest) (*PreparedRun, err
|
||||
}
|
||||
|
||||
// PrepareExecution completely prepares a prompt request without calling the
|
||||
// configured LLMClient or reserving backend admission capacity.
|
||||
// configured LLMClient or reserving backend admission capacity. Validated
|
||||
// RunRequest.AppendedMessages are included in the frozen effective messages.
|
||||
//
|
||||
// The returned opaque handle is bound to this Engine and permits one
|
||||
// [Engine.RunPrepared] invocation. Preparation freezes the selected sources,
|
||||
@@ -633,24 +650,29 @@ func (e *Engine) PrepareExecution(ctx context.Context, req RunRequest) (*Prepare
|
||||
}
|
||||
|
||||
// Run prepares a request, invokes the configured LLMClient, and validates the
|
||||
// generated output.
|
||||
// generated output. Each call resolves current sources and composes a fresh,
|
||||
// stateless effective prompt with any validated RunRequest.AppendedMessages.
|
||||
//
|
||||
// A content-validation failure is a successful run whose
|
||||
// RunResult.Validation has Status ValidationFailed. An inability to perform
|
||||
// validation returns an error matching ErrValidation and no partial result.
|
||||
// The public Engine does not perform output repair, so validation is
|
||||
// single-pass even when OutputContract.RepairAttempts is positive.
|
||||
// RunResult.Validation has Status ValidationFailed. When its output contract
|
||||
// has a positive repair budget, a failed eligible validation can make bounded
|
||||
// additional model calls and stops at the first valid candidate. Exhaustion
|
||||
// returns the final failed validation result with cumulative usage and actual
|
||||
// repair attempts. An inability to generate or validate returns an error and
|
||||
// no partial result.
|
||||
//
|
||||
// Run can return every error category documented by [Engine.Prepare], plus
|
||||
// ErrCapacityExceeded and ErrLLMGenerate. An engine admission rejection is
|
||||
// discoverable as [CapacityError] and still matches ErrCapacityExceeded. It
|
||||
// occurs before artifacts, schemas, rendering, or model generation because the
|
||||
// selected backend's admission capacity is full; it does not match
|
||||
// ErrInvalidRequest or ErrLLMGenerate. Errors from injected clients remain
|
||||
// available through errors.Is. Cancellation while waiting for model-generation
|
||||
// capacity matches both ErrLLMGenerate and the context error. Cancellation
|
||||
// otherwise follows the active collaborator's documented behavior. A nil
|
||||
// Engine returns ErrInvalidConfig. Run returns no partial result on error.
|
||||
// ErrInvalidRequest or ErrLLMGenerate. A built-in OpenAI-compatible non-2xx
|
||||
// response is discoverable as [GenerationError]. Errors from injected clients
|
||||
// remain available through errors.Is. Cancellation while waiting for
|
||||
// model-generation capacity matches both ErrLLMGenerate and the context error.
|
||||
// Cancellation otherwise follows the active collaborator's documented
|
||||
// behavior. A nil Engine returns ErrInvalidConfig. Run returns no partial
|
||||
// result on error.
|
||||
func (e *Engine) Run(ctx context.Context, req RunRequest) (*RunResult, error) {
|
||||
if e == nil || e.runner == nil {
|
||||
return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig)
|
||||
@@ -681,16 +703,17 @@ func (e *Engine) Run(ctx context.Context, req RunRequest) (*RunResult, error) {
|
||||
//
|
||||
// The supplied context governs this execution attempt independently of the
|
||||
// preparation context. It covers credential revalidation, admission,
|
||||
// generation, validation, and any internal repair. Result timing begins after
|
||||
// the claim and excludes preparation and consumer-held delay.
|
||||
// generation, validation, and any bounded output repair. Result timing begins
|
||||
// after the claim and excludes preparation and consumer-held delay.
|
||||
//
|
||||
// RunPrepared can return ErrInvalidRequest, ErrAPIKeyEnvMissing,
|
||||
// ErrCapacityExceeded, ErrLLMGenerate, or ErrValidation as applicable while
|
||||
// preserving documented collaborator and context identities. An engine
|
||||
// admission rejection is discoverable as [CapacityError] and still matches
|
||||
// ErrCapacityExceeded. A completed content-validation rejection is returned
|
||||
// in RunResult, not as an operational error. An operational error returns no
|
||||
// partial RunResult.
|
||||
// ErrCapacityExceeded. A built-in OpenAI-compatible non-2xx response is
|
||||
// discoverable as [GenerationError]. A completed content-validation rejection,
|
||||
// including repair exhaustion, is returned in RunResult, not as an operational
|
||||
// error. An operational error returns no partial RunResult.
|
||||
func (e *Engine) RunPrepared(ctx context.Context, prepared *PreparedExecution) (*RunResult, error) {
|
||||
if e == nil || e.runner == nil {
|
||||
return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig)
|
||||
|
||||
202
engine_test.go
202
engine_test.go
@@ -220,6 +220,14 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) {
|
||||
"transcript": promptkit.InlineWithURI(inputURI, inputBody),
|
||||
},
|
||||
Vars: map[string]string{"audience": variableValue},
|
||||
AppendedMessages: []promptkit.RenderedMessage{{
|
||||
Role: "private-role-sentinel",
|
||||
Content: "private-appended-content-sentinel",
|
||||
CacheControl: &promptkit.CacheControl{
|
||||
Type: "private-cache-sentinel",
|
||||
TTL: "private-ttl-sentinel",
|
||||
},
|
||||
}},
|
||||
}
|
||||
|
||||
for _, formatted := range []string{
|
||||
@@ -229,7 +237,7 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) {
|
||||
fmt.Sprintf("%+v", req),
|
||||
fmt.Sprintf("%#v", req),
|
||||
} {
|
||||
for _, privateValue := range []string{secret, inputURI, inputBody, variableValue} {
|
||||
for _, privateValue := range []string{secret, inputURI, inputBody, variableValue, "private-role-sentinel", "private-appended-content-sentinel", "private-cache-sentinel", "private-ttl-sentinel"} {
|
||||
if strings.Contains(formatted, privateValue) {
|
||||
t.Fatalf("formatted RunRequest leaked private value %q: %s", privateValue, formatted)
|
||||
}
|
||||
@@ -240,6 +248,7 @@ func TestRunRequestFormattingRedactsDirectAPIKey(t *testing.T) {
|
||||
"APIKeySet:true",
|
||||
"Inputs:1",
|
||||
"Vars:1",
|
||||
"AppendedMessages:1",
|
||||
} {
|
||||
if !strings.Contains(formatted, summary) {
|
||||
t.Fatalf("formatted RunRequest omitted structural summary %q: %s", summary, formatted)
|
||||
@@ -977,37 +986,72 @@ func TestPrepareDirectAPIKeyBypassesMissingEnvWithoutLeakingOrHashing(t *testing
|
||||
}
|
||||
}
|
||||
|
||||
func TestMissingCredentialsFailClearlyWhenProfileRequiresAuth(t *testing.T) {
|
||||
func TestOptionalMissingCredentialsReachUpstream(t *testing.T) {
|
||||
const missingEnv = "PROMPTKIT_PUBLIC_AUTH_MISSING"
|
||||
const providerBody = `{"error":{"message":"authentication failed","type":"authentication_error","code":"invalid_api_key"}}`
|
||||
t.Setenv(missingEnv, "")
|
||||
|
||||
profileDir := t.TempDir()
|
||||
writePublicProfileFileWithAPIKeyEnv(t, profileDir, "requires-auth", "http://localhost:8000/v1", "test-model", missingEnv)
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: frameworkPromptDir,
|
||||
ProfileDir: profileDir,
|
||||
SchemaDir: frameworkSchemaDir,
|
||||
})
|
||||
called := false
|
||||
config := promptkit.Config{
|
||||
PromptDir: frameworkPromptDir,
|
||||
SchemaDir: frameworkSchemaDir,
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
called = true
|
||||
if values := req.Header.Values("Authorization"); len(values) != 0 {
|
||||
t.Fatalf("Authorization values = %q, want absent", values)
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusUnauthorized,
|
||||
ContentLength: int64(len(providerBody)),
|
||||
Body: io.NopCloser(strings.NewReader(providerBody)),
|
||||
}, nil
|
||||
})},
|
||||
}
|
||||
engine, err := promptkit.NewEngine(config,
|
||||
promptkit.WithBackend(promptkit.Backend{
|
||||
ID: "optional-auth",
|
||||
Endpoint: "http://provider.test/v1",
|
||||
APIKeyEnv: missingEnv,
|
||||
}),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "optional-auth-profile",
|
||||
BackendID: "optional-auth",
|
||||
Model: "test-model",
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("expected engine construction to succeed, got %v", err)
|
||||
}
|
||||
|
||||
_, err = engine.Prepare(context.Background(), promptkit.RunRequest{
|
||||
result, err := engine.Run(context.Background(), promptkit.RunRequest{
|
||||
PromptID: frameworkMarkdownSummaryPromptID,
|
||||
ProfileID: "requires-auth",
|
||||
ProfileID: "optional-auth-profile",
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"transcript": promptkit.Inline("Rin opens the gate."),
|
||||
"glossary": promptkit.Inline("gate: A guarded passage."),
|
||||
},
|
||||
})
|
||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("expected invalid request for missing credentials, got %v", err)
|
||||
if !called {
|
||||
t.Fatal("optional missing credential did not reach upstream")
|
||||
}
|
||||
if !errors.Is(err, promptkit.ErrAPIKeyEnvMissing) {
|
||||
t.Fatalf("expected missing credential environment error, got %v", err)
|
||||
if result != nil {
|
||||
t.Fatalf("result = %+v, want nil", result)
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), missingEnv) {
|
||||
t.Fatalf("expected missing env name in error, got %v", err)
|
||||
if errors.Is(err, promptkit.ErrInvalidRequest) || errors.Is(err, promptkit.ErrAPIKeyEnvMissing) {
|
||||
t.Fatalf("error = %v, want upstream generation error without credential identities", err)
|
||||
}
|
||||
if !errors.Is(err, promptkit.ErrLLMGenerate) {
|
||||
t.Fatalf("error = %v, want ErrLLMGenerate", err)
|
||||
}
|
||||
var generationErr *promptkit.GenerationError
|
||||
if !errors.As(err, &generationErr) {
|
||||
t.Fatalf("error = %v, want GenerationError", err)
|
||||
}
|
||||
if generationErr.StatusCode() != http.StatusUnauthorized ||
|
||||
generationErr.ProviderType() != "authentication_error" ||
|
||||
generationErr.ProviderCode() != "invalid_api_key" ||
|
||||
generationErr.ProviderMessage() != "authentication failed" {
|
||||
t.Fatalf("GenerationError = %+v, want structured upstream authentication failure", generationErr)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1164,8 +1208,9 @@ func TestArtifactReaderFailuresPreserveArtifactLoadErrors(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRunAddsLLMGenerateToCollaboratorPublicError(t *testing.T) {
|
||||
injectedErr := errors.New("injected model client failure")
|
||||
engine := newContractEngineWithOptions(t, frameworkSchemaDir,
|
||||
promptkit.WithLLMClient(&fakeLLMClient{err: promptkit.ErrArtifactLoad}),
|
||||
promptkit.WithLLMClient(&fakeLLMClient{err: injectedErr}),
|
||||
)
|
||||
|
||||
_, err := engine.Run(context.Background(), promptkit.RunRequest{
|
||||
@@ -1178,8 +1223,12 @@ func TestRunAddsLLMGenerateToCollaboratorPublicError(t *testing.T) {
|
||||
if !errors.Is(err, promptkit.ErrLLMGenerate) {
|
||||
t.Fatalf("expected ErrLLMGenerate, got %v", err)
|
||||
}
|
||||
if !errors.Is(err, promptkit.ErrArtifactLoad) {
|
||||
t.Fatalf("expected preserved ErrArtifactLoad, got %v", err)
|
||||
if !errors.Is(err, injectedErr) {
|
||||
t.Fatalf("expected preserved injected error, got %v", err)
|
||||
}
|
||||
var generationErr *promptkit.GenerationError
|
||||
if errors.As(err, &generationErr) {
|
||||
t.Fatalf("injected error became GenerationError: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1492,43 +1541,51 @@ func TestSelectedProfileRepositoryReadFailureMapsToProfileLoad(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestPrepareUsesBuiltInProfileWithoutProfileDir(t *testing.T) {
|
||||
t.Setenv("OPENROUTER_API_KEY", "test-key")
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: frameworkPromptDir,
|
||||
SchemaDir: frameworkSchemaDir,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected engine construction to succeed, got %v", err)
|
||||
tests := []struct {
|
||||
name string
|
||||
profileID string
|
||||
backendID string
|
||||
}{
|
||||
{
|
||||
name: "OpenRouter",
|
||||
profileID: "mistral-small-3",
|
||||
backendID: promptkit.BackendOpenRouter,
|
||||
},
|
||||
{
|
||||
name: "Rakestrawhome",
|
||||
profileID: "rakestrawhome-gemma-4-31b",
|
||||
backendID: promptkit.BackendRakestrawHome,
|
||||
},
|
||||
}
|
||||
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{
|
||||
PromptID: frameworkMarkdownSummaryPromptID,
|
||||
ProfileID: "mistral-small-3",
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"transcript": promptkit.Inline("Rin opens the gate."),
|
||||
"glossary": promptkit.Inline("gate: A guarded passage."),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected built-in profile prepare to succeed, got %v", err)
|
||||
}
|
||||
if prepared.SelectedProfileID != "mistral-small-3" {
|
||||
t.Fatalf("unexpected selected profile: %q", prepared.SelectedProfileID)
|
||||
}
|
||||
if prepared.SelectedBackendID != promptkit.BackendOpenRouter {
|
||||
t.Fatalf("unexpected selected backend: %q", prepared.SelectedBackendID)
|
||||
}
|
||||
if prepared.EffectiveModelParams.BackendID != promptkit.BackendOpenRouter {
|
||||
t.Fatalf("unexpected effective backend: %q", prepared.EffectiveModelParams.BackendID)
|
||||
}
|
||||
if prepared.EffectiveModelParams.Endpoint != "https://openrouter.ai/api/v1" {
|
||||
t.Fatalf("unexpected built-in endpoint: %q", prepared.EffectiveModelParams.Endpoint)
|
||||
}
|
||||
if prepared.EffectiveModelParams.APIKeyEnv != "OPENROUTER_API_KEY" {
|
||||
t.Fatalf("unexpected built-in api key environment name: %q", prepared.EffectiveModelParams.APIKeyEnv)
|
||||
}
|
||||
if prepared.EffectiveModelParams.Model != "mistralai/mistral-small-3.2-24b-instruct" {
|
||||
t.Fatalf("unexpected built-in model: %q", prepared.EffectiveModelParams.Model)
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: frameworkPromptDir,
|
||||
SchemaDir: frameworkSchemaDir,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected engine construction to succeed, got %v", err)
|
||||
}
|
||||
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{
|
||||
PromptID: frameworkMarkdownSummaryPromptID,
|
||||
ProfileID: tc.profileID,
|
||||
APIKey: "test-key",
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"transcript": promptkit.Inline("Rin opens the gate."),
|
||||
"glossary": promptkit.Inline("gate: A guarded passage."),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected built-in profile prepare to succeed, got %v", err)
|
||||
}
|
||||
if prepared.SelectedProfileID != tc.profileID ||
|
||||
prepared.SelectedBackendID != tc.backendID ||
|
||||
prepared.EffectiveModelParams.BackendID != tc.backendID {
|
||||
t.Fatalf("unexpected built-in preparation: %#v", prepared)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1958,6 +2015,25 @@ func TestWithProfilesRejectsDuplicateIDs(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithProfilesAcceptsDerivedDefinitionWithoutTargetFields(t *testing.T) {
|
||||
_, err := promptkit.NewEngine(promptkit.Config{PromptDir: frameworkPromptDir},
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: " derived-profile ",
|
||||
BaseProfileID: " base-profile ",
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("derived profile should be accepted during construction: %v", err)
|
||||
}
|
||||
|
||||
_, err = promptkit.NewEngine(promptkit.Config{PromptDir: frameworkPromptDir},
|
||||
promptkit.WithProfiles(promptkit.Profile{ID: "standalone-profile"}),
|
||||
)
|
||||
if !errors.Is(err, promptkit.ErrInvalidConfig) {
|
||||
t.Fatalf("standalone incomplete profile error = %v, want ErrInvalidConfig", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithProfilesRejectsInvalidExecutionSettings(t *testing.T) {
|
||||
type testCase struct {
|
||||
name string
|
||||
@@ -2110,6 +2186,7 @@ func TestOpenAICompatibleProfileMapsEveryField(t *testing.T) {
|
||||
extraParams := map[string]any{"provider_option": "distinct-extra-params"}
|
||||
got := promptkit.OpenAICompatibleProfile(promptkit.OpenAICompatibleProfileConfig{
|
||||
ID: "distinct-id",
|
||||
BaseProfileID: "distinct-base",
|
||||
BackendID: "distinct-backend",
|
||||
Endpoint: "https://distinct.example/v1",
|
||||
Model: "distinct-model",
|
||||
@@ -2124,6 +2201,7 @@ func TestOpenAICompatibleProfileMapsEveryField(t *testing.T) {
|
||||
})
|
||||
want := promptkit.Profile{
|
||||
ID: "distinct-id",
|
||||
BaseProfileID: "distinct-base",
|
||||
BackendID: "distinct-backend",
|
||||
Endpoint: "https://distinct.example/v1",
|
||||
Model: "distinct-model",
|
||||
@@ -3405,9 +3483,10 @@ extra_params:
|
||||
}
|
||||
|
||||
type fakeLLMClient struct {
|
||||
response *promptkit.GenerateResponse
|
||||
err error
|
||||
requests []promptkit.GenerateRequest
|
||||
response *promptkit.GenerateResponse
|
||||
responses []*promptkit.GenerateResponse
|
||||
err error
|
||||
requests []promptkit.GenerateRequest
|
||||
}
|
||||
|
||||
type recordingArtifactReader struct {
|
||||
@@ -3583,5 +3662,12 @@ func (f *fakeLLMClient) Generate(_ context.Context, req promptkit.GenerateReques
|
||||
if f.err != nil {
|
||||
return nil, f.err
|
||||
}
|
||||
if len(f.responses) > 0 {
|
||||
index := len(f.requests) - 1
|
||||
if index >= len(f.responses) {
|
||||
return nil, fmt.Errorf("no response configured for generation %d", index+1)
|
||||
}
|
||||
return f.responses[index], nil
|
||||
}
|
||||
return f.response, nil
|
||||
}
|
||||
|
||||
14
errors.go
14
errors.go
@@ -6,6 +6,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/capacity"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/promptdef"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/usecase"
|
||||
@@ -21,6 +22,19 @@ func mapPublicError(err error) error {
|
||||
return &CapacityError{BackendID: internalCapacityError.BackendID}
|
||||
}
|
||||
publicErr := publicErrorFor(err)
|
||||
var providerHTTPError *llm.ProviderHTTPError
|
||||
if errors.As(err, &providerHTTPError) && providerHTTPError != nil {
|
||||
generationErr := newGenerationError(
|
||||
providerHTTPError.StatusCode(),
|
||||
providerHTTPError.ProviderCode(),
|
||||
providerHTTPError.ProviderType(),
|
||||
providerHTTPError.ProviderMessage(),
|
||||
)
|
||||
if publicErr != nil && !errors.Is(publicErr, ErrLLMGenerate) {
|
||||
return fmt.Errorf("%w: %w", publicErr, generationErr)
|
||||
}
|
||||
return generationErr
|
||||
}
|
||||
if publicErr == nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/usecase"
|
||||
)
|
||||
|
||||
@@ -48,3 +49,27 @@ func TestMapPublicErrorTranslatesCapacityError(t *testing.T) {
|
||||
t.Fatalf("mapped backend ID changed with source error: %q", publicErr.BackendID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapPublicErrorPreservesValidationAroundGenerationError(t *testing.T) {
|
||||
internalErr := fmt.Errorf(
|
||||
"%w: %w",
|
||||
usecase.ErrValidation,
|
||||
&llm.ProviderHTTPError{},
|
||||
)
|
||||
|
||||
err := mapPublicError(internalErr)
|
||||
if !errors.Is(err, ErrValidation) {
|
||||
t.Fatalf("mapped error=%v, want ErrValidation", err)
|
||||
}
|
||||
if !errors.Is(err, ErrLLMGenerate) {
|
||||
t.Fatalf("mapped error=%v, want ErrLLMGenerate", err)
|
||||
}
|
||||
var generationErr *GenerationError
|
||||
if !errors.As(err, &generationErr) || generationErr == nil {
|
||||
t.Fatalf("mapped error=%v, want GenerationError", err)
|
||||
}
|
||||
var leakedInternalErr *llm.ProviderHTTPError
|
||||
if errors.As(err, &leakedInternalErr) {
|
||||
t.Fatalf("mapped error exposes internal ProviderHTTPError: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,7 +18,7 @@ func (r RunRequest) GoString() string {
|
||||
|
||||
func (r RunRequest) redactedString() string {
|
||||
return fmt.Sprintf(
|
||||
"promptkit.RunRequest{PromptID:%q PromptVersion:%q ProfileID:%q APIKeySet:%t Inputs:%d Vars:%d ExecutionSet:%t ValidationSet:%t}",
|
||||
"promptkit.RunRequest{PromptID:%q PromptVersion:%q ProfileID:%q APIKeySet:%t Inputs:%d Vars:%d ExecutionSet:%t ValidationSet:%t AppendedMessages:%d}",
|
||||
r.PromptID,
|
||||
r.PromptVersion,
|
||||
r.ProfileID,
|
||||
@@ -27,6 +27,7 @@ func (r RunRequest) redactedString() string {
|
||||
len(r.Vars),
|
||||
r.Execution != nil,
|
||||
r.Validation != nil,
|
||||
len(r.AppendedMessages),
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
87
generation_error.go
Normal file
87
generation_error.go
Normal file
@@ -0,0 +1,87 @@
|
||||
package promptkit
|
||||
|
||||
import "fmt"
|
||||
|
||||
// GenerationError reports a non-2xx response from Promptkit's built-in
|
||||
// OpenAI-compatible client during [Engine.Run] or [Engine.RunPrepared].
|
||||
//
|
||||
// Engine-produced values are immutable, caller-owned values. Use errors.Is to
|
||||
// match [ErrLLMGenerate] and errors.As with a *GenerationError target to obtain
|
||||
// this type. The four provider accessors expose untrusted provider-controlled
|
||||
// values that can contain sensitive request or schema fragments. Applications
|
||||
// must apply their own disclosure policy before logging, displaying, or
|
||||
// returning them to another caller.
|
||||
//
|
||||
// Accessors, Error, GoString, and Unwrap are safe on a nil receiver and a zero
|
||||
// value. Default and Go-syntax formatting deliberately redact provider details.
|
||||
// GenerationError has no stable JSON representation.
|
||||
type GenerationError struct {
|
||||
statusCode int
|
||||
providerCode string
|
||||
providerType string
|
||||
providerMessage string
|
||||
}
|
||||
|
||||
func newGenerationError(statusCode int, providerCode, providerType, providerMessage string) *GenerationError {
|
||||
return &GenerationError{
|
||||
statusCode: statusCode,
|
||||
providerCode: providerCode,
|
||||
providerType: providerType,
|
||||
providerMessage: providerMessage,
|
||||
}
|
||||
}
|
||||
|
||||
// StatusCode returns the received provider HTTP status code, or zero for a nil
|
||||
// receiver or zero value.
|
||||
func (e *GenerationError) StatusCode() int {
|
||||
if e == nil {
|
||||
return 0
|
||||
}
|
||||
return e.statusCode
|
||||
}
|
||||
|
||||
// ProviderCode returns the normalized provider error code, if present. Its
|
||||
// value is untrusted and may contain sensitive data.
|
||||
func (e *GenerationError) ProviderCode() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
return e.providerCode
|
||||
}
|
||||
|
||||
// ProviderType returns the normalized provider error type, if present. Its
|
||||
// value is untrusted and may contain sensitive data.
|
||||
func (e *GenerationError) ProviderType() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
return e.providerType
|
||||
}
|
||||
|
||||
// ProviderMessage returns the bounded normalized provider diagnostic, if
|
||||
// present. Its value is untrusted and may contain sensitive data.
|
||||
func (e *GenerationError) ProviderMessage() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
return e.providerMessage
|
||||
}
|
||||
|
||||
// Error returns a redacted diagnostic that is not a parsing contract.
|
||||
func (e *GenerationError) Error() string {
|
||||
if e == nil || e.statusCode == 0 {
|
||||
return ErrLLMGenerate.Error()
|
||||
}
|
||||
return fmt.Sprintf("%s: provider returned HTTP status %d", ErrLLMGenerate, e.statusCode)
|
||||
}
|
||||
|
||||
// GoString returns the same redacted diagnostic as Error.
|
||||
func (e *GenerationError) GoString() string {
|
||||
return e.Error()
|
||||
}
|
||||
|
||||
// Unwrap returns ErrLLMGenerate. It is safe to call on a nil receiver or zero
|
||||
// value.
|
||||
func (e *GenerationError) Unwrap() error {
|
||||
return ErrLLMGenerate
|
||||
}
|
||||
150
generation_error_contract_test.go
Normal file
150
generation_error_contract_test.go
Normal file
@@ -0,0 +1,150 @@
|
||||
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 TestBuiltInGenerationErrorWithAppendedMessages(t *testing.T) {
|
||||
const messageMarker = "combined-request-provider-marker"
|
||||
engine := newBuiltInGenerationErrorEngine(t, http.StatusUnprocessableEntity,
|
||||
`{"error":{"message":"`+messageMarker+`"}}`)
|
||||
request := generationErrorRunRequest()
|
||||
request.AppendedMessages = []promptkit.RenderedMessage{
|
||||
{Role: promptkit.RoleAssistant, Content: "previous response"},
|
||||
{Role: promptkit.RoleUser, Content: "consumer correction"},
|
||||
}
|
||||
|
||||
result, err := engine.Run(context.Background(), request)
|
||||
if result != nil {
|
||||
t.Fatalf("Run result = %#v, want nil", result)
|
||||
}
|
||||
assertGenerationError(t, err, http.StatusUnprocessableEntity, "", "", messageMarker)
|
||||
}
|
||||
|
||||
func TestBuiltInRepairGenerationError(t *testing.T) {
|
||||
const (
|
||||
codeMarker = "repair-code-marker"
|
||||
typeMarker = "repair-type-marker"
|
||||
messageMarker = "repair-message-marker"
|
||||
)
|
||||
calls := 0
|
||||
config := contractConfig(frameworkSchemaDir)
|
||||
config.HTTPClient = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
calls++
|
||||
if calls == 1 {
|
||||
body := `{"choices":[{"message":{"content":"not-json"}}]}`
|
||||
return &http.Response{StatusCode: http.StatusOK, ContentLength: int64(len(body)), Body: io.NopCloser(strings.NewReader(body))}, nil
|
||||
}
|
||||
body := `{"error":{"code":"` + codeMarker + `","type":"` + typeMarker + `","message":"` + messageMarker + `"}}`
|
||||
return &http.Response{StatusCode: http.StatusUnprocessableEntity, ContentLength: int64(len(body)), Body: io.NopCloser(strings.NewReader(body))}, nil
|
||||
})}
|
||||
engine, err := promptkit.NewEngine(config)
|
||||
if err != nil {
|
||||
t.Fatalf("NewEngine: %v", err)
|
||||
}
|
||||
req := generationErrorRunRequest()
|
||||
req.Validation = &promptkit.OutputContract{
|
||||
Format: promptkit.FormatJSON,
|
||||
ValidationMode: promptkit.ValidationJSON,
|
||||
RepairAttempts: 1,
|
||||
}
|
||||
|
||||
result, err := engine.Run(context.Background(), req)
|
||||
if result != nil {
|
||||
t.Fatalf("Run result = %#v, want nil", result)
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls)
|
||||
}
|
||||
assertGenerationError(t, err, http.StatusUnprocessableEntity, codeMarker, typeMarker, messageMarker)
|
||||
}
|
||||
|
||||
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."),
|
||||
},
|
||||
}
|
||||
}
|
||||
32
generation_error_internal_test.go
Normal file
32
generation_error_internal_test.go
Normal file
@@ -0,0 +1,32 @@
|
||||
package promptkit
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGenerationErrorNilAndZeroValue(t *testing.T) {
|
||||
var nilError *GenerationError
|
||||
zeroError := &GenerationError{}
|
||||
|
||||
for name, err := range map[string]*GenerationError{
|
||||
"nil": nilError,
|
||||
"zero": zeroError,
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if err.StatusCode() != 0 || err.ProviderCode() != "" || err.ProviderType() != "" || err.ProviderMessage() != "" {
|
||||
t.Fatalf("accessors returned provider details: %#v", err)
|
||||
}
|
||||
if err.Error() != "failed to generate output" || err.GoString() != "failed to generate output" {
|
||||
t.Fatalf("redacted formatting = (%q, %q)", err.Error(), err.GoString())
|
||||
}
|
||||
if fmt.Sprintf("%v", err) != "failed to generate output" || fmt.Sprintf("%#v", err) != "failed to generate output" {
|
||||
t.Fatalf("formatted error = (%q, %q)", fmt.Sprintf("%v", err), fmt.Sprintf("%#v", err))
|
||||
}
|
||||
if !errors.Is(err, ErrLLMGenerate) {
|
||||
t.Fatalf("errors.Is(%v, ErrLLMGenerate) = false", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
2
go.mod
2
go.mod
@@ -3,6 +3,8 @@ module gitea.maximumdirect.net/eric/promptkit
|
||||
go 1.25.5
|
||||
|
||||
require (
|
||||
gitea.maximumdirect.net/eric/promptkit-backend-openrouter v1.0.0
|
||||
gitea.maximumdirect.net/eric/promptkit-backend-rakestrawhome v1.0.0
|
||||
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2
|
||||
gopkg.in/yaml.v3 v3.0.1
|
||||
)
|
||||
|
||||
4
go.sum
4
go.sum
@@ -1,3 +1,7 @@
|
||||
gitea.maximumdirect.net/eric/promptkit-backend-openrouter v1.0.0 h1:lc062euk2qseO//D762i3JaFyulDNML3eQQX7DkYTho=
|
||||
gitea.maximumdirect.net/eric/promptkit-backend-openrouter v1.0.0/go.mod h1:AIa7kAu2mfrRQgcspe4L+DW51WqgnALQT60lqkEywJI=
|
||||
gitea.maximumdirect.net/eric/promptkit-backend-rakestrawhome v1.0.0 h1:j9YY7wsTVjzke2kHH4YAzpU0oUpM+x+nXwl1IeS+2eg=
|
||||
gitea.maximumdirect.net/eric/promptkit-backend-rakestrawhome v1.0.0/go.mod h1:4RNS+LILDg4JbS4Ts9Lwy1C92wauXJIbeQaalps4Koo=
|
||||
github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI=
|
||||
github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
|
||||
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2 h1:KRzFb2m7YtdldCEkzs6KqmJw4nqEVZGK7IN2kJkjTuQ=
|
||||
|
||||
@@ -18,12 +18,11 @@ const (
|
||||
// OpenRouterID is the reserved ID of Promptkit's built-in OpenRouter
|
||||
// backend.
|
||||
OpenRouterID = "openrouter"
|
||||
// RakestrawHomeID is the reserved ID of Promptkit's built-in Rakestrawhome
|
||||
// backend.
|
||||
RakestrawHomeID = "rakestrawhome"
|
||||
|
||||
openRouterEndpoint = "https://openrouter.ai/api/v1"
|
||||
openRouterAPIKeyEnv = "OPENROUTER_API_KEY"
|
||||
|
||||
openRouterConcurrencyLimit = 16
|
||||
defaultQueueCapacity = 1024
|
||||
defaultQueueCapacity = 1024
|
||||
)
|
||||
|
||||
// ErrBackendNotFound identifies a registry lookup for an unknown backend ID.
|
||||
@@ -36,20 +35,15 @@ type Registry struct {
|
||||
backends map[string]domain.Backend
|
||||
}
|
||||
|
||||
// NewRegistry constructs a registry containing the built-in OpenRouter
|
||||
// definition followed by the supplied additions. Every ID must be unique.
|
||||
func NewRegistry(additions []domain.Backend) (*Registry, error) {
|
||||
// NewRegistry constructs a registry containing maintained definitions followed
|
||||
// by consumer additions. Every ID must be unique across both groups.
|
||||
func NewRegistry(maintained, additions []domain.Backend) (*Registry, error) {
|
||||
registry := &Registry{
|
||||
backends: make(map[string]domain.Backend, len(additions)+1),
|
||||
backends: make(map[string]domain.Backend, len(maintained)+len(additions)),
|
||||
}
|
||||
|
||||
definitions := make([]domain.Backend, 0, len(additions)+1)
|
||||
definitions = append(definitions, domain.Backend{
|
||||
ID: OpenRouterID,
|
||||
Endpoint: openRouterEndpoint,
|
||||
APIKeyEnv: openRouterAPIKeyEnv,
|
||||
ConcurrencyLimit: openRouterConcurrencyLimit,
|
||||
})
|
||||
definitions := make([]domain.Backend, 0, len(maintained)+len(additions))
|
||||
definitions = append(definitions, maintained...)
|
||||
definitions = append(definitions, additions...)
|
||||
|
||||
for _, definition := range definitions {
|
||||
@@ -61,7 +55,7 @@ func NewRegistry(additions []domain.Backend) (*Registry, error) {
|
||||
return nil, fmt.Errorf("backend ID %q is already registered", definition.ID)
|
||||
}
|
||||
|
||||
normalized, err := normalizeBackend(definition)
|
||||
normalized, err := NormalizeDefinition(definition)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -107,7 +101,8 @@ func (r *Registry) CapacityPolicies() map[string]domain.BackendCapacityPolicy {
|
||||
return policies
|
||||
}
|
||||
|
||||
func normalizeBackend(definition domain.Backend) (domain.Backend, error) {
|
||||
// NormalizeDefinition validates and defensively copies one backend definition.
|
||||
func NormalizeDefinition(definition domain.Backend) (domain.Backend, error) {
|
||||
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(definition.Endpoint)
|
||||
if err != nil {
|
||||
return domain.Backend{}, fmt.Errorf("backend %q endpoint: %w", definition.ID, err)
|
||||
|
||||
@@ -2,6 +2,7 @@ package backend_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -11,32 +12,39 @@ import (
|
||||
|
||||
const validEndpoint = "https://backend.example/v1"
|
||||
|
||||
func TestRegistryIncludesExactOpenRouterDefinition(t *testing.T) {
|
||||
registry, err := backend.NewRegistry(nil)
|
||||
func TestRegistryIncludesMaintainedDefinitions(t *testing.T) {
|
||||
maintained := []domain.Backend{
|
||||
{ID: backend.OpenRouterID, Endpoint: validEndpoint, ConcurrencyLimit: 2, QueueCapacity: 3, QueueCapacitySet: true},
|
||||
{ID: backend.RakestrawHomeID, Endpoint: "https://second.example/v1", ConcurrencyLimit: 4, QueueCapacity: 5, QueueCapacitySet: true},
|
||||
}
|
||||
registry, err := backend.NewRegistry(maintained, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("construct registry: %v", err)
|
||||
}
|
||||
|
||||
definition, err := registry.GetBackend(backend.OpenRouterID)
|
||||
if err != nil {
|
||||
t.Fatalf("look up OpenRouter: %v", err)
|
||||
}
|
||||
if definition.ID != "openrouter" ||
|
||||
definition.Endpoint != "https://openrouter.ai/api/v1" ||
|
||||
definition.APIKeyEnv != "OPENROUTER_API_KEY" ||
|
||||
definition.ConcurrencyLimit != 16 ||
|
||||
definition.QueueCapacity != 1024 ||
|
||||
!definition.QueueCapacitySet ||
|
||||
definition.ExtraParams != nil {
|
||||
t.Fatalf("unexpected OpenRouter definition: %#v", definition)
|
||||
for _, expected := range maintained {
|
||||
t.Run(expected.ID, func(t *testing.T) {
|
||||
definition, err := registry.GetBackend(expected.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("look up maintained definition: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(definition, expected) {
|
||||
t.Fatalf("unexpected maintained definition: %#v", definition)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
policies := registry.CapacityPolicies()
|
||||
if len(policies) != 1 ||
|
||||
policies["openrouter"] != (domain.BackendCapacityPolicy{
|
||||
ConcurrencyLimit: 16,
|
||||
QueueCapacity: 1024,
|
||||
if len(policies) != 2 ||
|
||||
policies[backend.OpenRouterID] != (domain.BackendCapacityPolicy{
|
||||
ConcurrencyLimit: 2,
|
||||
QueueCapacity: 3,
|
||||
}) ||
|
||||
policies[backend.RakestrawHomeID] != (domain.BackendCapacityPolicy{
|
||||
ConcurrencyLimit: 4,
|
||||
QueueCapacity: 5,
|
||||
}) {
|
||||
t.Fatalf("unexpected OpenRouter capacity policies: %#v", policies)
|
||||
t.Fatalf("unexpected built-in capacity policies: %#v", policies)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -46,7 +54,7 @@ func TestRegistryNormalizesUniqueAdditionsAndIsolatesMutations(t *testing.T) {
|
||||
"count": int64(7),
|
||||
"nested": nested,
|
||||
}
|
||||
registry, err := backend.NewRegistry([]domain.Backend{
|
||||
registry, err := backend.NewRegistry(nil, []domain.Backend{
|
||||
{
|
||||
ID: " custom ",
|
||||
Endpoint: " https://custom.example/openai/v1 ",
|
||||
@@ -109,11 +117,11 @@ func TestRegistryNormalizesUniqueAdditionsAndIsolatesMutations(t *testing.T) {
|
||||
}
|
||||
|
||||
policies := registry.CapacityPolicies()
|
||||
if len(policies) != 2 {
|
||||
if len(policies) != 1 {
|
||||
t.Fatalf("unexpected capacity policy count: %#v", policies)
|
||||
}
|
||||
policies["custom"] = domain.BackendCapacityPolicy{}
|
||||
delete(policies, backend.OpenRouterID)
|
||||
delete(policies, "custom")
|
||||
againPolicies := registry.CapacityPolicies()
|
||||
if againPolicies["custom"] != (domain.BackendCapacityPolicy{
|
||||
ConcurrencyLimit: 3,
|
||||
@@ -121,9 +129,6 @@ func TestRegistryNormalizesUniqueAdditionsAndIsolatesMutations(t *testing.T) {
|
||||
}) {
|
||||
t.Fatalf("capacity policy map mutated registry state: %#v", againPolicies)
|
||||
}
|
||||
if _, ok := againPolicies[backend.OpenRouterID]; !ok {
|
||||
t.Fatalf("capacity policy deletion mutated registry state: %#v", againPolicies)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewRegistryNormalizesCapacityPolicy(t *testing.T) {
|
||||
@@ -199,7 +204,7 @@ func TestNewRegistryNormalizesCapacityPolicy(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
tc.definition.ID = "custom"
|
||||
tc.definition.Endpoint = validEndpoint
|
||||
registry, err := backend.NewRegistry([]domain.Backend{tc.definition})
|
||||
registry, err := backend.NewRegistry(nil, []domain.Backend{tc.definition})
|
||||
if tc.wantError {
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid capacity policy error")
|
||||
@@ -242,11 +247,14 @@ func TestNewRegistryRejectsDuplicateIDs(t *testing.T) {
|
||||
wantID string
|
||||
}{
|
||||
{
|
||||
name: "built-in collision after normalization",
|
||||
additions: []domain.Backend{{
|
||||
ID: " openrouter ",
|
||||
}},
|
||||
wantID: "openrouter",
|
||||
name: "OpenRouter collision after normalization",
|
||||
additions: []domain.Backend{{ID: " openrouter "}},
|
||||
wantID: backend.OpenRouterID,
|
||||
},
|
||||
{
|
||||
name: "Rakestrawhome collision after normalization",
|
||||
additions: []domain.Backend{{ID: " rakestrawhome "}},
|
||||
wantID: backend.RakestrawHomeID,
|
||||
},
|
||||
{
|
||||
name: "consumer collision after normalization",
|
||||
@@ -260,7 +268,7 @@ func TestNewRegistryRejectsDuplicateIDs(t *testing.T) {
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := backend.NewRegistry(tc.additions)
|
||||
_, err := backend.NewRegistry(nil, tc.additions)
|
||||
if err == nil {
|
||||
t.Fatal("expected duplicate ID error")
|
||||
}
|
||||
@@ -274,7 +282,7 @@ func TestNewRegistryRejectsDuplicateIDs(t *testing.T) {
|
||||
func TestNewRegistryValidatesIDs(t *testing.T) {
|
||||
for _, id := range []string{"", " \t\n "} {
|
||||
t.Run(id, func(t *testing.T) {
|
||||
_, err := backend.NewRegistry([]domain.Backend{{
|
||||
_, err := backend.NewRegistry(nil, []domain.Backend{{
|
||||
ID: id,
|
||||
Endpoint: validEndpoint,
|
||||
}})
|
||||
@@ -303,7 +311,7 @@ func TestNewRegistryValidatesEndpoints(t *testing.T) {
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := backend.NewRegistry([]domain.Backend{{
|
||||
_, err := backend.NewRegistry(nil, []domain.Backend{{
|
||||
ID: "custom",
|
||||
Endpoint: tc.endpoint,
|
||||
}})
|
||||
@@ -317,7 +325,7 @@ func TestNewRegistryValidatesEndpoints(t *testing.T) {
|
||||
func TestNewRegistryValidatesEnvironmentVariableNames(t *testing.T) {
|
||||
for _, name := range []string{"1API_KEY", "API-KEY", "API KEY", "ÅPI_KEY"} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
_, err := backend.NewRegistry([]domain.Backend{{
|
||||
_, err := backend.NewRegistry(nil, []domain.Backend{{
|
||||
ID: "custom",
|
||||
Endpoint: validEndpoint,
|
||||
APIKeyEnv: name,
|
||||
@@ -340,7 +348,7 @@ func TestNewRegistryRejectsInvalidAndReservedExtraParameters(t *testing.T) {
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := backend.NewRegistry([]domain.Backend{{
|
||||
_, err := backend.NewRegistry(nil, []domain.Backend{{
|
||||
ID: "custom",
|
||||
Endpoint: validEndpoint,
|
||||
ExtraParams: tc.extraParams,
|
||||
@@ -353,7 +361,7 @@ func TestNewRegistryRejectsInvalidAndReservedExtraParameters(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRegistryLookupReportsNotFound(t *testing.T) {
|
||||
registry, err := backend.NewRegistry(nil)
|
||||
registry, err := backend.NewRegistry(nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("construct registry: %v", err)
|
||||
}
|
||||
|
||||
246
internal/catalog/catalog.go
Normal file
246
internal/catalog/catalog.go
Normal file
@@ -0,0 +1,246 @@
|
||||
// Package catalog validates immutable maintained backend catalog assets.
|
||||
package catalog
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"path"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/backend"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
)
|
||||
|
||||
// Source identifies one immutable backend catalog asset tree.
|
||||
type Source struct {
|
||||
Name string
|
||||
ExpectedBackendID string
|
||||
FS fs.FS
|
||||
Root string
|
||||
}
|
||||
|
||||
// Set is the validated maintained backend and raw profile catalog.
|
||||
type Set struct {
|
||||
Backends []domain.Backend
|
||||
Profiles profile.Repository
|
||||
profileIDs []string
|
||||
}
|
||||
|
||||
// Load validates and combines immutable catalog sources in source order.
|
||||
func Load(sources ...Source) (Set, error) {
|
||||
if len(sources) == 0 {
|
||||
return Set{}, errors.New("at least one catalog source is required")
|
||||
}
|
||||
loaded := Set{Backends: make([]domain.Backend, 0, len(sources))}
|
||||
names := map[string]bool{}
|
||||
backendIDs := map[string]bool{}
|
||||
profileIDs := map[string]bool{}
|
||||
for _, source := range sources {
|
||||
if err := validateSource(source, names); err != nil {
|
||||
return Set{}, err
|
||||
}
|
||||
names[source.Name] = true
|
||||
definition, err := loadBackend(source)
|
||||
if err != nil {
|
||||
return Set{}, err
|
||||
}
|
||||
if backendIDs[definition.ID] {
|
||||
return Set{}, fmt.Errorf("catalog %s: backend ID duplicates an earlier catalog", source.Name)
|
||||
}
|
||||
backendIDs[definition.ID] = true
|
||||
repository, metadata, err := profile.LoadFSRepository(context.Background(), source.FS, path.Join(source.Root, "profiles"))
|
||||
if err != nil {
|
||||
return Set{}, fmt.Errorf("catalog %s: %w", source.Name, err)
|
||||
}
|
||||
if len(metadata) == 0 {
|
||||
return Set{}, fmt.Errorf("catalog %s: profiles must not be empty", source.Name)
|
||||
}
|
||||
resolvingRepository := profile.NewResolvingRepository(repository)
|
||||
for _, entry := range metadata {
|
||||
if profileIDs[entry.ID] {
|
||||
return Set{}, fmt.Errorf("catalog %s: %s duplicates an earlier profile ID", source.Name, entry.Path)
|
||||
}
|
||||
if containsField(entry.ExplicitFields, "endpoint") || containsField(entry.ExplicitFields, "api_key_env") {
|
||||
return Set{}, fmt.Errorf("catalog %s: %s contains connection metadata", source.Name, entry.Path)
|
||||
}
|
||||
value, err := repository.GetProfile(context.Background(), entry.ID)
|
||||
if err != nil {
|
||||
return Set{}, fmt.Errorf("catalog %s: %s: %w", source.Name, entry.Path, err)
|
||||
}
|
||||
if err := rejectSecretKeys(value.ExtraParams); err != nil {
|
||||
return Set{}, fmt.Errorf("catalog %s: %s: prohibited extra parameter key", source.Name, entry.Path)
|
||||
}
|
||||
resolved, err := resolvingRepository.GetProfile(context.Background(), entry.ID)
|
||||
if err != nil {
|
||||
return Set{}, fmt.Errorf("catalog %s: %s has invalid profile inheritance", source.Name, entry.Path)
|
||||
}
|
||||
if resolved.BackendID != definition.ID {
|
||||
return Set{}, fmt.Errorf("catalog %s: %s selects a different backend", source.Name, entry.Path)
|
||||
}
|
||||
profileIDs[entry.ID] = true
|
||||
loaded.profileIDs = append(loaded.profileIDs, entry.ID)
|
||||
}
|
||||
if loaded.Profiles == nil {
|
||||
loaded.Profiles = repository
|
||||
} else {
|
||||
loaded.Profiles = profile.NewOverlayRepository(loaded.Profiles, repository)
|
||||
}
|
||||
loaded.Backends = append(loaded.Backends, definition)
|
||||
}
|
||||
sort.Strings(loaded.profileIDs)
|
||||
return loaded, nil
|
||||
}
|
||||
|
||||
func validateSource(source Source, names map[string]bool) error {
|
||||
if strings.TrimSpace(source.Name) == "" || source.Name != strings.TrimSpace(source.Name) {
|
||||
return errors.New("catalog source name must not be blank")
|
||||
}
|
||||
if names[source.Name] {
|
||||
return fmt.Errorf("duplicate catalog source name %q", source.Name)
|
||||
}
|
||||
if source.FS == nil {
|
||||
return fmt.Errorf("catalog %s: filesystem is nil", source.Name)
|
||||
}
|
||||
if source.Root == "." || source.Root == "" || !fs.ValidPath(source.Root) {
|
||||
return fmt.Errorf("catalog %s: asset root is invalid", source.Name)
|
||||
}
|
||||
if strings.TrimSpace(source.ExpectedBackendID) == "" {
|
||||
return fmt.Errorf("catalog %s: expected backend ID is blank", source.Name)
|
||||
}
|
||||
return validateLayout(source)
|
||||
}
|
||||
|
||||
func validateLayout(source Source) error {
|
||||
manifestPath := path.Join(source.Root, "backend.json")
|
||||
profilesRoot := path.Join(source.Root, "profiles")
|
||||
return fs.WalkDir(source.FS, source.Root, func(assetPath string, entry fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if assetPath == source.Root {
|
||||
if !entry.IsDir() {
|
||||
return fmt.Errorf("catalog %s: invalid asset path %s", source.Name, assetPath)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if entry.IsDir() {
|
||||
if assetPath == profilesRoot || strings.HasPrefix(assetPath, profilesRoot+"/") {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("catalog %s: invalid asset path %s", source.Name, assetPath)
|
||||
}
|
||||
if assetPath == manifestPath && entry.Type().IsRegular() {
|
||||
return nil
|
||||
}
|
||||
if strings.HasPrefix(assetPath, profilesRoot+"/") && entry.Type().IsRegular() && strings.HasSuffix(assetPath, ".yml") {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("catalog %s: invalid asset path %s", source.Name, assetPath)
|
||||
})
|
||||
}
|
||||
|
||||
func loadBackend(source Source) (domain.Backend, error) {
|
||||
data, err := fs.ReadFile(source.FS, path.Join(source.Root, "backend.json"))
|
||||
if err != nil {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: %w", source.Name, err)
|
||||
}
|
||||
type manifest struct {
|
||||
SchemaVersion *int `json:"schema_version"`
|
||||
ID *string `json:"id"`
|
||||
Endpoint *string `json:"endpoint"`
|
||||
APIKeyEnv *string `json:"api_key_env"`
|
||||
ConcurrencyLimit *int `json:"concurrency_limit"`
|
||||
QueueCapacity *int `json:"queue_capacity"`
|
||||
ExtraParams json.RawMessage `json:"extra_params"`
|
||||
}
|
||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||
decoder.DisallowUnknownFields()
|
||||
var value manifest
|
||||
if err := decoder.Decode(&value); err != nil {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: invalid manifest", source.Name)
|
||||
}
|
||||
var trailing any
|
||||
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: invalid manifest", source.Name)
|
||||
}
|
||||
if value.SchemaVersion == nil || *value.SchemaVersion != 1 || value.ID == nil || value.Endpoint == nil || value.APIKeyEnv == nil || value.ConcurrencyLimit == nil || value.QueueCapacity == nil || value.ExtraParams == nil {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: required field is missing or unsupported", source.Name)
|
||||
}
|
||||
if *value.ID != source.ExpectedBackendID {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend ID does not match expected ID", source.Name)
|
||||
}
|
||||
if strings.TrimSpace(*value.APIKeyEnv) == "" {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: api key environment variable must not be blank", source.Name)
|
||||
}
|
||||
extraParams, err := decodeExtraParams(value.ExtraParams)
|
||||
if err != nil {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: invalid extra parameters", source.Name)
|
||||
}
|
||||
if err := rejectSecretKeys(extraParams); err != nil {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: prohibited extra parameter key", source.Name)
|
||||
}
|
||||
normalized, err := backend.NormalizeDefinition(domain.Backend{ID: *value.ID, Endpoint: *value.Endpoint, APIKeyEnv: *value.APIKeyEnv, ExtraParams: extraParams, ConcurrencyLimit: *value.ConcurrencyLimit, QueueCapacity: *value.QueueCapacity, QueueCapacitySet: true})
|
||||
if err != nil {
|
||||
return domain.Backend{}, fmt.Errorf("catalog %s: backend.json: invalid backend definition", source.Name)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func decodeExtraParams(data []byte) (map[string]any, error) {
|
||||
if bytes.Equal(bytes.TrimSpace(data), []byte("null")) {
|
||||
return nil, nil
|
||||
}
|
||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||
decoder.UseNumber()
|
||||
var value map[string]any
|
||||
if err := decoder.Decode(&value); err != nil || value == nil {
|
||||
return nil, errors.New("extra parameters must be an object or null")
|
||||
}
|
||||
var trailing any
|
||||
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
|
||||
return nil, errors.New("extra parameters must contain exactly one JSON value")
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func containsField(fields []string, target string) bool {
|
||||
for _, field := range fields {
|
||||
if field == target {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func rejectSecretKeys(value map[string]any) error {
|
||||
return rejectSecretValue(value)
|
||||
}
|
||||
|
||||
func rejectSecretValue(value any) error {
|
||||
switch value := value.(type) {
|
||||
case map[string]any:
|
||||
for key, child := range value {
|
||||
switch strings.ToLower(key) {
|
||||
case "api_key", "apikey", "authorization", "credential", "credentials", "password", "secret", "token", "access_token":
|
||||
return errors.New("prohibited key")
|
||||
}
|
||||
if err := rejectSecretValue(child); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
case []any:
|
||||
for _, child := range value {
|
||||
if err := rejectSecretValue(child); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
364
internal/catalog/catalog_test.go
Normal file
364
internal/catalog/catalog_test.go
Normal file
@@ -0,0 +1,364 @@
|
||||
package catalog
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
|
||||
openrouter "gitea.maximumdirect.net/eric/promptkit-backend-openrouter"
|
||||
rakestrawhome "gitea.maximumdirect.net/eric/promptkit-backend-rakestrawhome"
|
||||
)
|
||||
|
||||
func TestLoadPublishedCatalogsMatchCompatibilityFixture(t *testing.T) {
|
||||
loaded, err := Load(
|
||||
Source{Name: "OpenRouter", ExpectedBackendID: "openrouter", FS: openrouter.FS(), Root: openrouter.Root},
|
||||
Source{Name: "Rakestrawhome", ExpectedBackendID: "rakestrawhome", FS: rakestrawhome.FS(), Root: rakestrawhome.Root},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("load published catalogs: %v", err)
|
||||
}
|
||||
expected := loadCompatibilityFixture(t)
|
||||
expectedIDs := fixtureProfileIDs(t, expected)
|
||||
if !reflect.DeepEqual(loaded.profileIDs, expectedIDs) {
|
||||
t.Fatalf("published profile IDs differ from compatibility fixture: got %q, want %q", loaded.profileIDs, expectedIDs)
|
||||
}
|
||||
actual := catalogValue(t, loaded)
|
||||
actualJSON, err := json.Marshal(actual)
|
||||
if err != nil {
|
||||
t.Fatalf("encode loaded catalogs: %v", err)
|
||||
}
|
||||
expectedJSON, err := json.Marshal(expected)
|
||||
if err != nil {
|
||||
t.Fatalf("encode compatibility fixture: %v", err)
|
||||
}
|
||||
if !bytes.Equal(actualJSON, expectedJSON) {
|
||||
t.Fatalf("published catalogs differ from compatibility fixture: got %#v, want %#v", actual, expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidSources(t *testing.T) {
|
||||
for name, sources := range map[string][]Source{
|
||||
"none": nil,
|
||||
"blank name": {{Name: " ", ExpectedBackendID: "openrouter", FS: openrouter.FS(), Root: openrouter.Root}},
|
||||
"duplicate name": {
|
||||
{Name: "same", ExpectedBackendID: "openrouter", FS: openrouter.FS(), Root: openrouter.Root},
|
||||
{Name: "same", ExpectedBackendID: "rakestrawhome", FS: rakestrawhome.FS(), Root: rakestrawhome.Root},
|
||||
},
|
||||
"nil filesystem": {{Name: "missing", ExpectedBackendID: "openrouter", Root: openrouter.Root}},
|
||||
"invalid root": {{Name: "invalid-root", ExpectedBackendID: "openrouter", FS: openrouter.FS(), Root: "."}},
|
||||
"blank expected backend": {{Name: "blank-backend", ExpectedBackendID: " ", FS: openrouter.FS(), Root: openrouter.Root}},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if _, err := Load(sources...); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidLayouts(t *testing.T) {
|
||||
tests := map[string]func(fstest.MapFS){
|
||||
"unexpected file": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/notes.txt"] = &fstest.MapFile{Data: []byte("unexpected")}
|
||||
},
|
||||
"unexpected directory": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/unexpected"] = &fstest.MapFile{Mode: fs.ModeDir}
|
||||
},
|
||||
"nonregular manifest": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/backend.json"].Mode = fs.ModeSymlink
|
||||
},
|
||||
"wrong profile extension": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/extra.yaml"] = &fstest.MapFile{Data: []byte(validProfile("extra", "one"))}
|
||||
},
|
||||
"nonregular profile": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Mode = fs.ModeSymlink
|
||||
},
|
||||
}
|
||||
for name, mutate := range tests {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
fsys := validCatalogFS("one")
|
||||
mutate(fsys)
|
||||
if _, err := Load(testSource("one", fsys)); err == nil {
|
||||
t.Fatal("expected invalid layout error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidManifests(t *testing.T) {
|
||||
tests := map[string]string{
|
||||
"malformed": `{`,
|
||||
"trailing value": validManifest("one", "TEST_API_KEY", "null") + `{}`,
|
||||
"missing fields": `{"schema_version":1,"id":"one"}`,
|
||||
"unsupported version": strings.Replace(validManifest("one", "TEST_API_KEY", "null"), `"schema_version":1`, `"schema_version":2`, 1),
|
||||
"unknown field": strings.Replace(validManifest("one", "TEST_API_KEY", "null"), `"extra_params":null`, `"extra_params":null,"unknown":true`, 1),
|
||||
"blank API key env": validManifest("one", " ", "null"),
|
||||
"invalid API key env": validManifest("one", "LEAK-MARKER", "null"),
|
||||
"invalid endpoint": strings.Replace(validManifest("one", "TEST_API_KEY", "null"), `https://one.example/v1`, `ftp://leak-marker.invalid/v1`, 1),
|
||||
"zero concurrency": strings.Replace(validManifest("one", "TEST_API_KEY", "null"), `"concurrency_limit":2`, `"concurrency_limit":0`, 1),
|
||||
"non-object parameters": validManifest("one", "TEST_API_KEY", `[]`),
|
||||
"secret parameter": validManifest("one", "TEST_API_KEY", `{"nested":{"token":"leak-marker"}}`),
|
||||
}
|
||||
for name, manifest := range tests {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
fsys := validCatalogFS("one")
|
||||
fsys["catalog/backend.json"].Data = []byte(manifest)
|
||||
_, err := Load(testSource("one", fsys))
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid manifest error")
|
||||
}
|
||||
if strings.Contains(strings.ToLower(err.Error()), "leak-marker") {
|
||||
t.Fatalf("catalog error exposed manifest content: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadPreservesManifestJSONNumbers(t *testing.T) {
|
||||
fsys := validCatalogFS("one")
|
||||
fsys["catalog/backend.json"].Data = []byte(validManifest(
|
||||
"one",
|
||||
"TEST_API_KEY",
|
||||
`{"large":9007199254740993,"nested":[1.25]}`,
|
||||
))
|
||||
loaded, err := Load(testSource("one", fsys))
|
||||
if err != nil {
|
||||
t.Fatalf("load catalog: %v", err)
|
||||
}
|
||||
if got := loaded.Backends[0].ExtraParams["large"]; got != json.Number("9007199254740993") {
|
||||
t.Fatalf("large JSON integer = %#v, want preserved json.Number", got)
|
||||
}
|
||||
nested := loaded.Backends[0].ExtraParams["nested"].([]any)
|
||||
if nested[0] != json.Number("1.25") {
|
||||
t.Fatalf("nested JSON number = %#v, want preserved json.Number", nested[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidCatalogProfiles(t *testing.T) {
|
||||
tests := map[string]func(fstest.MapFS){
|
||||
"empty": func(fsys fstest.MapFS) {
|
||||
delete(fsys, "catalog/profiles/one-profile.yml")
|
||||
fsys["catalog/profiles"] = &fstest.MapFile{Mode: fs.ModeDir}
|
||||
},
|
||||
"malformed": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte("id: [")
|
||||
},
|
||||
"raw API key": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte(validProfile("one-profile", "one") + "api_key: leak-marker\n")
|
||||
},
|
||||
"endpoint field": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte(validProfile("one-profile", "one") + "endpoint: ''\n")
|
||||
},
|
||||
"API key environment field": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte(validProfile("one-profile", "one") + "api_key_env: ''\n")
|
||||
},
|
||||
"owner mismatch": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte(validProfile("one-profile", "other"))
|
||||
},
|
||||
"missing base": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte("id: one-profile\nbase_profile: leak-marker\n")
|
||||
},
|
||||
"cyclic base": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte("id: one-profile\nbase_profile: second\n")
|
||||
fsys["catalog/profiles/second.yml"] = &fstest.MapFile{Data: []byte("id: second\nbase_profile: one-profile\n")}
|
||||
},
|
||||
"secret profile parameter": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte(validProfile("one-profile", "one") + "extra_params:\n nested:\n password: leak-marker\n")
|
||||
},
|
||||
"unknown field is redacted": func(fsys fstest.MapFS) {
|
||||
fsys["catalog/profiles/one-profile.yml"].Data = []byte(validProfile("one-profile", "one") + "leak_marker: leak-marker\n")
|
||||
},
|
||||
}
|
||||
for name, mutate := range tests {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
fsys := validCatalogFS("one")
|
||||
mutate(fsys)
|
||||
_, err := Load(testSource("one", fsys))
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid profile error")
|
||||
}
|
||||
if strings.Contains(strings.ToLower(err.Error()), "leak-marker") {
|
||||
t.Fatalf("catalog error exposed profile content: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsCrossCatalogConflicts(t *testing.T) {
|
||||
t.Run("duplicate backend", func(t *testing.T) {
|
||||
second := catalogFS("one", map[string]string{
|
||||
"catalog/profiles/second.yml": validProfile("second", "one"),
|
||||
}, "null")
|
||||
_, err := Load(
|
||||
Source{Name: "first", ExpectedBackendID: "one", FS: validCatalogFS("one"), Root: "catalog"},
|
||||
Source{Name: "second", ExpectedBackendID: "one", FS: second, Root: "catalog"},
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatal("expected duplicate backend error")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("duplicate profile", func(t *testing.T) {
|
||||
first := catalogFS("one", map[string]string{
|
||||
"catalog/profiles/shared.yml": validProfile("shared", "one"),
|
||||
}, "null")
|
||||
second := catalogFS("two", map[string]string{
|
||||
"catalog/profiles/shared.yml": validProfile("shared", "two"),
|
||||
}, "null")
|
||||
_, err := Load(testSource("one", first), testSource("two", second))
|
||||
if err == nil {
|
||||
t.Fatal("expected duplicate profile error")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cross-catalog base", func(t *testing.T) {
|
||||
first := catalogFS("one", map[string]string{
|
||||
"catalog/profiles/base.yml": validProfile("base", "one"),
|
||||
}, "null")
|
||||
second := catalogFS("two", map[string]string{
|
||||
"catalog/profiles/child.yml": "id: child\nbase_profile: base\n",
|
||||
}, "null")
|
||||
_, err := Load(testSource("one", first), testSource("two", second))
|
||||
if err == nil {
|
||||
t.Fatal("expected cross-catalog base error")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestLoadReturnsDefensiveCatalogValues(t *testing.T) {
|
||||
fsys := catalogFS("one", map[string]string{
|
||||
"catalog/profiles/one-profile.yml": validProfile("one-profile", "one") + "extra_params:\n nested:\n value: profile\n",
|
||||
}, `{"nested":{"value":"backend"}}`)
|
||||
loaded, err := Load(testSource("one", fsys))
|
||||
if err != nil {
|
||||
t.Fatalf("load catalog: %v", err)
|
||||
}
|
||||
loaded.Backends[0].ExtraParams["nested"].(map[string]any)["value"] = "changed"
|
||||
profileValue, err := loaded.Profiles.GetProfile(context.Background(), "one-profile")
|
||||
if err != nil {
|
||||
t.Fatalf("load profile: %v", err)
|
||||
}
|
||||
profileValue.ExtraParams["nested"].(map[string]any)["value"] = "changed"
|
||||
|
||||
again, err := Load(testSource("one", fsys))
|
||||
if err != nil {
|
||||
t.Fatalf("reload catalog: %v", err)
|
||||
}
|
||||
if got := again.Backends[0].ExtraParams["nested"].(map[string]any)["value"]; got != "backend" {
|
||||
t.Fatalf("backend mutation escaped returned set: %#v", got)
|
||||
}
|
||||
againProfile, err := loaded.Profiles.GetProfile(context.Background(), "one-profile")
|
||||
if err != nil {
|
||||
t.Fatalf("reload profile: %v", err)
|
||||
}
|
||||
if got := againProfile.ExtraParams["nested"].(map[string]any)["value"]; got != "profile" {
|
||||
t.Fatalf("profile mutation escaped returned value: %#v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func loadCompatibilityFixture(t *testing.T) map[string]any {
|
||||
t.Helper()
|
||||
data, err := os.ReadFile(filepath.Join("..", "..", "testdata", "builtin-catalog-v1.json"))
|
||||
if err != nil {
|
||||
t.Fatalf("read fixture: %v", err)
|
||||
}
|
||||
var value map[string]any
|
||||
if err := json.Unmarshal(data, &value); err != nil {
|
||||
t.Fatalf("decode fixture: %v", err)
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func fixtureProfileIDs(t *testing.T, fixture map[string]any) []string {
|
||||
t.Helper()
|
||||
profiles, ok := fixture["profiles"].([]any)
|
||||
if !ok {
|
||||
t.Fatal("compatibility fixture profiles are malformed")
|
||||
}
|
||||
ids := make([]string, 0, len(profiles))
|
||||
for _, entry := range profiles {
|
||||
profileValue, ok := entry.(map[string]any)
|
||||
if !ok {
|
||||
t.Fatal("compatibility fixture profile is malformed")
|
||||
}
|
||||
id, ok := profileValue["id"].(string)
|
||||
if !ok {
|
||||
t.Fatal("compatibility fixture profile ID is malformed")
|
||||
}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Strings(ids)
|
||||
return ids
|
||||
}
|
||||
|
||||
func catalogValue(t *testing.T, loaded Set) map[string]any {
|
||||
t.Helper()
|
||||
backends := make([]any, 0, len(loaded.Backends))
|
||||
for _, backend := range loaded.Backends {
|
||||
backends = append(backends, map[string]any{"id": backend.ID, "endpoint": backend.Endpoint, "api_key_env": backend.APIKeyEnv, "extra_params": interfaceValue(backend.ExtraParams), "concurrency_limit": backend.ConcurrencyLimit, "queue_capacity": backend.QueueCapacity, "queue_capacity_set": backend.QueueCapacitySet})
|
||||
}
|
||||
sort.Slice(backends, func(left, right int) bool {
|
||||
return backends[left].(map[string]any)["id"].(string) < backends[right].(map[string]any)["id"].(string)
|
||||
})
|
||||
actualProfiles := make([]any, 0, len(loaded.profileIDs))
|
||||
for _, id := range loaded.profileIDs {
|
||||
profile, err := loaded.Profiles.GetProfile(context.Background(), id)
|
||||
if err != nil {
|
||||
t.Fatalf("load profile %q: %v", id, err)
|
||||
}
|
||||
actualProfiles = append(actualProfiles, map[string]any{"id": profile.ID, "base_profile": profile.BaseProfileID, "backend": profile.BackendID, "endpoint": profile.Endpoint, "model": profile.Model, "temperature": profile.Temperature, "max_tokens": profile.MaxTokens, "top_p": profile.TopP, "timeout_seconds": profile.TimeoutSeconds, "service_tier": profile.ServiceTier, "reasoning_effort": profile.ReasoningEffort, "api_key_env": profile.APIKeyEnv, "api_key_required": profile.APIKeyRequired, "extra_params": interfaceValue(profile.ExtraParams)})
|
||||
}
|
||||
return map[string]any{"backends": backends, "profiles": actualProfiles}
|
||||
}
|
||||
|
||||
func interfaceValue(value map[string]any) any {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func testSource(id string, fsys fs.FS) Source {
|
||||
return Source{Name: id, ExpectedBackendID: id, FS: fsys, Root: "catalog"}
|
||||
}
|
||||
|
||||
func validCatalogFS(id string) fstest.MapFS {
|
||||
return catalogFS(id, map[string]string{
|
||||
"catalog/profiles/" + id + "-profile.yml": validProfile(id+"-profile", id),
|
||||
}, "null")
|
||||
}
|
||||
|
||||
func catalogFS(id string, profiles map[string]string, extraParams string) fstest.MapFS {
|
||||
fsys := fstest.MapFS{
|
||||
"catalog/backend.json": &fstest.MapFile{Data: []byte(validManifest(id, "TEST_API_KEY", extraParams))},
|
||||
}
|
||||
for name, content := range profiles {
|
||||
fsys[name] = &fstest.MapFile{Data: []byte(content)}
|
||||
}
|
||||
return fsys
|
||||
}
|
||||
|
||||
func validManifest(id, apiKeyEnv, extraParams string) string {
|
||||
return fmt.Sprintf(
|
||||
`{"schema_version":1,"id":%q,"endpoint":%q,"api_key_env":%q,"concurrency_limit":2,"queue_capacity":3,"extra_params":%s}`,
|
||||
id,
|
||||
"https://"+id+".example/v1",
|
||||
apiKeyEnv,
|
||||
extraParams,
|
||||
)
|
||||
}
|
||||
|
||||
func validProfile(id, backendID string) string {
|
||||
return fmt.Sprintf("id: %s\nbackend: %s\nmodel: test-model\n", id, backendID)
|
||||
}
|
||||
|
||||
var _ fs.FS = openrouter.FS()
|
||||
@@ -60,15 +60,16 @@ type CacheControl struct {
|
||||
|
||||
// RunRequest represents a request to generate a single artifact.
|
||||
type RunRequest struct {
|
||||
PromptID string
|
||||
PromptVersion string
|
||||
ProfileID string
|
||||
SessionID string
|
||||
APIKey string `json:"-" yaml:"-"`
|
||||
Inputs map[string]ArtifactRef
|
||||
Vars map[string]string
|
||||
Execution *ExecutionTargetOverride
|
||||
Validation *OutputContract
|
||||
PromptID string
|
||||
PromptVersion string
|
||||
ProfileID string
|
||||
SessionID string
|
||||
APIKey string `json:"-" yaml:"-"`
|
||||
Inputs map[string]ArtifactRef
|
||||
Vars map[string]string
|
||||
Execution *ExecutionTargetOverride
|
||||
Validation *OutputContract
|
||||
AppendedMessages []RenderedMessage
|
||||
}
|
||||
|
||||
// RunResult represents the complete result of a prompt execution run.
|
||||
@@ -192,6 +193,7 @@ type BackendCapacityPolicy struct {
|
||||
// ExecutionProfile describes how and where to execute a model.
|
||||
type ExecutionProfile struct {
|
||||
ID string `yaml:"id"`
|
||||
BaseProfileID string `yaml:"base_profile"`
|
||||
BackendID string `yaml:"backend"`
|
||||
Endpoint string `yaml:"endpoint"`
|
||||
Model string `yaml:"model"`
|
||||
|
||||
82
internal/domain/message.go
Normal file
82
internal/domain/message.go
Normal file
@@ -0,0 +1,82 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
const (
|
||||
RoleDeveloper = "developer"
|
||||
RoleSystem = "system"
|
||||
RoleUser = "user"
|
||||
RoleAssistant = "assistant"
|
||||
)
|
||||
|
||||
// NormalizeMessageRole validates and canonicalizes a provider-bound chat role.
|
||||
func NormalizeMessageRole(role string) (string, error) {
|
||||
if !utf8.ValidString(role) {
|
||||
return "", errors.New("message role must be valid UTF-8")
|
||||
}
|
||||
|
||||
normalized := strings.ToLower(strings.TrimSpace(role))
|
||||
switch normalized {
|
||||
case RoleDeveloper, RoleSystem, RoleUser, RoleAssistant:
|
||||
return normalized, nil
|
||||
default:
|
||||
return "", errors.New("message role must be developer, system, user, or assistant")
|
||||
}
|
||||
}
|
||||
|
||||
// NormalizeCacheControl validates, canonicalizes, and copies cache metadata.
|
||||
func NormalizeCacheControl(control *CacheControl) (*CacheControl, error) {
|
||||
if control == nil {
|
||||
return nil, nil
|
||||
}
|
||||
if !utf8.ValidString(string(control.Type)) {
|
||||
return nil, errors.New("cache control type must be valid UTF-8")
|
||||
}
|
||||
if !utf8.ValidString(control.TTL) {
|
||||
return nil, errors.New("cache control ttl must be valid UTF-8")
|
||||
}
|
||||
|
||||
cacheType := strings.TrimSpace(string(control.Type))
|
||||
if cacheType == "" {
|
||||
return nil, errors.New("cache control type is required")
|
||||
}
|
||||
if CacheControlType(cacheType) != CacheControlEphemeral {
|
||||
return nil, errors.New("unsupported type")
|
||||
}
|
||||
|
||||
ttl := strings.TrimSpace(control.TTL)
|
||||
if ttl != "" && ttl != "1h" {
|
||||
return nil, errors.New("unsupported ttl")
|
||||
}
|
||||
|
||||
return &CacheControl{Type: CacheControlType(cacheType), TTL: ttl}, nil
|
||||
}
|
||||
|
||||
// CloneRenderedMessages returns a deep copy of rendered messages.
|
||||
func CloneRenderedMessages(messages []RenderedMessage) []RenderedMessage {
|
||||
cloned := make([]RenderedMessage, len(messages))
|
||||
copyRenderedMessages(cloned, messages)
|
||||
return cloned
|
||||
}
|
||||
|
||||
// ConcatRenderedMessages returns an independently owned concatenation of messages.
|
||||
func ConcatRenderedMessages(prefix, suffix []RenderedMessage) []RenderedMessage {
|
||||
messages := make([]RenderedMessage, len(prefix)+len(suffix))
|
||||
copyRenderedMessages(messages, prefix)
|
||||
copyRenderedMessages(messages[len(prefix):], suffix)
|
||||
return messages
|
||||
}
|
||||
|
||||
func copyRenderedMessages(destination, source []RenderedMessage) {
|
||||
for index, message := range source {
|
||||
destination[index] = message
|
||||
if message.CacheControl != nil {
|
||||
cacheControl := *message.CacheControl
|
||||
destination[index].CacheControl = &cacheControl
|
||||
}
|
||||
}
|
||||
}
|
||||
121
internal/domain/message_test.go
Normal file
121
internal/domain/message_test.go
Normal file
@@ -0,0 +1,121 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNormalizeMessageRole(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "developer", input: RoleDeveloper, want: RoleDeveloper},
|
||||
{name: "system", input: RoleSystem, want: RoleSystem},
|
||||
{name: "user", input: RoleUser, want: RoleUser},
|
||||
{name: "assistant", input: RoleAssistant, want: RoleAssistant},
|
||||
{name: "surrounding whitespace and mixed case", input: " \u2003UsEr\u2003 ", want: RoleUser},
|
||||
{name: "blank", input: " \t\n ", wantErr: true},
|
||||
{name: "tool", input: "tool", wantErr: true},
|
||||
{name: "function", input: "function", wantErr: true},
|
||||
{name: "custom", input: "custom-role", wantErr: true},
|
||||
{name: "invalid UTF-8", input: string([]byte{0xff}), wantErr: true},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
got, err := NormalizeMessageRole(test.input)
|
||||
if test.wantErr {
|
||||
if err == nil {
|
||||
t.Fatal("expected an error")
|
||||
}
|
||||
if strings.Contains(err.Error(), test.input) {
|
||||
t.Fatalf("error exposed the input: %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("NormalizeMessageRole() error = %v", err)
|
||||
}
|
||||
if got != test.want {
|
||||
t.Fatalf("NormalizeMessageRole() = %q, want %q", got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeCacheControl(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input *CacheControl
|
||||
want *CacheControl
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "nil", input: nil, want: nil},
|
||||
{
|
||||
name: "canonical values",
|
||||
input: &CacheControl{Type: CacheControlEphemeral, TTL: "1h"},
|
||||
want: &CacheControl{Type: CacheControlEphemeral, TTL: "1h"},
|
||||
},
|
||||
{
|
||||
name: "trims values",
|
||||
input: &CacheControl{Type: " ephemeral ", TTL: " 1h\t"},
|
||||
want: &CacheControl{Type: CacheControlEphemeral, TTL: "1h"},
|
||||
},
|
||||
{name: "empty type", input: &CacheControl{}, wantErr: true},
|
||||
{name: "unsupported type", input: &CacheControl{Type: "persistent"}, wantErr: true},
|
||||
{name: "unsupported ttl", input: &CacheControl{Type: CacheControlEphemeral, TTL: "5m"}, wantErr: true},
|
||||
{name: "invalid type UTF-8", input: &CacheControl{Type: CacheControlType(string([]byte{0xff}))}, wantErr: true},
|
||||
{name: "invalid ttl UTF-8", input: &CacheControl{Type: CacheControlEphemeral, TTL: string([]byte{0xff})}, wantErr: true},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
got, err := NormalizeCacheControl(test.input)
|
||||
if test.wantErr {
|
||||
if err == nil {
|
||||
t.Fatal("expected an error")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("NormalizeCacheControl() error = %v", err)
|
||||
}
|
||||
if got == nil || test.want == nil {
|
||||
if got != test.want {
|
||||
t.Fatalf("NormalizeCacheControl() = %#v, want %#v", got, test.want)
|
||||
}
|
||||
return
|
||||
}
|
||||
if *got != *test.want {
|
||||
t.Fatalf("NormalizeCacheControl() = %#v, want %#v", got, test.want)
|
||||
}
|
||||
if got == test.input {
|
||||
t.Fatal("normalized cache control aliases its input")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderedMessageCloning(t *testing.T) {
|
||||
prefix := []RenderedMessage{{Role: RoleSystem, Content: "prefix", CacheControl: &CacheControl{Type: CacheControlEphemeral, TTL: "1h"}}}
|
||||
suffix := []RenderedMessage{{Role: RoleUser, Content: " suffix "}, {Role: RoleAssistant, Content: "", CacheControl: &CacheControl{Type: CacheControlEphemeral}}}
|
||||
|
||||
cloned := CloneRenderedMessages(prefix)
|
||||
combined := ConcatRenderedMessages(prefix, suffix)
|
||||
if len(combined) != 3 || combined[0].Content != "prefix" || combined[1].Content != " suffix " || combined[2].Content != "" {
|
||||
t.Fatalf("unexpected combined messages: %#v", combined)
|
||||
}
|
||||
if cloned[0].CacheControl == prefix[0].CacheControl || combined[0].CacheControl == prefix[0].CacheControl || combined[2].CacheControl == suffix[1].CacheControl {
|
||||
t.Fatal("cloned cache controls alias their inputs")
|
||||
}
|
||||
|
||||
prefix[0].Content = "changed"
|
||||
prefix[0].CacheControl.TTL = ""
|
||||
suffix[1].CacheControl.Type = "changed"
|
||||
if cloned[0].Content != "prefix" || cloned[0].CacheControl.TTL != "1h" || combined[0].Content != "prefix" || combined[0].CacheControl.TTL != "1h" || combined[2].CacheControl.Type != CacheControlEphemeral {
|
||||
t.Fatalf("cloned messages changed with their inputs: cloned=%#v combined=%#v", cloned, combined)
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
const maxOutputRepairAttempts = 3
|
||||
|
||||
// ValidateOutputContract validates source-neutral output-contract invariants.
|
||||
func ValidateOutputContract(contract OutputContract) error {
|
||||
switch contract.Format {
|
||||
@@ -26,5 +28,11 @@ func ValidateOutputContract(contract OutputContract) error {
|
||||
if contract.RepairAttempts < 0 {
|
||||
return errors.New("repair_attempts must be greater than or equal to 0")
|
||||
}
|
||||
if contract.RepairAttempts > maxOutputRepairAttempts {
|
||||
return fmt.Errorf("repair_attempts must be less than or equal to %d", maxOutputRepairAttempts)
|
||||
}
|
||||
if contract.ValidationMode == ValidationNone && contract.RepairAttempts > 0 {
|
||||
return errors.New("repair_attempts requires basic, json, or json_schema validation")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -30,9 +30,28 @@ func TestValidateOutputContract(t *testing.T) {
|
||||
}},
|
||||
{name: "empty validation mode", change: func(c *OutputContract) { c.ValidationMode = "" }, wantErr: "validation mode"},
|
||||
{name: "unsupported validation mode", change: func(c *OutputContract) { c.ValidationMode = ValidationMode("unknown") }, wantErr: "validation mode"},
|
||||
{name: "negative repair attempts", change: func(c *OutputContract) { c.RepairAttempts = -1 }, wantErr: "repair_attempts"},
|
||||
{name: "negative repair attempts", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationBasic
|
||||
c.RepairAttempts = -1
|
||||
}, wantErr: "repair_attempts"},
|
||||
{name: "zero repair attempts", change: func(c *OutputContract) { c.RepairAttempts = 0 }},
|
||||
{name: "positive repair attempts", change: func(c *OutputContract) { c.RepairAttempts = 1 }},
|
||||
{name: "one repair attempt", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationBasic
|
||||
c.RepairAttempts = 1
|
||||
}},
|
||||
{name: "maximum repair attempts", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationJSON
|
||||
c.RepairAttempts = 3
|
||||
}},
|
||||
{name: "too many repair attempts", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationJSONSchema
|
||||
c.SchemaPath = "schema.json"
|
||||
c.RepairAttempts = 4
|
||||
}, wantErr: "repair_attempts"},
|
||||
{name: "none validation with repair attempts", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationNone
|
||||
c.RepairAttempts = 1
|
||||
}, wantErr: "repair_attempts"},
|
||||
{name: "json schema empty path", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationJSONSchema
|
||||
c.SchemaPath = ""
|
||||
|
||||
@@ -133,13 +133,18 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
||||
return nil, fmt.Errorf("%w: failed to create request: %v", ErrRequestFailed, err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
if apiKey := strings.TrimSpace(req.Target.APIKey); apiKey != "" {
|
||||
httpReq.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
} else if envName := strings.TrimSpace(req.Target.APIKeyEnv); envName != "" {
|
||||
apiKey := strings.TrimSpace(os.Getenv(envName))
|
||||
if apiKey == "" {
|
||||
apiKey := strings.TrimSpace(req.Target.APIKey)
|
||||
envName := strings.TrimSpace(req.Target.APIKeyEnv)
|
||||
if apiKey == "" && envName != "" {
|
||||
apiKey = strings.TrimSpace(os.Getenv(envName))
|
||||
}
|
||||
if apiKey == "" && req.Target.APIKeyRequired {
|
||||
if envName != "" {
|
||||
return nil, fmt.Errorf("%w: api key environment variable %q is not set", ErrInvalidRequest, envName)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: api key is required", ErrInvalidRequest)
|
||||
}
|
||||
if apiKey != "" {
|
||||
httpReq.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
}
|
||||
|
||||
@@ -155,8 +160,11 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
||||
defer httpResp.Body.Close()
|
||||
|
||||
if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 {
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(httpResp.Body, 4096))
|
||||
return nil, fmt.Errorf("%w: status=%d", ErrUnexpectedStatus, httpResp.StatusCode)
|
||||
return nil, providerHTTPErrorFromBody(
|
||||
httpResp.StatusCode,
|
||||
httpResp.ContentLength,
|
||||
httpResp.Body,
|
||||
)
|
||||
}
|
||||
if httpResp.ContentLength > maxOpenAIChatResponseBytes {
|
||||
return nil, openAIChatResponseTooLargeError()
|
||||
@@ -171,12 +179,12 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
||||
return nil, fmt.Errorf("%w: no choices returned", ErrMalformedResponse)
|
||||
}
|
||||
content := wireResp.Choices[0].Message.Content
|
||||
if content == "" {
|
||||
return nil, fmt.Errorf("%w: first choice has empty message content", ErrMalformedResponse)
|
||||
if content == nil {
|
||||
return nil, fmt.Errorf("%w: first choice has missing message content", ErrMalformedResponse)
|
||||
}
|
||||
|
||||
return &domain.GenerateResponse{
|
||||
Content: content,
|
||||
Content: *content,
|
||||
Usage: domain.TokenUsage{
|
||||
PromptTokens: wireResp.Usage.PromptTokens,
|
||||
CompletionTokens: wireResp.Usage.CompletionTokens,
|
||||
@@ -371,8 +379,8 @@ type openAICacheControl struct {
|
||||
}
|
||||
|
||||
type openAIChatResponseMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
Role string `json:"role"`
|
||||
Content *string `json:"content"`
|
||||
}
|
||||
|
||||
type openAIChatResponse struct {
|
||||
|
||||
@@ -363,6 +363,7 @@ func TestOpenAICompatibleClientRequestMapping(t *testing.T) {
|
||||
run func(*testing.T)
|
||||
}{
|
||||
{name: "complete request and response mapping", run: checkCompleteRequestAndResponseMapping},
|
||||
{name: "canonical role and exact content", run: checkCanonicalRoleAndExactContentMapping},
|
||||
{name: "cache-controlled message", run: checkCacheControlledMessageMapping},
|
||||
{name: "empty cache-control TTL", run: checkEmptyCacheControlTTLOmission},
|
||||
{name: "session ID", run: checkSessionIDMapping},
|
||||
@@ -377,6 +378,34 @@ func TestOpenAICompatibleClientRequestMapping(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func checkCanonicalRoleAndExactContentMapping(t *testing.T) {
|
||||
provider := newRecordingProvider(t)
|
||||
client := newProviderClient(t, provider, OpenAICompatibleConfig{})
|
||||
const content = " exact content\nwith whitespace "
|
||||
_, err := client.Generate(context.Background(), domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{
|
||||
Role: domain.RoleDeveloper,
|
||||
Content: content,
|
||||
}}},
|
||||
Target: domain.ExecutionTarget{Model: "model"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Generate() error = %v", err)
|
||||
}
|
||||
|
||||
var messages []struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
request := provider.lastRequest(t)
|
||||
if !request.decodeField(t, "messages", &messages) || len(messages) != 1 {
|
||||
t.Fatalf("messages = %#v, want one message", messages)
|
||||
}
|
||||
if messages[0].Role != domain.RoleDeveloper || messages[0].Content != content {
|
||||
t.Fatalf("message = %#v, want role %q and exact content %q", messages[0], domain.RoleDeveloper, content)
|
||||
}
|
||||
}
|
||||
|
||||
func checkCompleteRequestAndResponseMapping(t *testing.T) {
|
||||
provider := newRecordingProvider(t)
|
||||
provider.respond(http.StatusOK, `{
|
||||
@@ -501,36 +530,61 @@ func checkCompleteRequestAndResponseMapping(t *testing.T) {
|
||||
|
||||
func TestOpenAICompatibleClientAuthentication(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
configureEnv func(*testing.T)
|
||||
target domain.ExecutionTarget
|
||||
wantAuth string
|
||||
wantErr error
|
||||
wantCallCount int
|
||||
name string
|
||||
configureEnv func(*testing.T)
|
||||
target domain.ExecutionTarget
|
||||
wantAuthorization string
|
||||
wantError error
|
||||
wantCallCount int
|
||||
}{
|
||||
{
|
||||
name: "direct key takes precedence over environment",
|
||||
configureEnv: func(t *testing.T) {
|
||||
t.Setenv("PROMPTKIT_TEST_API_KEY", "env-key")
|
||||
t.Setenv("PROMPTKIT_TEST_API_KEY", " env-key ")
|
||||
},
|
||||
target: domain.ExecutionTarget{
|
||||
APIKeyEnv: "PROMPTKIT_TEST_API_KEY",
|
||||
APIKey: "direct-llm-key",
|
||||
APIKey: " direct-llm-key ",
|
||||
},
|
||||
wantAuth: "Bearer direct-llm-key",
|
||||
wantCallCount: 1,
|
||||
wantAuthorization: "Bearer direct-llm-key",
|
||||
wantCallCount: 1,
|
||||
},
|
||||
{
|
||||
name: "environment key supplies authorization",
|
||||
configureEnv: func(t *testing.T) {
|
||||
t.Setenv("PROMPTKIT_TEST_API_KEY", " env-key ")
|
||||
},
|
||||
target: domain.ExecutionTarget{APIKeyEnv: "PROMPTKIT_TEST_API_KEY"},
|
||||
wantAuthorization: "Bearer env-key",
|
||||
wantCallCount: 1,
|
||||
},
|
||||
{
|
||||
name: "no key omits authorization",
|
||||
wantCallCount: 1,
|
||||
},
|
||||
{
|
||||
name: "missing environment key fails before transport",
|
||||
name: "optional missing environment omits authorization",
|
||||
configureEnv: func(t *testing.T) {
|
||||
t.Setenv("PROMPTKIT_MISSING_KEY", "")
|
||||
},
|
||||
target: domain.ExecutionTarget{APIKeyEnv: "PROMPTKIT_MISSING_KEY"},
|
||||
wantErr: ErrInvalidRequest,
|
||||
target: domain.ExecutionTarget{APIKeyEnv: "PROMPTKIT_MISSING_KEY"},
|
||||
wantCallCount: 1,
|
||||
},
|
||||
{
|
||||
name: "required missing environment fails before transport",
|
||||
configureEnv: func(t *testing.T) {
|
||||
t.Setenv("PROMPTKIT_MISSING_KEY", "")
|
||||
},
|
||||
target: domain.ExecutionTarget{
|
||||
APIKeyEnv: "PROMPTKIT_MISSING_KEY",
|
||||
APIKeyRequired: true,
|
||||
},
|
||||
wantError: ErrInvalidRequest,
|
||||
},
|
||||
{
|
||||
name: "required target without source fails before transport",
|
||||
target: domain.ExecutionTarget{APIKeyRequired: true},
|
||||
wantError: ErrInvalidRequest,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -545,9 +599,9 @@ func TestOpenAICompatibleClientAuthentication(t *testing.T) {
|
||||
request.Target = tc.target
|
||||
|
||||
_, err := client.Generate(context.Background(), request)
|
||||
if tc.wantErr != nil {
|
||||
if !errors.Is(err, tc.wantErr) {
|
||||
t.Fatalf("error = %v, want %v", err, tc.wantErr)
|
||||
if tc.wantError != nil {
|
||||
if !errors.Is(err, tc.wantError) {
|
||||
t.Fatalf("error = %v, want %v", err, tc.wantError)
|
||||
}
|
||||
} else if err != nil {
|
||||
t.Fatalf("generate: %v", err)
|
||||
@@ -556,8 +610,13 @@ func TestOpenAICompatibleClientAuthentication(t *testing.T) {
|
||||
t.Fatalf("provider calls = %d, want %d", got, tc.wantCallCount)
|
||||
}
|
||||
if tc.wantCallCount == 1 {
|
||||
if got := provider.lastRequest(t).header.Get("Authorization"); got != tc.wantAuth {
|
||||
t.Fatalf("Authorization = %q, want %q", got, tc.wantAuth)
|
||||
values := provider.lastRequest(t).header.Values("Authorization")
|
||||
if tc.wantAuthorization == "" {
|
||||
if len(values) != 0 {
|
||||
t.Fatalf("Authorization values = %q, want absent", values)
|
||||
}
|
||||
} else if len(values) != 1 || values[0] != tc.wantAuthorization {
|
||||
t.Fatalf("Authorization values = %q, want [%q]", values, tc.wantAuthorization)
|
||||
}
|
||||
}
|
||||
})
|
||||
@@ -735,6 +794,7 @@ func TestOpenAICompatibleClientResponseFraming(t *testing.T) {
|
||||
run func(*testing.T)
|
||||
}{
|
||||
{name: "usage mapping", run: checkCacheUsageMapping},
|
||||
{name: "content presence", run: checkContentPresence},
|
||||
{name: "common response failures", run: checkCommonResponseFailures},
|
||||
{name: "successful response byte boundary", run: checkSuccessfulResponseByteBoundary},
|
||||
{name: "continuing oversized response", run: checkContinuingOversizedResponse},
|
||||
@@ -746,6 +806,53 @@ func TestOpenAICompatibleClientResponseFraming(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func checkContentPresence(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
content string
|
||||
}{
|
||||
{name: "explicit empty string", content: ""},
|
||||
{name: "whitespace string", content: " \n\t "},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
provider := newRecordingProvider(t)
|
||||
provider.respond(http.StatusOK, `{
|
||||
"choices": [{"message": {"content": `+strconv.Quote(tc.content)+`}}],
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"total_tokens": 30,
|
||||
"prompt_tokens_details": {"cached_tokens": 4},
|
||||
"cache_write_tokens": 5
|
||||
}
|
||||
}`)
|
||||
client := newProviderClient(t, provider, OpenAICompatibleConfig{Model: "model"})
|
||||
|
||||
response, err := client.Generate(context.Background(), ordinaryGenerateRequest())
|
||||
if err != nil {
|
||||
t.Fatalf("generate: %v", err)
|
||||
}
|
||||
if response == nil {
|
||||
t.Fatal("expected response")
|
||||
}
|
||||
if response.Content != tc.content {
|
||||
t.Fatalf("content = %q, want %q", response.Content, tc.content)
|
||||
}
|
||||
if response.Usage != (domain.TokenUsage{
|
||||
PromptTokens: 10,
|
||||
CompletionTokens: 20,
|
||||
TotalTokens: 30,
|
||||
CachedTokens: 4,
|
||||
CacheWriteTokens: 5,
|
||||
}) {
|
||||
t.Fatalf("usage = %+v", response.Usage)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func checkCacheUsageMapping(t *testing.T) {
|
||||
provider := newRecordingProvider(t)
|
||||
provider.respond(http.StatusOK, `{
|
||||
@@ -1099,6 +1206,9 @@ func checkCommonResponseFailures(t *testing.T) {
|
||||
},
|
||||
{name: "invalid JSON", statusCode: http.StatusOK, body: `{not valid json`, wantErr: ErrMalformedResponse},
|
||||
{name: "missing choices", statusCode: http.StatusOK, body: `{"choices": []}`, wantErr: ErrMalformedResponse},
|
||||
{name: "missing content", statusCode: http.StatusOK, body: `{"choices": [{"message": {}}]}`, wantErr: ErrMalformedResponse},
|
||||
{name: "null content", statusCode: http.StatusOK, body: `{"choices": [{"message": {"content": null}}]}`, wantErr: ErrMalformedResponse},
|
||||
{name: "non-string content", statusCode: http.StatusOK, body: `{"choices": [{"message": {"content": 1}}]}`, wantErr: ErrMalformedResponse},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
@@ -1114,6 +1224,15 @@ func checkCommonResponseFailures(t *testing.T) {
|
||||
if !errors.Is(err, tc.wantErr) {
|
||||
t.Fatalf("error = %v, want %v", err, tc.wantErr)
|
||||
}
|
||||
if tc.statusCode < http.StatusOK || tc.statusCode >= http.StatusMultipleChoices {
|
||||
var providerHTTPError *ProviderHTTPError
|
||||
if !errors.As(err, &providerHTTPError) {
|
||||
t.Fatalf("error = %T, want *ProviderHTTPError", err)
|
||||
}
|
||||
if got := providerHTTPError.StatusCode(); got != tc.statusCode {
|
||||
t.Fatalf("provider status = %d, want %d", got, tc.statusCode)
|
||||
}
|
||||
}
|
||||
if tc.wantText != "" && !strings.Contains(err.Error(), tc.wantText) {
|
||||
t.Fatalf("error %q does not contain %q", err, tc.wantText)
|
||||
}
|
||||
|
||||
189
internal/llm/provider_http_error.go
Normal file
189
internal/llm/provider_http_error.go
Normal file
@@ -0,0 +1,189 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
const (
|
||||
maxProviderErrorResponseBytes int64 = 64 << 10
|
||||
maxProviderErrorIdentifierRunes = 256
|
||||
maxProviderErrorMessageRunes = 4096
|
||||
)
|
||||
|
||||
// ProviderHTTPError describes a non-success response from an LLM provider.
|
||||
type ProviderHTTPError struct {
|
||||
statusCode int
|
||||
providerCode string
|
||||
providerType string
|
||||
providerMessage string
|
||||
}
|
||||
|
||||
func (e *ProviderHTTPError) StatusCode() int {
|
||||
if e == nil {
|
||||
return 0
|
||||
}
|
||||
return e.statusCode
|
||||
}
|
||||
|
||||
func (e *ProviderHTTPError) ProviderCode() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
return e.providerCode
|
||||
}
|
||||
|
||||
func (e *ProviderHTTPError) ProviderType() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
return e.providerType
|
||||
}
|
||||
|
||||
func (e *ProviderHTTPError) ProviderMessage() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
return e.providerMessage
|
||||
}
|
||||
|
||||
func (e *ProviderHTTPError) Error() string {
|
||||
if e == nil || e.statusCode == 0 {
|
||||
return ErrUnexpectedStatus.Error()
|
||||
}
|
||||
return fmt.Sprintf("%s: status=%d", ErrUnexpectedStatus, e.statusCode)
|
||||
}
|
||||
|
||||
func (e *ProviderHTTPError) GoString() string {
|
||||
return e.Error()
|
||||
}
|
||||
|
||||
func (e *ProviderHTTPError) Unwrap() error {
|
||||
return ErrUnexpectedStatus
|
||||
}
|
||||
|
||||
type providerErrorDetails struct {
|
||||
providerCode string
|
||||
providerType string
|
||||
providerMessage string
|
||||
}
|
||||
|
||||
func newProviderHTTPError(statusCode int, details providerErrorDetails) *ProviderHTTPError {
|
||||
return &ProviderHTTPError{
|
||||
statusCode: statusCode,
|
||||
providerCode: details.providerCode,
|
||||
providerType: details.providerType,
|
||||
providerMessage: details.providerMessage,
|
||||
}
|
||||
}
|
||||
|
||||
func providerHTTPErrorFromBody(statusCode int, contentLength int64, body io.Reader) *ProviderHTTPError {
|
||||
if contentLength > maxProviderErrorResponseBytes {
|
||||
return newProviderHTTPError(statusCode, providerErrorDetails{})
|
||||
}
|
||||
|
||||
limited := &io.LimitedReader{
|
||||
R: body,
|
||||
N: maxProviderErrorResponseBytes + 1,
|
||||
}
|
||||
contents, err := io.ReadAll(limited)
|
||||
if err != nil || limited.N == 0 {
|
||||
return newProviderHTTPError(statusCode, providerErrorDetails{})
|
||||
}
|
||||
return newProviderHTTPError(statusCode, parseProviderErrorEnvelope(contents))
|
||||
}
|
||||
|
||||
func parseProviderErrorEnvelope(body []byte) providerErrorDetails {
|
||||
decoder := json.NewDecoder(strings.NewReader(string(body)))
|
||||
decoder.UseNumber()
|
||||
|
||||
var envelope map[string]json.RawMessage
|
||||
if err := decoder.Decode(&envelope); err != nil {
|
||||
return providerErrorDetails{}
|
||||
}
|
||||
var trailing any
|
||||
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
|
||||
return providerErrorDetails{}
|
||||
}
|
||||
|
||||
rawError, ok := envelope["error"]
|
||||
if !ok {
|
||||
return providerErrorDetails{}
|
||||
}
|
||||
var providerError map[string]json.RawMessage
|
||||
if err := json.Unmarshal(rawError, &providerError); err != nil || providerError == nil {
|
||||
return providerErrorDetails{}
|
||||
}
|
||||
|
||||
var details providerErrorDetails
|
||||
if raw, ok := providerError["message"]; ok {
|
||||
var value string
|
||||
if json.Unmarshal(raw, &value) == nil {
|
||||
details.providerMessage = normalizeProviderErrorMessage(value)
|
||||
}
|
||||
}
|
||||
if raw, ok := providerError["type"]; ok {
|
||||
var value string
|
||||
if json.Unmarshal(raw, &value) == nil {
|
||||
details.providerType = normalizeProviderErrorIdentifier(value)
|
||||
}
|
||||
}
|
||||
if raw, ok := providerError["code"]; ok {
|
||||
var value any
|
||||
fieldDecoder := json.NewDecoder(strings.NewReader(string(raw)))
|
||||
fieldDecoder.UseNumber()
|
||||
if fieldDecoder.Decode(&value) == nil {
|
||||
switch value := value.(type) {
|
||||
case string:
|
||||
details.providerCode = normalizeProviderErrorIdentifier(value)
|
||||
case json.Number:
|
||||
details.providerCode = normalizeProviderErrorIdentifier(value.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return details
|
||||
}
|
||||
|
||||
func normalizeProviderErrorIdentifier(value string) string {
|
||||
normalized := normalizeProviderErrorText(value)
|
||||
if len([]rune(normalized)) > maxProviderErrorIdentifierRunes {
|
||||
return ""
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func normalizeProviderErrorMessage(value string) string {
|
||||
normalized := normalizeProviderErrorText(value)
|
||||
runes := []rune(normalized)
|
||||
if len(runes) <= maxProviderErrorMessageRunes {
|
||||
return normalized
|
||||
}
|
||||
return string(runes[:maxProviderErrorMessageRunes-1]) + "…"
|
||||
}
|
||||
|
||||
func normalizeProviderErrorText(value string) string {
|
||||
value = strings.ToValidUTF8(value, "<22>")
|
||||
|
||||
var result strings.Builder
|
||||
result.Grow(len(value))
|
||||
separatorPending := false
|
||||
for _, r := range value {
|
||||
if unicode.IsSpace(r) || unicode.IsControl(r) || unicode.In(r, unicode.Cf) {
|
||||
if result.Len() > 0 {
|
||||
separatorPending = true
|
||||
}
|
||||
continue
|
||||
}
|
||||
if separatorPending {
|
||||
result.WriteByte(' ')
|
||||
separatorPending = false
|
||||
}
|
||||
result.WriteRune(r)
|
||||
}
|
||||
return result.String()
|
||||
}
|
||||
255
internal/llm/provider_http_error_test.go
Normal file
255
internal/llm/provider_http_error_test.go
Normal file
@@ -0,0 +1,255 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
type guardedReader struct {
|
||||
reader io.Reader
|
||||
remaining int64
|
||||
bytes int64
|
||||
violated bool
|
||||
}
|
||||
|
||||
func (r *guardedReader) Read(buffer []byte) (int, error) {
|
||||
if int64(len(buffer)) > r.remaining {
|
||||
r.violated = true
|
||||
return 0, errors.New("reader was read past its allowed boundary")
|
||||
}
|
||||
n, err := r.reader.Read(buffer)
|
||||
r.bytes += int64(n)
|
||||
r.remaining -= int64(n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
type failingReader struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (r failingReader) Read([]byte) (int, error) {
|
||||
return 0, r.err
|
||||
}
|
||||
|
||||
func TestProviderHTTPErrorEnvelopeParsing(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
want providerErrorDetails
|
||||
}{
|
||||
{
|
||||
name: "all supported string fields",
|
||||
body: `{"error":{"message":"diagnostic","type":"invalid_request_error","code":"unsupported_parameter"}}`,
|
||||
want: providerErrorDetails{providerMessage: "diagnostic", providerType: "invalid_request_error", providerCode: "unsupported_parameter"},
|
||||
},
|
||||
{
|
||||
name: "integer code",
|
||||
body: `{"error":{"code":17}}`,
|
||||
want: providerErrorDetails{providerCode: "17"},
|
||||
},
|
||||
{
|
||||
name: "fractional code",
|
||||
body: `{"error":{"code":1.25}}`,
|
||||
want: providerErrorDetails{providerCode: "1.25"},
|
||||
},
|
||||
{
|
||||
name: "exponent code",
|
||||
body: `{"error":{"code":6.02e+23}}`,
|
||||
want: providerErrorDetails{providerCode: "6.02e+23"},
|
||||
},
|
||||
{
|
||||
name: "invalid fields do not discard valid fields",
|
||||
body: `{"error":{"message":null,"type":"invalid_request_error","code":false}}`,
|
||||
want: providerErrorDetails{providerType: "invalid_request_error"},
|
||||
},
|
||||
{
|
||||
name: "unknown fields are ignored",
|
||||
body: `{"trace":"do not retain","error":{"param":"temperature","metadata":{"secret":"x"}}}`,
|
||||
want: providerErrorDetails{},
|
||||
},
|
||||
{name: "missing error", body: `{}`, want: providerErrorDetails{}},
|
||||
{name: "null error", body: `{"error":null}`, want: providerErrorDetails{}},
|
||||
{name: "scalar error", body: `{"error":"nope"}`, want: providerErrorDetails{}},
|
||||
{name: "empty error", body: `{"error":{}}`, want: providerErrorDetails{}},
|
||||
{name: "malformed", body: `{"error":`, want: providerErrorDetails{}},
|
||||
{name: "truncated", body: `{"error":{"message":"x"`, want: providerErrorDetails{}},
|
||||
{name: "trailing garbage", body: `{"error":{"message":"x"}} garbage`, want: providerErrorDetails{}},
|
||||
{name: "second document", body: `{"error":{"message":"x"}} {}`, want: providerErrorDetails{}},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := parseProviderErrorEnvelope([]byte(tc.body)); !reflect.DeepEqual(got, tc.want) {
|
||||
t.Fatalf("parseProviderErrorEnvelope() = %#v, want %#v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderErrorTextNormalizationAndLimits(t *testing.T) {
|
||||
validIdentifier := strings.Repeat("界", maxProviderErrorIdentifierRunes)
|
||||
validMessage := strings.Repeat("界", maxProviderErrorMessageRunes)
|
||||
tests := []struct {
|
||||
name string
|
||||
got string
|
||||
want string
|
||||
}{
|
||||
{name: "multibyte text", got: "Grüße 世界", want: "Grüße 世界"},
|
||||
{name: "invalid UTF-8", got: string([]byte{'a', 0xff, 'b'}), want: "a<>b"},
|
||||
{name: "whitespace control and format runs", got: " \n\talpha\x00\u200b\u200bbeta \r ", want: "alpha beta"},
|
||||
{name: "blank normalization", got: "\t\u200b\n", want: ""},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := normalizeProviderErrorText(tc.got); got != tc.want {
|
||||
t.Fatalf("normalizeProviderErrorText() = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
if got := normalizeProviderErrorIdentifier(validIdentifier); got != validIdentifier {
|
||||
t.Fatalf("exact identifier boundary = %q, want retained value", got)
|
||||
}
|
||||
if got := normalizeProviderErrorIdentifier(validIdentifier + "界"); got != "" {
|
||||
t.Fatalf("overlong identifier = %q, want empty", got)
|
||||
}
|
||||
if got := normalizeProviderErrorMessage(validMessage); got != validMessage {
|
||||
t.Fatalf("exact message boundary = %q, want retained value", got)
|
||||
}
|
||||
wantTruncatedMessage := strings.Repeat("界", maxProviderErrorMessageRunes-1) + "…"
|
||||
if got := normalizeProviderErrorMessage(validMessage + "界"); got != wantTruncatedMessage {
|
||||
t.Fatalf("overlong message length = %d, want %d", utf8.RuneCountInString(got), maxProviderErrorMessageRunes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderHTTPErrorIdentityAndFormatting(t *testing.T) {
|
||||
const marker = "provider-secret-marker"
|
||||
err := newProviderHTTPError(429, providerErrorDetails{
|
||||
providerCode: marker + "-code",
|
||||
providerType: marker + "-type",
|
||||
providerMessage: marker + "-message",
|
||||
})
|
||||
|
||||
if err.StatusCode() != 429 || err.ProviderCode() != marker+"-code" || err.ProviderType() != marker+"-type" || err.ProviderMessage() != marker+"-message" {
|
||||
t.Fatalf("accessors returned unexpected values: %#v", err)
|
||||
}
|
||||
if !errors.Is(err, ErrUnexpectedStatus) {
|
||||
t.Fatalf("errors.Is(%v, ErrUnexpectedStatus) = false", err)
|
||||
}
|
||||
for _, rendered := range []string{fmt.Sprintf("%v", err), fmt.Sprintf("%+v", err), fmt.Sprintf("%#v", err)} {
|
||||
if rendered != "llm returned non-success status: status=429" {
|
||||
t.Fatalf("formatted error = %q", rendered)
|
||||
}
|
||||
if strings.Contains(rendered, marker) {
|
||||
t.Fatalf("formatted error exposed provider marker: %q", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
var nilError *ProviderHTTPError
|
||||
if nilError.StatusCode() != 0 || nilError.ProviderCode() != "" || nilError.ProviderType() != "" || nilError.ProviderMessage() != "" {
|
||||
t.Fatal("nil accessors returned provider values")
|
||||
}
|
||||
if nilError.Error() != "llm returned non-success status" || nilError.GoString() != "llm returned non-success status" || !errors.Is(nilError, ErrUnexpectedStatus) {
|
||||
t.Fatalf("nil error behavior is not safe: %v", nilError)
|
||||
}
|
||||
|
||||
zero := &ProviderHTTPError{}
|
||||
if zero.Error() != "llm returned non-success status" || zero.GoString() != "llm returned non-success status" || !errors.Is(zero, ErrUnexpectedStatus) {
|
||||
t.Fatalf("zero error behavior is not safe: %v", zero)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderHTTPErrorBodyBounds(t *testing.T) {
|
||||
const (
|
||||
statusCode = 502
|
||||
marker = "provider-body-marker"
|
||||
)
|
||||
|
||||
ordinaryBody := `{"error":{"message":"` + marker + `"}}`
|
||||
exactLimitBody := ordinaryBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)-len(ordinaryBody))
|
||||
overLimitBody := ordinaryBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)+1-len(ordinaryBody))
|
||||
tests := []struct {
|
||||
name string
|
||||
contentLength int64
|
||||
reader io.Reader
|
||||
wantRead int64
|
||||
wantMessage string
|
||||
}{
|
||||
{
|
||||
name: "recognized envelope",
|
||||
contentLength: int64(len(ordinaryBody)),
|
||||
reader: strings.NewReader(ordinaryBody),
|
||||
wantRead: int64(len(ordinaryBody)),
|
||||
wantMessage: marker,
|
||||
},
|
||||
{
|
||||
name: "exact limit",
|
||||
contentLength: maxProviderErrorResponseBytes,
|
||||
reader: strings.NewReader(exactLimitBody),
|
||||
wantRead: maxProviderErrorResponseBytes,
|
||||
wantMessage: marker,
|
||||
},
|
||||
{
|
||||
name: "declared oversize does not read",
|
||||
contentLength: maxProviderErrorResponseBytes + 1,
|
||||
reader: strings.NewReader(ordinaryBody),
|
||||
wantRead: 0,
|
||||
},
|
||||
{
|
||||
name: "unknown length oversize",
|
||||
contentLength: -1,
|
||||
reader: strings.NewReader(overLimitBody),
|
||||
wantRead: maxProviderErrorResponseBytes + 1,
|
||||
},
|
||||
{
|
||||
name: "underreported oversize",
|
||||
contentLength: maxProviderErrorResponseBytes,
|
||||
reader: strings.NewReader(overLimitBody),
|
||||
wantRead: maxProviderErrorResponseBytes + 1,
|
||||
},
|
||||
{
|
||||
name: "read failure",
|
||||
contentLength: -1,
|
||||
reader: failingReader{err: errors.New("read failure")},
|
||||
wantRead: 0,
|
||||
},
|
||||
{
|
||||
name: "empty body",
|
||||
contentLength: 0,
|
||||
reader: strings.NewReader(""),
|
||||
wantRead: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
reader := &guardedReader{
|
||||
reader: tc.reader,
|
||||
remaining: maxProviderErrorResponseBytes + 1,
|
||||
}
|
||||
err := providerHTTPErrorFromBody(statusCode, tc.contentLength, reader)
|
||||
if err == nil || err.StatusCode() != statusCode {
|
||||
t.Fatalf("error status = %v, want %d", err, statusCode)
|
||||
}
|
||||
if reader.bytes != tc.wantRead {
|
||||
t.Fatalf("body bytes read = %d, want %d", reader.bytes, tc.wantRead)
|
||||
}
|
||||
if reader.violated {
|
||||
t.Fatal("body reader was asked to read beyond the overflow probe")
|
||||
}
|
||||
if got := err.ProviderMessage(); got != tc.wantMessage {
|
||||
t.Fatalf("provider message = %q, want %q", got, tc.wantMessage)
|
||||
}
|
||||
if tc.wantMessage == "" {
|
||||
if err.ProviderCode() != "" || err.ProviderType() != "" || strings.Contains(err.Error(), marker) {
|
||||
t.Fatalf("discarded details were retained: %#v", err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
145
internal/llm/provider_http_error_transport_test.go
Normal file
145
internal/llm/provider_http_error_transport_test.go
Normal file
@@ -0,0 +1,145 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestOpenAICompatibleClientStructuredNonSuccessResponse(t *testing.T) {
|
||||
body := `{"error":{"message":" provider\nmessage\u200b","type":"invalid\ttype","code":1.5e+4}}`
|
||||
responseBody := &countingReadCloser{reader: strings.NewReader(body)}
|
||||
client := newNonSuccessResponseClient(t, http.StatusBadRequest, int64(len(body)), responseBody)
|
||||
|
||||
response, err := client.Generate(context.Background(), ordinaryGenerateRequest())
|
||||
if response != nil {
|
||||
t.Fatalf("response = %#v, want nil", response)
|
||||
}
|
||||
if !errors.Is(err, ErrUnexpectedStatus) {
|
||||
t.Fatalf("errors.Is(%v, ErrUnexpectedStatus) = false", err)
|
||||
}
|
||||
var providerHTTPError *ProviderHTTPError
|
||||
if !errors.As(err, &providerHTTPError) {
|
||||
t.Fatalf("error = %T, want *ProviderHTTPError", err)
|
||||
}
|
||||
if providerHTTPError.StatusCode() != http.StatusBadRequest || providerHTTPError.ProviderCode() != "1.5e+4" || providerHTTPError.ProviderType() != "invalid type" || providerHTTPError.ProviderMessage() != "provider message" {
|
||||
t.Fatalf("provider error = %#v", providerHTTPError)
|
||||
}
|
||||
if !responseBody.closed {
|
||||
t.Fatal("non-success response body was not closed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientNonSuccessBodyOwnership(t *testing.T) {
|
||||
const marker = "provider-body-marker"
|
||||
normalBody := `{"error":{"message":"` + marker + `"}}`
|
||||
overLimitBody := normalBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)+1-len(normalBody))
|
||||
tests := []struct {
|
||||
name string
|
||||
contentLength int64
|
||||
reader io.Reader
|
||||
wantRead int64
|
||||
wantMessage string
|
||||
}{
|
||||
{
|
||||
name: "normal",
|
||||
contentLength: int64(len(normalBody)),
|
||||
reader: strings.NewReader(normalBody),
|
||||
wantRead: int64(len(normalBody)),
|
||||
wantMessage: marker,
|
||||
},
|
||||
{
|
||||
name: "declared oversize",
|
||||
contentLength: maxProviderErrorResponseBytes + 1,
|
||||
reader: strings.NewReader(normalBody),
|
||||
wantRead: 0,
|
||||
},
|
||||
{
|
||||
name: "streamed oversize",
|
||||
contentLength: -1,
|
||||
reader: &guardedReader{
|
||||
reader: strings.NewReader(overLimitBody),
|
||||
remaining: maxProviderErrorResponseBytes + 1,
|
||||
},
|
||||
wantRead: maxProviderErrorResponseBytes + 1,
|
||||
},
|
||||
{
|
||||
name: "underreported oversize",
|
||||
contentLength: maxProviderErrorResponseBytes,
|
||||
reader: &guardedReader{
|
||||
reader: strings.NewReader(overLimitBody),
|
||||
remaining: maxProviderErrorResponseBytes + 1,
|
||||
},
|
||||
wantRead: maxProviderErrorResponseBytes + 1,
|
||||
},
|
||||
{
|
||||
name: "malformed",
|
||||
contentLength: 1,
|
||||
reader: strings.NewReader("{"),
|
||||
wantRead: 1,
|
||||
},
|
||||
{
|
||||
name: "read failure",
|
||||
contentLength: -1,
|
||||
reader: failingReader{err: errors.New("response read failed")},
|
||||
wantRead: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
body := &countingReadCloser{reader: tc.reader}
|
||||
client := newNonSuccessResponseClient(t, http.StatusBadGateway, tc.contentLength, body)
|
||||
|
||||
response, err := client.Generate(context.Background(), ordinaryGenerateRequest())
|
||||
if response != nil {
|
||||
t.Fatalf("response = %#v, want nil", response)
|
||||
}
|
||||
var providerHTTPError *ProviderHTTPError
|
||||
if !errors.As(err, &providerHTTPError) {
|
||||
t.Fatalf("error = %T, want *ProviderHTTPError", err)
|
||||
}
|
||||
if !body.closed {
|
||||
t.Fatal("response body was not closed")
|
||||
}
|
||||
if body.bytesRead != tc.wantRead {
|
||||
t.Fatalf("body bytes read = %d, want %d", body.bytesRead, tc.wantRead)
|
||||
}
|
||||
if body.bytesRead > maxProviderErrorResponseBytes+1 {
|
||||
t.Fatalf("body bytes read = %d, exceeds overflow probe", body.bytesRead)
|
||||
}
|
||||
if guarded, ok := tc.reader.(*guardedReader); ok && guarded.violated {
|
||||
t.Fatal("body reader was asked to read beyond the overflow probe")
|
||||
}
|
||||
if got := providerHTTPError.ProviderMessage(); got != tc.wantMessage {
|
||||
t.Fatalf("provider message = %q, want %q", got, tc.wantMessage)
|
||||
}
|
||||
if tc.wantMessage == "" && (providerHTTPError.ProviderCode() != "" || providerHTTPError.ProviderType() != "" || strings.Contains(providerHTTPError.Error(), marker)) {
|
||||
t.Fatalf("discarded details were retained: %#v", providerHTTPError)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func newNonSuccessResponseClient(t *testing.T, statusCode int, contentLength int64, body io.ReadCloser) *OpenAICompatibleClient {
|
||||
t.Helper()
|
||||
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
BaseURL: "https://provider.example/v1",
|
||||
Model: "m",
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: statusCode,
|
||||
ContentLength: contentLength,
|
||||
Body: body,
|
||||
}, nil
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("construct client: %v", err)
|
||||
}
|
||||
return client
|
||||
}
|
||||
@@ -1,8 +0,0 @@
|
||||
id: aion-2
|
||||
backend: openrouter
|
||||
model: aion-labs/aion-2.0
|
||||
temperature: 0.72
|
||||
reasoning_effort: high
|
||||
top_p: 0.95
|
||||
timeout_seconds: 180
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: claude-fable-latest
|
||||
backend: openrouter
|
||||
model: "~anthropic/claude-fable-latest"
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 600
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: claude-haiku-latest
|
||||
backend: openrouter
|
||||
model: "~anthropic/claude-haiku-latest"
|
||||
reasoning_effort: medium
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: claude-opus-latest
|
||||
backend: openrouter
|
||||
model: "~anthropic/claude-opus-latest"
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: claude-sonnet-latest
|
||||
backend: openrouter
|
||||
model: "~anthropic/claude-sonnet-latest"
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: deepseek-3-2
|
||||
backend: openrouter
|
||||
model: deepseek/deepseek-v3.2
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 180
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: deepseek-4-flash
|
||||
backend: openrouter
|
||||
model: deepseek/deepseek-v4-flash
|
||||
#reasoning_effort: medium
|
||||
timeout_seconds: 180
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: deepseek-4-pro
|
||||
backend: openrouter
|
||||
model: deepseek/deepseek-v4-pro
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 180
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: gemini-2-flash-lite
|
||||
backend: openrouter
|
||||
model: "google/gemini-2.5-flash-lite"
|
||||
#temperature: 0.15
|
||||
reasoning_effort: high
|
||||
#top_p: 0.98
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: gemini-2-flash
|
||||
backend: openrouter
|
||||
model: "google/gemini-2.5-flash"
|
||||
#temperature: 0.15
|
||||
reasoning_effort: high
|
||||
#top_p: 0.98
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: gemini-2-pro
|
||||
backend: openrouter
|
||||
model: "google/gemini-2.5-pro"
|
||||
#temperature: 0.15
|
||||
reasoning_effort: high
|
||||
#top_p: 0.98
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: gemini-3-flash-lite
|
||||
backend: openrouter
|
||||
model: "google/gemini-3.1-flash-lite"
|
||||
#temperature: 0.15
|
||||
reasoning_effort: high
|
||||
#top_p: 0.98
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: gemini-flash-latest
|
||||
backend: openrouter
|
||||
model: "~google/gemini-flash-latest"
|
||||
#temperature: 0.15
|
||||
reasoning_effort: high
|
||||
#top_p: 0.98
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: gemini-pro-latest
|
||||
backend: openrouter
|
||||
model: "~google/gemini-pro-latest"
|
||||
#temperature: 0.15
|
||||
reasoning_effort: high
|
||||
#top_p: 0.98
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: gemma-4-31b
|
||||
backend: openrouter
|
||||
model: google/gemma-4-31b-it:exacto
|
||||
temperature: 0.15
|
||||
reasoning_effort: high
|
||||
top_p: 0.98
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: minimax-m2
|
||||
backend: openrouter
|
||||
model: minimax/minimax-m2.5
|
||||
temperature: 0.5
|
||||
reasoning_effort: high
|
||||
top_p: 0.95
|
||||
timeout_seconds: 180
|
||||
service_tier: flex
|
||||
@@ -1,8 +0,0 @@
|
||||
id: minimax-m3
|
||||
backend: openrouter
|
||||
model: minimax/minimax-m3
|
||||
#temperature: 0.5
|
||||
reasoning_effort: high
|
||||
#top_p: 0.95
|
||||
timeout_seconds: 180
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: mistral-large-2512
|
||||
backend: openrouter
|
||||
model: mistralai/mistral-large-2512
|
||||
temperature: 0.15
|
||||
top_p: 0.98
|
||||
timeout_seconds: 180
|
||||
@@ -1,7 +0,0 @@
|
||||
id: mistral-medium-3-5
|
||||
backend: openrouter
|
||||
model: mistralai/mistral-medium-3-5
|
||||
temperature: 0.15
|
||||
reasoning_effort: high
|
||||
top_p: 0.98
|
||||
timeout_seconds: 180
|
||||
@@ -1,6 +0,0 @@
|
||||
id: mistral-small-3
|
||||
backend: openrouter
|
||||
model: mistralai/mistral-small-3.2-24b-instruct
|
||||
temperature: 0.05
|
||||
top_p: 1.0
|
||||
timeout_seconds: 180
|
||||
@@ -1,7 +0,0 @@
|
||||
id: mistral-small-4
|
||||
backend: openrouter
|
||||
model: mistralai/mistral-small-2603
|
||||
temperature: 0.1
|
||||
reasoning_effort: high
|
||||
top_p: 0.98
|
||||
timeout_seconds: 180
|
||||
@@ -1,6 +0,0 @@
|
||||
id: nemotron-3-ultra
|
||||
backend: openrouter
|
||||
model: nvidia/nemotron-3-ultra-550b-a55b
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 180
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: gpt-5-mini
|
||||
backend: openrouter
|
||||
model: "openai/gpt-5.4-mini"
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,6 +0,0 @@
|
||||
id: gpt-5-nano
|
||||
backend: openrouter
|
||||
model: "openai/gpt-5.4-nano"
|
||||
reasoning_effort: high
|
||||
timeout_seconds: 240
|
||||
service_tier: flex
|
||||
@@ -1,16 +0,0 @@
|
||||
package builtin
|
||||
|
||||
import (
|
||||
"embed"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
)
|
||||
|
||||
const assetRoot = "assets"
|
||||
|
||||
//go:embed assets/**/*.yml
|
||||
var assets embed.FS
|
||||
|
||||
func NewRepository() profile.Repository {
|
||||
return profile.NewFSRepository(assets, assetRoot)
|
||||
}
|
||||
@@ -1,90 +0,0 @@
|
||||
package builtin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io/fs"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/backend"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
func TestBuiltInProfilesValidateThroughRepository(t *testing.T) {
|
||||
repo := NewRepository()
|
||||
ids := loadBuiltInProfileIDs(t)
|
||||
if len(ids) == 0 {
|
||||
t.Fatal("expected built-in profiles")
|
||||
}
|
||||
|
||||
for id := range ids {
|
||||
t.Run(id, func(t *testing.T) {
|
||||
p, err := repo.GetProfile(context.Background(), id)
|
||||
if err != nil {
|
||||
t.Fatalf("expected built-in profile %q to load, got %v", id, err)
|
||||
}
|
||||
if p.ID != id {
|
||||
t.Fatalf("expected profile id %q, got %q", id, p.ID)
|
||||
}
|
||||
if p.BackendID != backend.OpenRouterID {
|
||||
t.Fatalf("expected profile %q to select %q, got %q", id, backend.OpenRouterID, p.BackendID)
|
||||
}
|
||||
if p.Endpoint != "" || p.APIKeyEnv != "" {
|
||||
t.Fatalf("expected profile %q to inherit backend connection settings, got endpoint=%q api_key_env=%q", id, p.Endpoint, p.APIKeyEnv)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuiltInProfilesDoNotContainDuplicateIDsOrRawAPIKeys(t *testing.T) {
|
||||
loadBuiltInProfileIDs(t)
|
||||
}
|
||||
|
||||
func loadBuiltInProfileIDs(t *testing.T) map[string]string {
|
||||
t.Helper()
|
||||
|
||||
ids := map[string]string{}
|
||||
err := fs.WalkDir(assets, assetRoot, func(name string, d fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if d.IsDir() || !strings.HasSuffix(name, ".yml") {
|
||||
return nil
|
||||
}
|
||||
|
||||
data, err := assets.ReadFile(name)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read built-in profile %s: %v", name, err)
|
||||
}
|
||||
|
||||
var raw map[string]any
|
||||
if err := yaml.Unmarshal(data, &raw); err != nil {
|
||||
t.Fatalf("failed to decode built-in profile %s: %v", name, err)
|
||||
}
|
||||
if _, ok := raw["api_key"]; ok {
|
||||
t.Fatalf("built-in profile %s contains raw api_key", name)
|
||||
}
|
||||
if raw["backend"] != backend.OpenRouterID {
|
||||
t.Fatalf("built-in profile %s does not select %q", name, backend.OpenRouterID)
|
||||
}
|
||||
if _, ok := raw["endpoint"]; ok {
|
||||
t.Fatalf("built-in profile %s repeats endpoint", name)
|
||||
}
|
||||
if _, ok := raw["api_key_env"]; ok {
|
||||
t.Fatalf("built-in profile %s repeats api_key_env", name)
|
||||
}
|
||||
id, ok := raw["id"].(string)
|
||||
if !ok || strings.TrimSpace(id) == "" {
|
||||
t.Fatalf("built-in profile %s has missing id", name)
|
||||
}
|
||||
if previous, ok := ids[id]; ok {
|
||||
t.Fatalf("duplicate built-in profile id %q in %s and %s", id, previous, name)
|
||||
}
|
||||
ids[id] = name
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to walk built-in profiles: %v", err)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
47
internal/profile/definition.go
Normal file
47
internal/profile/definition.go
Normal file
@@ -0,0 +1,47 @@
|
||||
package profile
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
// NormalizeAndValidateDefinition normalizes and validates one source-local
|
||||
// profile definition without resolving a base profile.
|
||||
func NormalizeAndValidateDefinition(profile *domain.ExecutionProfile) error {
|
||||
if profile == nil {
|
||||
return errors.New("profile is required")
|
||||
}
|
||||
|
||||
profile.ID = strings.TrimSpace(profile.ID)
|
||||
profile.BaseProfileID = strings.TrimSpace(profile.BaseProfileID)
|
||||
profile.BackendID = strings.TrimSpace(profile.BackendID)
|
||||
profile.Endpoint = strings.TrimSpace(profile.Endpoint)
|
||||
|
||||
if profile.ID == "" {
|
||||
return errors.New("id is required")
|
||||
}
|
||||
if profile.Endpoint != "" {
|
||||
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(profile.Endpoint)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
profile.Endpoint = endpoint
|
||||
}
|
||||
if profile.BaseProfileID == "" {
|
||||
if profile.BackendID == "" && profile.Endpoint == "" {
|
||||
return errors.New("backend or endpoint is required")
|
||||
}
|
||||
if strings.TrimSpace(profile.Model) == "" {
|
||||
return errors.New("model is required")
|
||||
}
|
||||
}
|
||||
|
||||
return domain.ValidateExecutionTargetSettings(domain.ExecutionTarget{
|
||||
Temperature: profile.Temperature,
|
||||
MaxTokens: profile.MaxTokens,
|
||||
TopP: profile.TopP,
|
||||
TimeoutSeconds: profile.TimeoutSeconds,
|
||||
})
|
||||
}
|
||||
98
internal/profile/eager_repository.go
Normal file
98
internal/profile/eager_repository.go
Normal file
@@ -0,0 +1,98 @@
|
||||
package profile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/filecatalog"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
)
|
||||
|
||||
// LoadedProfileMetadata identifies one profile accepted by LoadFSRepository.
|
||||
type LoadedProfileMetadata struct {
|
||||
// ID is the normalized profile ID.
|
||||
ID string
|
||||
// Path is the safe root-relative source path.
|
||||
Path string
|
||||
// ExplicitFields lists the sorted top-level YAML fields present in source.
|
||||
ExplicitFields []string
|
||||
}
|
||||
|
||||
// LoadFSRepository eagerly validates every profile under root and returns an
|
||||
// immutable raw repository and independently owned source metadata.
|
||||
func LoadFSRepository(ctx context.Context, fsys fs.FS, root string) (Repository, []LoadedProfileMetadata, error) {
|
||||
if fsys == nil {
|
||||
return nil, nil, fmt.Errorf("failed to read profile directory: filesystem is nil")
|
||||
}
|
||||
paths, err := filecatalog.FindFSYAMLFiles(ctx, fsys, root)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to read profile directory: %w", err)
|
||||
}
|
||||
|
||||
repository := &loadedRepository{profiles: make(map[string]domain.ExecutionProfile, len(paths))}
|
||||
metadata := make([]LoadedProfileMetadata, 0, len(paths))
|
||||
for _, path := range paths {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
data, err := fs.ReadFile(fsys, path)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to read profile file %s: %w", filecatalog.DisplayPath(root, path), err)
|
||||
}
|
||||
fileMetadata, err := readProfileFileMetadata(data)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("%w: %s", ErrInvalidYAML, filecatalog.DisplayPath(root, path))
|
||||
}
|
||||
if fileMetadata.hasRawAPIKey {
|
||||
return nil, nil, fmt.Errorf("%w: %s", ErrRawAPIKeyNotAllowed, filecatalog.DisplayPath(root, path))
|
||||
}
|
||||
definition, err := decodeProfile(data)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("%w: %s", ErrInvalidYAML, filecatalog.DisplayPath(root, path))
|
||||
}
|
||||
definition.ExtraParams, err = jsonvalue.CopyMap(definition.ExtraParams)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("%w: %s", ErrInvalidProfile, filecatalog.DisplayPath(root, path))
|
||||
}
|
||||
if err := NormalizeAndValidateDefinition(definition); err != nil {
|
||||
return nil, nil, fmt.Errorf("%w: %s", ErrInvalidProfile, filecatalog.DisplayPath(root, path))
|
||||
}
|
||||
if _, exists := repository.profiles[definition.ID]; exists {
|
||||
return nil, nil, fmt.Errorf("%w: %s: duplicate profile ID", ErrInvalidProfile, filecatalog.DisplayPath(root, path))
|
||||
}
|
||||
repository.profiles[definition.ID] = *definition
|
||||
fields := append([]string(nil), fileMetadata.explicitFields...)
|
||||
sort.Strings(fields)
|
||||
metadata = append(metadata, LoadedProfileMetadata{
|
||||
ID: definition.ID,
|
||||
Path: filecatalog.DisplayPath(root, path),
|
||||
ExplicitFields: fields,
|
||||
})
|
||||
}
|
||||
sort.Slice(metadata, func(left, right int) bool { return metadata[left].ID < metadata[right].ID })
|
||||
return repository, metadata, nil
|
||||
}
|
||||
|
||||
type loadedRepository struct {
|
||||
profiles map[string]domain.ExecutionProfile
|
||||
}
|
||||
|
||||
func (r *loadedRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
definition, found := r.profiles[strings.TrimSpace(id)]
|
||||
if !found {
|
||||
return nil, ErrProfileNotFound
|
||||
}
|
||||
extraParams, err := jsonvalue.CopyMap(definition.ExtraParams)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("copy loaded profile %q: %w", definition.ID, err)
|
||||
}
|
||||
definition.ExtraParams = extraParams
|
||||
return &definition, nil
|
||||
}
|
||||
110
internal/profile/eager_repository_test.go
Normal file
110
internal/profile/eager_repository_test.go
Normal file
@@ -0,0 +1,110 @@
|
||||
package profile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io/fs"
|
||||
"reflect"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
)
|
||||
|
||||
func TestLoadFSRepositoryLoadsSortedIndependentProfiles(t *testing.T) {
|
||||
fsys := fstest.MapFS{
|
||||
"profiles/z.yml": {Data: []byte("id: z\nbackend: local\nmodel: z-model\nextra_params:\n nested:\n value: one\n")},
|
||||
"profiles/a.yml": {Data: []byte("id: a\nbackend: local\nmodel: a-model\nendpoint: ''\n")},
|
||||
}
|
||||
|
||||
repository, metadata, err := LoadFSRepository(context.Background(), fsys, "profiles")
|
||||
if err != nil {
|
||||
t.Fatalf("load repository: %v", err)
|
||||
}
|
||||
if len(metadata) != 2 || metadata[0].ID != "a" || metadata[1].ID != "z" || metadata[0].Path != "a.yml" {
|
||||
t.Fatalf("unexpected metadata: %#v", metadata)
|
||||
}
|
||||
if len(metadata[0].ExplicitFields) != 4 || metadata[0].ExplicitFields[0] != "backend" || metadata[0].ExplicitFields[1] != "endpoint" {
|
||||
t.Fatalf("expected explicitly empty endpoint metadata, got %#v", metadata[0].ExplicitFields)
|
||||
}
|
||||
metadata[1].ExplicitFields[0] = "changed"
|
||||
first, err := repository.GetProfile(context.Background(), "z")
|
||||
if err != nil {
|
||||
t.Fatalf("load profile: %v", err)
|
||||
}
|
||||
first.ExtraParams["nested"].(map[string]any)["value"] = "changed"
|
||||
second, err := repository.GetProfile(context.Background(), "z")
|
||||
if err != nil {
|
||||
t.Fatalf("reload profile: %v", err)
|
||||
}
|
||||
if second.ExtraParams["nested"].(map[string]any)["value"] != "one" {
|
||||
t.Fatalf("profile value was not defensively copied: %#v", second)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadFSRepositoryRejectsInvalidProfiles(t *testing.T) {
|
||||
tests := map[string]fstest.MapFS{
|
||||
"duplicate IDs": {
|
||||
"profiles/one.yml": {Data: []byte("id: duplicate\nbackend: local\nmodel: one\n")},
|
||||
"profiles/two.yml": {Data: []byte("id: duplicate\nbackend: local\nmodel: two\n")},
|
||||
},
|
||||
"raw API key": {
|
||||
"profiles/one.yml": {Data: []byte("id: one\nbackend: local\nmodel: one\napi_key: forbidden\n")},
|
||||
},
|
||||
"multiple documents": {
|
||||
"profiles/one.yml": {Data: []byte("id: one\nbackend: local\nmodel: one\n---\nid: two\n")},
|
||||
},
|
||||
}
|
||||
for name, fsys := range tests {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
_, _, err := LoadFSRepository(context.Background(), fsys, "profiles")
|
||||
if err == nil {
|
||||
t.Fatal("expected load failure")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadedRepositoryHonorsCancellationAndMissingProfiles(t *testing.T) {
|
||||
repository, _, err := LoadFSRepository(context.Background(), fstest.MapFS{
|
||||
"profiles/one.yml": {Data: []byte("id: one\nbackend: local\nmodel: one\n")},
|
||||
}, "profiles")
|
||||
if err != nil {
|
||||
t.Fatalf("load repository: %v", err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if _, err := repository.GetProfile(ctx, "one"); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected cancellation, got %v", err)
|
||||
}
|
||||
if _, err := repository.GetProfile(context.Background(), "missing"); !errors.Is(err, ErrProfileNotFound) {
|
||||
t.Fatalf("expected missing profile, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadFSRepositoryMatchesPointLookupForRawProfiles(t *testing.T) {
|
||||
fsys := fstest.MapFS{
|
||||
"profiles/base.yml": {Data: []byte("id: base\nbackend: local\nmodel: base-model\n")},
|
||||
"profiles/derived.yml": {Data: []byte("id: derived\nbase_profile: base\nreasoning_effort: high\n")},
|
||||
}
|
||||
eager, _, err := LoadFSRepository(context.Background(), fsys, "profiles")
|
||||
if err != nil {
|
||||
t.Fatalf("load eager repository: %v", err)
|
||||
}
|
||||
pointLookup := NewFSRepository(fsys, "profiles")
|
||||
for _, id := range []string{"base", "derived"} {
|
||||
t.Run(id, func(t *testing.T) {
|
||||
got, err := eager.GetProfile(context.Background(), id)
|
||||
if err != nil {
|
||||
t.Fatalf("load eager profile: %v", err)
|
||||
}
|
||||
want, err := pointLookup.GetProfile(context.Background(), id)
|
||||
if err != nil {
|
||||
t.Fatalf("load point-in-time profile: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("raw profiles differ: got %#v, want %#v", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
var _ fs.FS = fstest.MapFS{}
|
||||
@@ -127,12 +127,11 @@ func loadProfile(ctx context.Context, fsys fs.FS, root string, id string) (*doma
|
||||
if prof.ID != id {
|
||||
continue
|
||||
}
|
||||
prof.BackendID = strings.TrimSpace(prof.BackendID)
|
||||
prof.ExtraParams, err = jsonvalue.CopyMap(prof.ExtraParams)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidProfile, relPath, err)
|
||||
}
|
||||
if err := normalizeAndValidateProfile(prof); err != nil {
|
||||
if err := NormalizeAndValidateDefinition(prof); err != nil {
|
||||
if errors.Is(err, ErrRawAPIKeyNotAllowed) {
|
||||
return nil, fmt.Errorf("%w: %s", err, relPath)
|
||||
}
|
||||
@@ -165,8 +164,9 @@ type profileMatch struct {
|
||||
}
|
||||
|
||||
type profileFileMetadata struct {
|
||||
ids []string
|
||||
hasRawAPIKey bool
|
||||
ids []string
|
||||
hasRawAPIKey bool
|
||||
explicitFields []string
|
||||
}
|
||||
|
||||
func readProfileFileMetadata(data []byte) (profileFileMetadata, error) {
|
||||
@@ -207,6 +207,7 @@ func profileMetadataFromNode(node *yaml.Node) profileFileMetadata {
|
||||
for i := 0; i+1 < len(mapping.Content); i += 2 {
|
||||
key := mapping.Content[i]
|
||||
value := mapping.Content[i+1]
|
||||
metadata.explicitFields = append(metadata.explicitFields, key.Value)
|
||||
switch key.Value {
|
||||
case "id":
|
||||
metadata.ids = append(metadata.ids, strings.TrimSpace(value.Value))
|
||||
@@ -229,6 +230,7 @@ func (m profileFileMetadata) matchesID(id string) bool {
|
||||
func (m *profileFileMetadata) merge(other profileFileMetadata) {
|
||||
m.ids = append(m.ids, other.ids...)
|
||||
m.hasRawAPIKey = m.hasRawAPIKey || other.hasRawAPIKey
|
||||
m.explicitFields = append(m.explicitFields, other.explicitFields...)
|
||||
}
|
||||
|
||||
func decodeProfile(data []byte) (*domain.ExecutionProfile, error) {
|
||||
@@ -255,30 +257,3 @@ func requireYAMLStreamEnd(decoder *yaml.Decoder) error {
|
||||
}
|
||||
return errors.New("profile file must contain exactly one YAML document")
|
||||
}
|
||||
|
||||
func normalizeAndValidateProfile(p *domain.ExecutionProfile) error {
|
||||
if strings.TrimSpace(p.ID) == "" {
|
||||
return errors.New("id is required")
|
||||
}
|
||||
p.Endpoint = strings.TrimSpace(p.Endpoint)
|
||||
if strings.TrimSpace(p.BackendID) == "" && p.Endpoint == "" {
|
||||
return errors.New("backend or endpoint is required")
|
||||
}
|
||||
if p.Endpoint != "" {
|
||||
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(p.Endpoint)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.Endpoint = endpoint
|
||||
}
|
||||
if strings.TrimSpace(p.Model) == "" {
|
||||
return errors.New("model is required")
|
||||
}
|
||||
|
||||
return domain.ValidateExecutionTargetSettings(domain.ExecutionTarget{
|
||||
Temperature: p.Temperature,
|
||||
MaxTokens: p.MaxTokens,
|
||||
TopP: p.TopP,
|
||||
TimeoutSeconds: p.TimeoutSeconds,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -858,6 +858,107 @@ top_p: .inf
|
||||
})
|
||||
}
|
||||
|
||||
func TestProfileRepositoriesValidateDerivedDefinitions(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
files map[string]string
|
||||
wantError error
|
||||
wantBaseID string
|
||||
wantProfile bool
|
||||
}{
|
||||
{
|
||||
name: "alias is locally valid and normalizes base id",
|
||||
files: map[string]string{"alias.yaml": `
|
||||
id: selected-profile
|
||||
base_profile: " base-profile "
|
||||
`},
|
||||
wantBaseID: "base-profile",
|
||||
wantProfile: true,
|
||||
},
|
||||
{
|
||||
name: "derived endpoint remains valid",
|
||||
files: map[string]string{"invalid.yaml": `
|
||||
id: selected-profile
|
||||
base_profile: base-profile
|
||||
endpoint: /v1
|
||||
`},
|
||||
wantError: ErrInvalidProfile,
|
||||
},
|
||||
{
|
||||
name: "derived settings remain valid",
|
||||
files: map[string]string{"invalid.yaml": `
|
||||
id: selected-profile
|
||||
base_profile: base-profile
|
||||
top_p: 1.1
|
||||
`},
|
||||
wantError: ErrInvalidProfile,
|
||||
},
|
||||
{
|
||||
name: "derived extra params remain valid",
|
||||
files: map[string]string{"invalid.yaml": `
|
||||
id: selected-profile
|
||||
base_profile: base-profile
|
||||
extra_params:
|
||||
timestamp: 2026-08-11T12:34:56Z
|
||||
`},
|
||||
wantError: ErrInvalidProfile,
|
||||
},
|
||||
{
|
||||
name: "derived raw key remains prohibited",
|
||||
files: map[string]string{"invalid.yaml": `
|
||||
id: selected-profile
|
||||
base_profile: base-profile
|
||||
api_key: secret
|
||||
`},
|
||||
wantError: ErrRawAPIKeyNotAllowed,
|
||||
},
|
||||
{
|
||||
name: "derived duplicate id remains invalid",
|
||||
files: map[string]string{
|
||||
"first.yaml": "id: selected-profile\nbase_profile: first-base\n",
|
||||
"second.yaml": "id: selected-profile\nbase_profile: second-base\n",
|
||||
},
|
||||
wantError: ErrInvalidProfile,
|
||||
},
|
||||
{
|
||||
name: "derived extra document remains invalid",
|
||||
files: map[string]string{"invalid.yaml": `
|
||||
id: selected-profile
|
||||
base_profile: base-profile
|
||||
---
|
||||
id: other
|
||||
`},
|
||||
wantError: ErrInvalidYAML,
|
||||
},
|
||||
{
|
||||
name: "standalone profile remains complete",
|
||||
files: map[string]string{"invalid.yaml": "id: selected-profile\n"},
|
||||
wantError: ErrInvalidProfile,
|
||||
},
|
||||
}
|
||||
|
||||
for _, source := range profileRepositorySources() {
|
||||
for _, tc := range tests {
|
||||
t.Run(source.name+"/"+tc.name, func(t *testing.T) {
|
||||
repo := source.newRepository(t, tc.files)
|
||||
got, err := repo.GetProfile(context.Background(), "selected-profile")
|
||||
if tc.wantError != nil {
|
||||
if !errors.Is(err, tc.wantError) {
|
||||
t.Fatalf("error = %v, want %v", err, tc.wantError)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil || !tc.wantProfile {
|
||||
t.Fatalf("profile = %+v, error = %v, want valid derived definition", got, err)
|
||||
}
|
||||
if got.BaseProfileID != tc.wantBaseID {
|
||||
t.Fatalf("BaseProfileID = %q, want %q", got.BaseProfileID, tc.wantBaseID)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOverlayRepository(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
primaryProfile := &domain.ExecutionProfile{ID: "shared", Endpoint: "http://primary", Model: "primary"}
|
||||
|
||||
160
internal/profile/resolving_repository.go
Normal file
160
internal/profile/resolving_repository.go
Normal file
@@ -0,0 +1,160 @@
|
||||
package profile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
)
|
||||
|
||||
const maximumProfileChainLength = 32
|
||||
|
||||
type resolvingRepository struct {
|
||||
source Repository
|
||||
}
|
||||
|
||||
// NewResolvingRepository resolves inherited profile definitions from source.
|
||||
func NewResolvingRepository(source Repository) Repository {
|
||||
return &resolvingRepository{source: source}
|
||||
}
|
||||
|
||||
func (r *resolvingRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
||||
if r == nil || r.source == nil {
|
||||
return nil, fmt.Errorf("%w: profile repository is required", ErrInvalidProfile)
|
||||
}
|
||||
requestedID := strings.TrimSpace(id)
|
||||
if requestedID == "" {
|
||||
return nil, fmt.Errorf("%w: profile id is required", ErrInvalidProfile)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
profile, err := r.getRawProfile(ctx, requestedID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if profile == nil {
|
||||
return nil, fmt.Errorf("%w: selected profile %q is nil", ErrInvalidProfile, requestedID)
|
||||
}
|
||||
|
||||
chain := []*domain.ExecutionProfile{profile}
|
||||
chainIDs := []string{requestedID}
|
||||
visited := map[string]struct{}{requestedID: {}}
|
||||
current := profile
|
||||
|
||||
for {
|
||||
baseID := strings.TrimSpace(current.BaseProfileID)
|
||||
if baseID == "" {
|
||||
break
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, seen := visited[baseID]; seen {
|
||||
return nil, fmt.Errorf("%w: profile inheritance cycle %s", ErrInvalidProfile, joinProfileChain(chainIDs, baseID))
|
||||
}
|
||||
if len(chain) >= maximumProfileChainLength {
|
||||
return nil, fmt.Errorf("%w: profile inheritance chain exceeds %d profiles: %s", ErrInvalidProfile, maximumProfileChainLength, joinProfileChain(chainIDs, baseID))
|
||||
}
|
||||
|
||||
base, err := r.getRawProfile(ctx, baseID)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrProfileNotFound) {
|
||||
return nil, fmt.Errorf("%w: base profile %q is missing in chain %s", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID))
|
||||
}
|
||||
return nil, fmt.Errorf("%w: failed to load base profile %q in chain %s: %w", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID), err)
|
||||
}
|
||||
if base == nil {
|
||||
return nil, fmt.Errorf("%w: base profile %q is nil in chain %s", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID))
|
||||
}
|
||||
|
||||
chain = append(chain, base)
|
||||
chainIDs = append(chainIDs, baseID)
|
||||
visited[baseID] = struct{}{}
|
||||
current = base
|
||||
}
|
||||
|
||||
resolved, err := mergeProfileChain(chain)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: resolved profile chain %s: %w", ErrInvalidProfile, strings.Join(chainIDs, " -> "), err)
|
||||
}
|
||||
if err := validateResolvedProfile(resolved); err != nil {
|
||||
return nil, fmt.Errorf("%w: resolved profile chain %s: %w", ErrInvalidProfile, strings.Join(chainIDs, " -> "), err)
|
||||
}
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
func (r *resolvingRepository) getRawProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
||||
profile, err := r.source.GetProfile(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return profile, nil
|
||||
}
|
||||
|
||||
func joinProfileChain(chain []string, next string) string {
|
||||
return strings.Join(append(append([]string(nil), chain...), next), " -> ")
|
||||
}
|
||||
|
||||
func mergeProfileChain(chain []*domain.ExecutionProfile) (*domain.ExecutionProfile, error) {
|
||||
resolved := &domain.ExecutionProfile{ID: chain[0].ID}
|
||||
for index := len(chain) - 1; index >= 0; index-- {
|
||||
definition := chain[index]
|
||||
if strings.TrimSpace(definition.BackendID) != "" {
|
||||
resolved.BackendID = definition.BackendID
|
||||
}
|
||||
if strings.TrimSpace(definition.Endpoint) != "" {
|
||||
resolved.Endpoint = definition.Endpoint
|
||||
}
|
||||
if strings.TrimSpace(definition.Model) != "" {
|
||||
resolved.Model = definition.Model
|
||||
}
|
||||
if definition.Temperature != 0 {
|
||||
resolved.Temperature = definition.Temperature
|
||||
}
|
||||
if definition.MaxTokens != 0 {
|
||||
resolved.MaxTokens = definition.MaxTokens
|
||||
}
|
||||
if definition.TopP != 0 {
|
||||
resolved.TopP = definition.TopP
|
||||
}
|
||||
if definition.TimeoutSeconds != 0 {
|
||||
resolved.TimeoutSeconds = definition.TimeoutSeconds
|
||||
}
|
||||
if strings.TrimSpace(definition.ServiceTier) != "" {
|
||||
resolved.ServiceTier = definition.ServiceTier
|
||||
}
|
||||
if strings.TrimSpace(definition.ReasoningEffort) != "" {
|
||||
resolved.ReasoningEffort = definition.ReasoningEffort
|
||||
}
|
||||
if strings.TrimSpace(definition.APIKeyEnv) != "" {
|
||||
resolved.APIKeyEnv = definition.APIKeyEnv
|
||||
}
|
||||
resolved.APIKeyRequired = resolved.APIKeyRequired || definition.APIKeyRequired
|
||||
if len(definition.ExtraParams) != 0 {
|
||||
extraParams, err := jsonvalue.CopyMap(definition.ExtraParams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resolved.ExtraParams = extraParams
|
||||
}
|
||||
}
|
||||
resolved.ID = chain[0].ID
|
||||
resolved.BaseProfileID = ""
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
func validateResolvedProfile(profile *domain.ExecutionProfile) error {
|
||||
if profile == nil {
|
||||
return errors.New("resolved profile is required")
|
||||
}
|
||||
profile.BaseProfileID = ""
|
||||
return NormalizeAndValidateDefinition(profile)
|
||||
}
|
||||
418
internal/profile/resolving_repository_test.go
Normal file
418
internal/profile/resolving_repository_test.go
Normal file
@@ -0,0 +1,418 @@
|
||||
package profile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
func TestResolvingRepositoryMergesProfileChain(t *testing.T) {
|
||||
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"leaf": {
|
||||
ID: "leaf",
|
||||
BaseProfileID: "middle",
|
||||
BackendID: "leaf-backend",
|
||||
TopP: 0.8,
|
||||
TimeoutSeconds: 45,
|
||||
ReasoningEffort: "high",
|
||||
},
|
||||
"middle": {
|
||||
ID: "middle",
|
||||
BaseProfileID: "root",
|
||||
Endpoint: "https://middle.example/v1",
|
||||
Model: "middle-model",
|
||||
MaxTokens: 256,
|
||||
APIKeyEnv: "MIDDLE_API_KEY",
|
||||
APIKeyRequired: true,
|
||||
ExtraParams: map[string]any{"middle": map[string]any{"value": "middle"}},
|
||||
},
|
||||
"root": {
|
||||
ID: "root",
|
||||
BackendID: "root-backend",
|
||||
Endpoint: "https://root.example/v1",
|
||||
Model: "root-model",
|
||||
Temperature: 0.3,
|
||||
ServiceTier: "priority",
|
||||
ExtraParams: map[string]any{"root": "value"},
|
||||
},
|
||||
}}
|
||||
|
||||
got, err := NewResolvingRepository(repo).GetProfile(context.Background(), "leaf")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve profile: %v", err)
|
||||
}
|
||||
want := &domain.ExecutionProfile{
|
||||
ID: "leaf",
|
||||
BackendID: "leaf-backend",
|
||||
Endpoint: "https://middle.example/v1",
|
||||
Model: "middle-model",
|
||||
Temperature: 0.3,
|
||||
MaxTokens: 256,
|
||||
TopP: 0.8,
|
||||
TimeoutSeconds: 45,
|
||||
ServiceTier: "priority",
|
||||
ReasoningEffort: "high",
|
||||
APIKeyEnv: "MIDDLE_API_KEY",
|
||||
APIKeyRequired: true,
|
||||
ExtraParams: map[string]any{"middle": map[string]any{"value": "middle"}},
|
||||
}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("resolved profile:\n got %#v\nwant %#v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvingRepositoryRejectsMissingSourceAndProfileID(t *testing.T) {
|
||||
if _, err := NewResolvingRepository(nil).GetProfile(context.Background(), "profile"); !errors.Is(err, ErrInvalidProfile) {
|
||||
t.Fatalf("nil source error = %v, want ErrInvalidProfile", err)
|
||||
}
|
||||
|
||||
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{}}
|
||||
if _, err := NewResolvingRepository(repo).GetProfile(context.Background(), " \t "); !errors.Is(err, ErrInvalidProfile) {
|
||||
t.Fatalf("blank id error = %v, want ErrInvalidProfile", err)
|
||||
}
|
||||
if got := repo.callCount(" "); got != 0 {
|
||||
t.Fatalf("blank id looked up source %d times", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvingRepositoryCopiesExtraParams(t *testing.T) {
|
||||
baseParams := map[string]any{"nested": map[string]any{"value": "base"}}
|
||||
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"child": {ID: "child", BaseProfileID: "base"},
|
||||
"base": {
|
||||
ID: "base",
|
||||
Endpoint: "https://base.example/v1",
|
||||
Model: "model",
|
||||
ExtraParams: baseParams,
|
||||
},
|
||||
}}
|
||||
resolver := NewResolvingRepository(repo)
|
||||
|
||||
first, err := resolver.GetProfile(context.Background(), "child")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve inherited map: %v", err)
|
||||
}
|
||||
first.ExtraParams["nested"].(map[string]any)["value"] = "mutated"
|
||||
second, err := resolver.GetProfile(context.Background(), "child")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve inherited map again: %v", err)
|
||||
}
|
||||
if got := second.ExtraParams["nested"].(map[string]any)["value"]; got != "base" {
|
||||
t.Fatalf("later result retained mutation: %v", got)
|
||||
}
|
||||
if got := baseParams["nested"].(map[string]any)["value"]; got != "base" {
|
||||
t.Fatalf("source map retained mutation: %v", got)
|
||||
}
|
||||
|
||||
repo.set("child", &domain.ExecutionProfile{
|
||||
ID: "child",
|
||||
BaseProfileID: "base",
|
||||
ExtraParams: map[string]any{"child": "replacement"},
|
||||
})
|
||||
replaced, err := resolver.GetProfile(context.Background(), "child")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve replacement map: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(replaced.ExtraParams, map[string]any{"child": "replacement"}) {
|
||||
t.Fatalf("extra params = %#v, want complete child replacement", replaced.ExtraParams)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvingRepositoryUsesRawOverlayForEachLookup(t *testing.T) {
|
||||
leafSource := NewFSRepository(profileTestFS(map[string]string{
|
||||
"leaf.yaml": "id: leaf\nbase_profile: base\n",
|
||||
}), ".")
|
||||
fallback := NewFSRepository(profileTestFS(map[string]string{
|
||||
"base.yaml": "id: base\nendpoint: https://fallback.example/v1\nmodel: fallback-model\n",
|
||||
}), ".")
|
||||
overlay := NewOverlayRepository(leafSource, fallback)
|
||||
resolver := NewResolvingRepository(overlay)
|
||||
|
||||
got, err := resolver.GetProfile(context.Background(), "leaf")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve fallback base: %v", err)
|
||||
}
|
||||
if got.Model != "fallback-model" {
|
||||
t.Fatalf("fallback base model = %q", got.Model)
|
||||
}
|
||||
|
||||
shadowing := NewOverlayRepository(NewFSRepository(profileTestFS(map[string]string{
|
||||
"leaf.yaml": "id: leaf\nbase_profile: base\n",
|
||||
"base.yaml": "id: base\nendpoint: https://primary.example/v1\nmodel: primary-model\n",
|
||||
}), "."), fallback)
|
||||
got, err = NewResolvingRepository(shadowing).GetProfile(context.Background(), "leaf")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve shadowed base: %v", err)
|
||||
}
|
||||
if got.Model != "primary-model" || got.Endpoint != "https://primary.example/v1" {
|
||||
t.Fatalf("shadowed base = %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvingRepositoryReportsSafetyAndSourceErrors(t *testing.T) {
|
||||
sourceErr := errors.New("source failure")
|
||||
tests := []struct {
|
||||
name string
|
||||
repo *resolvingTestRepository
|
||||
id string
|
||||
want []error
|
||||
wantNot error
|
||||
contains []string
|
||||
}{
|
||||
{
|
||||
name: "missing selected profile preserves not found",
|
||||
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{}},
|
||||
id: "missing",
|
||||
want: []error{ErrProfileNotFound},
|
||||
wantNot: ErrInvalidProfile,
|
||||
},
|
||||
{
|
||||
name: "missing base is invalid but not not found",
|
||||
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"leaf": {ID: "leaf", BaseProfileID: "missing"},
|
||||
}},
|
||||
id: "leaf",
|
||||
want: []error{ErrInvalidProfile},
|
||||
wantNot: ErrProfileNotFound,
|
||||
contains: []string{"missing", "leaf -> missing"},
|
||||
},
|
||||
{
|
||||
name: "direct cycle",
|
||||
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"a": {ID: "a", BaseProfileID: "a"},
|
||||
}},
|
||||
id: "a",
|
||||
want: []error{ErrInvalidProfile},
|
||||
contains: []string{"a -> a"},
|
||||
},
|
||||
{
|
||||
name: "indirect cycle",
|
||||
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"a": {ID: "a", BaseProfileID: "b"},
|
||||
"b": {ID: "b", BaseProfileID: "c"},
|
||||
"c": {ID: "c", BaseProfileID: "a"},
|
||||
}},
|
||||
id: "a",
|
||||
want: []error{ErrInvalidProfile},
|
||||
contains: []string{"a -> b -> c -> a"},
|
||||
},
|
||||
{
|
||||
name: "nil result",
|
||||
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"leaf": nil,
|
||||
}},
|
||||
id: "leaf",
|
||||
want: []error{ErrInvalidProfile},
|
||||
},
|
||||
{
|
||||
name: "incomplete resolved profile",
|
||||
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"leaf": {ID: "leaf", BaseProfileID: "base"},
|
||||
"base": {ID: "base", Model: "model"},
|
||||
}},
|
||||
id: "leaf",
|
||||
want: []error{ErrInvalidProfile},
|
||||
},
|
||||
{
|
||||
name: "base source error is retained",
|
||||
repo: &resolvingTestRepository{
|
||||
profiles: map[string]*domain.ExecutionProfile{"leaf": {ID: "leaf", BaseProfileID: "base"}},
|
||||
errors: map[string]error{"base": sourceErr},
|
||||
},
|
||||
id: "leaf",
|
||||
want: []error{ErrInvalidProfile, sourceErr},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := NewResolvingRepository(tc.repo).GetProfile(context.Background(), tc.id)
|
||||
for _, want := range tc.want {
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want %v", err, want)
|
||||
}
|
||||
}
|
||||
if tc.wantNot != nil && errors.Is(err, tc.wantNot) {
|
||||
t.Fatalf("error = %v, must not match %v", err, tc.wantNot)
|
||||
}
|
||||
for _, fragment := range tc.contains {
|
||||
if !strings.Contains(err.Error(), fragment) {
|
||||
t.Fatalf("error = %v, want %q", err, fragment)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvingRepositoryEnforcesChainLength(t *testing.T) {
|
||||
for _, count := range []int{maximumProfileChainLength, maximumProfileChainLength + 1} {
|
||||
t.Run(fmt.Sprintf("%d profiles", count), func(t *testing.T) {
|
||||
profiles := make(map[string]*domain.ExecutionProfile, count)
|
||||
for index := 1; index <= count; index++ {
|
||||
id := fmt.Sprintf("profile-%d", index)
|
||||
definition := &domain.ExecutionProfile{ID: id}
|
||||
if index == count {
|
||||
definition.Endpoint = "https://root.example/v1"
|
||||
definition.Model = "model"
|
||||
} else {
|
||||
definition.BaseProfileID = fmt.Sprintf("profile-%d", index+1)
|
||||
}
|
||||
profiles[id] = definition
|
||||
}
|
||||
|
||||
got, err := NewResolvingRepository(&resolvingTestRepository{profiles: profiles}).GetProfile(context.Background(), "profile-1")
|
||||
if count == maximumProfileChainLength {
|
||||
if err != nil || got == nil {
|
||||
t.Fatalf("profile = %+v, error = %v, want accepted chain", got, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if !errors.Is(err, ErrInvalidProfile) {
|
||||
t.Fatalf("error = %v, want ErrInvalidProfile", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvingRepositoryIsFreshAndCancellationAware(t *testing.T) {
|
||||
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
|
||||
"leaf": {ID: "leaf", BaseProfileID: "base"},
|
||||
"base": {ID: "base", Endpoint: "https://base.example/v1", Model: "first", ExtraParams: map[string]any{"nested": map[string]any{"value": "first"}}},
|
||||
}}
|
||||
resolver := NewResolvingRepository(repo)
|
||||
|
||||
first, err := resolver.GetProfile(context.Background(), "leaf")
|
||||
if err != nil || first.Model != "first" {
|
||||
t.Fatalf("first result=(%+v, %v)", first, err)
|
||||
}
|
||||
repo.set("base", &domain.ExecutionProfile{ID: "base", Endpoint: "https://base.example/v1", Model: "second", ExtraParams: map[string]any{"nested": map[string]any{"value": "second"}}})
|
||||
second, err := resolver.GetProfile(context.Background(), "leaf")
|
||||
if err != nil || second.Model != "second" {
|
||||
t.Fatalf("second result=(%+v, %v)", second, err)
|
||||
}
|
||||
|
||||
canceled, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if _, err := resolver.GetProfile(canceled, "leaf"); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("canceled lookup error = %v", err)
|
||||
}
|
||||
if got := repo.callCount("leaf"); got != 2 {
|
||||
t.Fatalf("calls after canceled lookup = %d, want 2", got)
|
||||
}
|
||||
|
||||
duringTraversal, cancelDuringTraversal := context.WithCancel(context.Background())
|
||||
repo.afterGet = func(id string) {
|
||||
if id == "leaf" {
|
||||
cancelDuringTraversal()
|
||||
}
|
||||
}
|
||||
if _, err := resolver.GetProfile(duringTraversal, "leaf"); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("during traversal error = %v", err)
|
||||
}
|
||||
if got := repo.callCount("base"); got != 2 {
|
||||
t.Fatalf("base calls after cancellation = %d, want 2", got)
|
||||
}
|
||||
|
||||
terminalLookup, cancelTerminalLookup := context.WithCancel(context.Background())
|
||||
repo.afterGet = func(id string) {
|
||||
if id == "base" {
|
||||
cancelTerminalLookup()
|
||||
}
|
||||
}
|
||||
if _, err := resolver.GetProfile(terminalLookup, "leaf"); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("terminal lookup cancellation error = %v", err)
|
||||
}
|
||||
if got := repo.callCount("base"); got != 3 {
|
||||
t.Fatalf("base calls after terminal cancellation = %d, want 3", got)
|
||||
}
|
||||
|
||||
repo.afterGet = nil
|
||||
var wg sync.WaitGroup
|
||||
errors := make(chan error, 8)
|
||||
for index := 0; index < cap(errors); index++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
resolved, err := resolver.GetProfile(context.Background(), "leaf")
|
||||
if err != nil {
|
||||
errors <- err
|
||||
return
|
||||
}
|
||||
resolved.ExtraParams["nested"].(map[string]any)["value"] = "mutated"
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
close(errors)
|
||||
for err := range errors {
|
||||
t.Errorf("concurrent resolution: %v", err)
|
||||
}
|
||||
latest, err := resolver.GetProfile(context.Background(), "leaf")
|
||||
if err != nil || latest.ExtraParams["nested"].(map[string]any)["value"] != "second" {
|
||||
t.Fatalf("latest result=(%+v, %v)", latest, err)
|
||||
}
|
||||
}
|
||||
|
||||
type resolvingTestRepository struct {
|
||||
mu sync.Mutex
|
||||
profiles map[string]*domain.ExecutionProfile
|
||||
errors map[string]error
|
||||
calls map[string]int
|
||||
afterGet func(string)
|
||||
}
|
||||
|
||||
func (r *resolvingTestRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.mu.Lock()
|
||||
if r.calls == nil {
|
||||
r.calls = make(map[string]int)
|
||||
}
|
||||
r.calls[id]++
|
||||
err := r.errors[id]
|
||||
profile := r.profiles[id]
|
||||
afterGet := r.afterGet
|
||||
r.mu.Unlock()
|
||||
if afterGet != nil {
|
||||
afterGet(id)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if profile == nil {
|
||||
if _, exists := r.profiles[id]; exists {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, ErrProfileNotFound
|
||||
}
|
||||
copy := *profile
|
||||
return ©, nil
|
||||
}
|
||||
|
||||
func (r *resolvingTestRepository) set(id string, profile *domain.ExecutionProfile) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.profiles[id] = profile
|
||||
}
|
||||
|
||||
func (r *resolvingTestRepository) callCount(id string) int {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return r.calls[id]
|
||||
}
|
||||
|
||||
func profileTestFS(files map[string]string) fs.FS {
|
||||
fsys := make(fstest.MapFS, len(files))
|
||||
for name, content := range files {
|
||||
fsys[name] = profileMapFile(content)
|
||||
}
|
||||
return fsys
|
||||
}
|
||||
@@ -68,8 +68,9 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if tmplMsg.Role == "" {
|
||||
return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i)
|
||||
role, err := domain.NormalizeMessageRole(tmplMsg.Role)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidMessageRole, i, err)
|
||||
}
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
@@ -96,7 +97,7 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
|
||||
}
|
||||
|
||||
renderedMessages = append(renderedMessages, domain.RenderedMessage{
|
||||
Role: tmplMsg.Role,
|
||||
Role: role,
|
||||
Content: buf.String(),
|
||||
CacheControl: cloneCacheControl(tmplMsg.CacheControl),
|
||||
})
|
||||
|
||||
@@ -361,6 +361,35 @@ func TestGoRenderer_Render(t *testing.T) {
|
||||
t.Fatalf("expected ErrInvalidMessageRole, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("canonicalizes directly supplied message roles", func(t *testing.T) {
|
||||
def := &domain.PromptDefinition{
|
||||
Templates: []domain.PromptMessageTemplate{
|
||||
{Role: " \u2003DeVeLoPeR\u2003 ", Content: "Hello"},
|
||||
},
|
||||
}
|
||||
res, err := renderer.Render(ctx, def, nil, vars)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if got := res.Messages[0].Role; got != domain.RoleDeveloper {
|
||||
t.Fatalf("rendered role = %q, want %q", got, domain.RoleDeveloper)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("rejects unsupported directly supplied message roles", func(t *testing.T) {
|
||||
const unsupportedRole = "consumer-private-role"
|
||||
def := &domain.PromptDefinition{
|
||||
Templates: []domain.PromptMessageTemplate{{Role: unsupportedRole, Content: "Hello"}},
|
||||
}
|
||||
_, err := renderer.Render(ctx, def, nil, vars)
|
||||
if !errors.Is(err, ErrInvalidMessageRole) {
|
||||
t.Fatalf("expected ErrInvalidMessageRole, got %v", err)
|
||||
}
|
||||
if strings.Contains(err.Error(), unsupportedRole) {
|
||||
t.Fatalf("renderer error exposed unsupported role: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestGoRendererCancellation(t *testing.T) {
|
||||
|
||||
@@ -254,20 +254,27 @@ func normalizePromptDefinitionWithContent(raw *promptDefinitionFile, readContent
|
||||
|
||||
templates := make([]domain.PromptMessageTemplate, 0, len(raw.Messages))
|
||||
for i, msg := range raw.Messages {
|
||||
role := strings.TrimSpace(msg.Role)
|
||||
if role == "" {
|
||||
return nil, fmt.Errorf("message %d role is required", i)
|
||||
role, err := domain.NormalizeMessageRole(msg.Role)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("message %d role: %w", i, err)
|
||||
}
|
||||
|
||||
hasContent := strings.TrimSpace(msg.Content) != ""
|
||||
hasContentFile := strings.TrimSpace(msg.ContentFile) != ""
|
||||
if hasContent == hasContentFile {
|
||||
return nil, fmt.Errorf("message %d (%s) must set exactly one of content or content_file", i, role)
|
||||
return nil, fmt.Errorf("message %d must set exactly one of content or content_file", i)
|
||||
}
|
||||
|
||||
cacheControl, err := normalizeCacheControl(msg.CacheControl)
|
||||
var rawCacheControl *domain.CacheControl
|
||||
if msg.CacheControl != nil {
|
||||
rawCacheControl = &domain.CacheControl{
|
||||
Type: domain.CacheControlType(msg.CacheControl.Type),
|
||||
TTL: msg.CacheControl.TTL,
|
||||
}
|
||||
}
|
||||
cacheControl, err := domain.NormalizeCacheControl(rawCacheControl)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("message %d (%s) cache_control: %w", i, role, err)
|
||||
return nil, fmt.Errorf("message %d cache_control: %w", i, err)
|
||||
}
|
||||
|
||||
templateContent := msg.Content
|
||||
@@ -275,7 +282,7 @@ func normalizePromptDefinitionWithContent(raw *promptDefinitionFile, readContent
|
||||
if hasContentFile {
|
||||
body, resolvedPath, err := readContentFile(msg.ContentFile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prompt %q message %d (%s): failed to read content_file %q: %w", id, i, role, msg.ContentFile, err)
|
||||
return nil, fmt.Errorf("prompt %q message %d: failed to read content_file %q: %w", id, i, msg.ContentFile, err)
|
||||
}
|
||||
templateContent = body
|
||||
resolvedContentFile = resolvedPath
|
||||
@@ -319,27 +326,3 @@ func normalizePromptDefinitionWithContent(raw *promptDefinitionFile, readContent
|
||||
Validation: outputContract,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeCacheControl(raw *cacheControlFile) (*domain.CacheControl, error) {
|
||||
if raw == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
cacheType := strings.TrimSpace(raw.Type)
|
||||
if cacheType == "" {
|
||||
return nil, errors.New("type is required")
|
||||
}
|
||||
if domain.CacheControlType(cacheType) != domain.CacheControlEphemeral {
|
||||
return nil, fmt.Errorf("unsupported type %q", cacheType)
|
||||
}
|
||||
|
||||
ttl := strings.TrimSpace(raw.TTL)
|
||||
if ttl != "" && ttl != "1h" {
|
||||
return nil, fmt.Errorf("unsupported ttl %q", ttl)
|
||||
}
|
||||
|
||||
return &domain.CacheControl{
|
||||
Type: domain.CacheControlType(cacheType),
|
||||
TTL: ttl,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -892,6 +892,38 @@ output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
repair_attempts: -1
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "repair_attempts",
|
||||
},
|
||||
{
|
||||
name: "repair attempts above maximum",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: basic
|
||||
repair_attempts: 4
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "repair_attempts",
|
||||
},
|
||||
{
|
||||
name: "none validation with repair attempts",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
repair_attempts: 1
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "repair_attempts",
|
||||
@@ -1055,6 +1087,48 @@ output:
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptDefinitionMessageRoleNormalization(t *testing.T) {
|
||||
const invalidRole = "consumer-private-role"
|
||||
repo := NewFSRepository(fstest.MapFS{
|
||||
"canonical.yaml": {Data: []byte(`
|
||||
id: canonical
|
||||
version: "1"
|
||||
messages:
|
||||
- role: " \u2003SyStEm\u2003 "
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
"invalid.yaml": {Data: []byte(`
|
||||
id: invalid-role
|
||||
version: "1"
|
||||
messages:
|
||||
- role: consumer-private-role
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
}, ".")
|
||||
|
||||
definition, err := repo.GetPromptDefinition(context.Background(), "canonical", "")
|
||||
if err != nil {
|
||||
t.Fatalf("GetPromptDefinition() error = %v", err)
|
||||
}
|
||||
if got := definition.Templates[0].Role; got != domain.RoleSystem {
|
||||
t.Fatalf("normalized role = %q, want %q", got, domain.RoleSystem)
|
||||
}
|
||||
|
||||
_, err = repo.GetPromptDefinition(context.Background(), "invalid-role", "")
|
||||
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
||||
t.Fatalf("expected ErrInvalidPromptDefinition, got %v", err)
|
||||
}
|
||||
if strings.Contains(err.Error(), invalidRole) {
|
||||
t.Fatalf("invalid prompt error exposed the supplied role: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCacheControl(t *testing.T, got *domain.CacheControl, wantType domain.CacheControlType, wantTTL string) {
|
||||
t.Helper()
|
||||
if got == nil {
|
||||
|
||||
@@ -106,6 +106,8 @@ func TestRunnerPreparationRejectsInvalidOutputContractsBeforeCompletion(t *testi
|
||||
{name: "unsupported format", override: domain.OutputContract{Format: "binary", ValidationMode: domain.ValidationNone}},
|
||||
{name: "empty validation mode", override: domain.OutputContract{Format: domain.FormatText}},
|
||||
{name: "negative repair attempts", override: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone, RepairAttempts: -1}},
|
||||
{name: "repair attempts above maximum", override: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationBasic, RepairAttempts: 4}},
|
||||
{name: "none validation with repair attempts", override: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone, RepairAttempts: 1}},
|
||||
{name: "json schema without path", override: domain.OutputContract{Format: domain.FormatJSON, ValidationMode: domain.ValidationJSONSchema}},
|
||||
}
|
||||
|
||||
|
||||
@@ -213,16 +213,7 @@ func clonePreparedRun(source *domain.PreparedRun) (*domain.PreparedRun, error) {
|
||||
copied.InputHashes[name] = hash
|
||||
}
|
||||
}
|
||||
if source.Messages != nil {
|
||||
copied.Messages = make([]domain.RenderedMessage, len(source.Messages))
|
||||
for i, message := range source.Messages {
|
||||
copied.Messages[i] = message
|
||||
if message.CacheControl != nil {
|
||||
cacheControl := *message.CacheControl
|
||||
copied.Messages[i].CacheControl = &cacheControl
|
||||
}
|
||||
}
|
||||
}
|
||||
copied.Messages = domain.CloneRenderedMessages(source.Messages)
|
||||
|
||||
if source.StructuredOutput != nil {
|
||||
structuredOutput := *source.StructuredOutput
|
||||
|
||||
@@ -3,11 +3,12 @@ package usecase
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/validate"
|
||||
)
|
||||
|
||||
@@ -241,49 +242,76 @@ func excessivelyDeepPreparedJSONValue() any {
|
||||
return value
|
||||
}
|
||||
|
||||
func TestRunnerRunPreparedRechecksEnvironmentCredentialBeforeAdmission(t *testing.T) {
|
||||
func TestRunnerRunPreparedCredentialAvailabilityBeforeAdmission(t *testing.T) {
|
||||
const environmentName = "PROMPTKIT_PREPARED_EXECUTION_TEST_KEY"
|
||||
t.Setenv(environmentName, "available-during-preparation")
|
||||
|
||||
profile := defaultExecutionProfile()
|
||||
profile.APIKeyEnv = environmentName
|
||||
validator := &recordingValidationPreparer{plan: &recordingPreparedValidation{}}
|
||||
admitter := &fakeRunAdmitter{}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "unexpected"}}
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": profile}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
validator,
|
||||
admitter,
|
||||
)
|
||||
prepared, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
if err := os.Unsetenv(environmentName); err != nil {
|
||||
t.Fatalf("unset credential environment: %v", err)
|
||||
tests := []struct {
|
||||
name string
|
||||
apiKeyRequired bool
|
||||
profileEnv bool
|
||||
overrideEnv bool
|
||||
wantFailure bool
|
||||
}{
|
||||
{name: "optional environment becomes unavailable", profileEnv: true},
|
||||
{
|
||||
name: "required request environment becomes unavailable",
|
||||
apiKeyRequired: true,
|
||||
overrideEnv: true,
|
||||
wantFailure: true,
|
||||
},
|
||||
}
|
||||
|
||||
result, err := runner.RunPrepared(context.Background(), prepared)
|
||||
if result != nil {
|
||||
t.Fatalf("credential failure returned partial result: %+v", result)
|
||||
}
|
||||
if !errors.Is(err, ErrInvalidRequest) || !errors.Is(err, ErrAPIKeyEnvMissing) {
|
||||
t.Fatalf("credential error identities are missing: %v", err)
|
||||
}
|
||||
if len(admitter.backendIDs) != 0 || llmClient.calls != 0 {
|
||||
t.Fatalf("credential failure reached admission or generation: admission=%v generation=%d", admitter.backendIDs, llmClient.calls)
|
||||
}
|
||||
if _, err := runner.RunPrepared(context.Background(), prepared); !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("credential failure did not consume execution: %v", err)
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Setenv(environmentName, "available-during-preparation")
|
||||
|
||||
profile := defaultExecutionProfile()
|
||||
profile.APIKeyRequired = tc.apiKeyRequired
|
||||
if tc.profileEnv {
|
||||
profile.APIKeyEnv = environmentName
|
||||
}
|
||||
validator := &recordingValidationPreparer{plan: &recordingPreparedValidation{}}
|
||||
admitter := &fakeRunAdmitter{}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": profile}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
validator,
|
||||
admitter,
|
||||
)
|
||||
request := domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}
|
||||
if tc.overrideEnv {
|
||||
request.Execution = &domain.ExecutionTargetOverride{APIKeyEnv: environmentName}
|
||||
}
|
||||
prepared, err := runner.PrepareExecution(context.Background(), request)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
t.Setenv(environmentName, "")
|
||||
|
||||
result, err := runner.RunPrepared(context.Background(), prepared)
|
||||
if tc.wantFailure {
|
||||
if result != nil || !errors.Is(err, ErrInvalidRequest) || !errors.Is(err, ErrAPIKeyEnvMissing) {
|
||||
t.Fatalf("required credential result = (%+v, %v)", result, err)
|
||||
}
|
||||
if len(admitter.backendIDs) != 0 || llmClient.calls != 0 {
|
||||
t.Fatalf("required credential reached admission or generation: admission=%v generation=%d", admitter.backendIDs, llmClient.calls)
|
||||
}
|
||||
} else {
|
||||
if result == nil || err != nil {
|
||||
t.Fatalf("optional credential result = (%+v, %v), want success", result, err)
|
||||
}
|
||||
if len(admitter.backendIDs) != 1 || llmClient.calls != 1 {
|
||||
t.Fatalf("optional credential admission=%v generation=%d, want one each", admitter.backendIDs, llmClient.calls)
|
||||
}
|
||||
}
|
||||
if _, err := runner.RunPrepared(context.Background(), prepared); !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("execution outcome did not consume handle: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -359,19 +387,20 @@ func TestRunnerRunPreparedUsesFrozenValidationForInitialAndRepairOutputs(t *test
|
||||
admitter := &fakeRunAdmitter{}
|
||||
reader := defaultArtifactReader()
|
||||
renderer := defaultRenderer()
|
||||
client := &sequenceLLM{responses: []*domain.GenerateResponse{{
|
||||
Content: `{"broken":true}`,
|
||||
Usage: domain.TokenUsage{
|
||||
PromptTokens: 13, CompletionTokens: 17, TotalTokens: 19,
|
||||
CachedTokens: 23, CacheWriteTokens: 29,
|
||||
},
|
||||
}}}
|
||||
runner := NewRunnerWithRepairer(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 1)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
reader,
|
||||
renderer,
|
||||
&fakeLLM{resp: &domain.GenerateResponse{
|
||||
Content: `{"broken":true}`,
|
||||
Usage: domain.TokenUsage{
|
||||
PromptTokens: 13, CompletionTokens: 17, TotalTokens: 19,
|
||||
CachedTokens: 23, CacheWriteTokens: 29,
|
||||
},
|
||||
}},
|
||||
client,
|
||||
validator,
|
||||
repairer,
|
||||
admitter,
|
||||
@@ -395,6 +424,10 @@ func TestRunnerRunPreparedUsesFrozenValidationForInitialAndRepairOutputs(t *test
|
||||
if validator.directValidateCalls != 0 || repairer.calls != 1 {
|
||||
t.Fatalf("validation/repair calls=(direct=%d repair=%d), want (0, 1)", validator.directValidateCalls, repairer.calls)
|
||||
}
|
||||
if len(client.requests) != 1 || len(repairer.reqs) != 1 ||
|
||||
!reflect.DeepEqual(repairer.reqs[0].OriginalMessages, client.requests[0].Prompt.Messages) {
|
||||
t.Fatalf("initial and repair messages = (%#v, %#v)", client.requests, repairer.reqs)
|
||||
}
|
||||
if result.Validation.Status != domain.ValidationPassed || result.Validation.RepairAttempts != 1 {
|
||||
t.Fatalf("unexpected repaired validation result: %+v", result.Validation)
|
||||
}
|
||||
@@ -417,6 +450,7 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
generationFailure := errors.New("generation failed")
|
||||
validationFailure := errors.New("validation failed")
|
||||
repairFailure := errors.New("repair failed")
|
||||
repairInvalidRequest := fmt.Errorf("repair request: %w", llm.ErrInvalidRequest)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -424,6 +458,7 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
validation *recordingPreparedValidation
|
||||
repairer *fakeRepairer
|
||||
wantError error
|
||||
wantSource error
|
||||
}{
|
||||
{
|
||||
name: "generation failure",
|
||||
@@ -448,8 +483,51 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: repairFailure},
|
||||
wantError: ErrValidation,
|
||||
repairer: &fakeRepairer{err: repairFailure},
|
||||
wantError: ErrLLMGenerate,
|
||||
wantSource: repairFailure,
|
||||
},
|
||||
{
|
||||
name: "repair invalid request",
|
||||
validation: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: repairInvalidRequest},
|
||||
wantError: ErrInvalidRequest,
|
||||
wantSource: repairInvalidRequest,
|
||||
},
|
||||
{
|
||||
name: "repair cancellation",
|
||||
validation: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: context.Canceled},
|
||||
wantError: ErrLLMGenerate,
|
||||
wantSource: context.Canceled,
|
||||
},
|
||||
{
|
||||
name: "repair deadline",
|
||||
validation: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: context.DeadlineExceeded},
|
||||
wantError: ErrLLMGenerate,
|
||||
wantSource: context.DeadlineExceeded,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -488,6 +566,9 @@ func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
if result != nil || !errors.Is(err, test.wantError) {
|
||||
t.Fatalf("run prepared=(%+v, %v), want %v", result, err, test.wantError)
|
||||
}
|
||||
if test.wantSource != nil && !errors.Is(err, test.wantSource) {
|
||||
t.Fatalf("run prepared error = %v, want source %v", err, test.wantSource)
|
||||
}
|
||||
if len(admitter.backendIDs) != 1 || admitter.releaseCalls != 1 {
|
||||
t.Fatalf(
|
||||
"admission calls=%#v releases=%d, want one each",
|
||||
|
||||
@@ -2,19 +2,29 @@ package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||
)
|
||||
|
||||
const (
|
||||
maxRepairDiagnosticBytes = 64 * 1024
|
||||
omittedRepairDiagnostics = "additional validation diagnostics were omitted"
|
||||
)
|
||||
|
||||
// OutputRepairer generates a corrected candidate after validation fails.
|
||||
type OutputRepairer interface {
|
||||
Repair(ctx context.Context, req RepairRequest) (*domain.GenerateResponse, error)
|
||||
}
|
||||
|
||||
// RepairRequest contains the immutable execution state needed for one correction.
|
||||
type RepairRequest struct {
|
||||
OriginalMessages []domain.RenderedMessage
|
||||
PreviousOutput string
|
||||
ValidationErrors []string
|
||||
SessionID string
|
||||
@@ -30,6 +40,7 @@ type defaultOutputRepairer struct {
|
||||
llm llm.Client
|
||||
}
|
||||
|
||||
// NewDefaultOutputRepairer constructs the standard internal output repairer.
|
||||
func NewDefaultOutputRepairer(llmClient llm.Client) OutputRepairer {
|
||||
return &defaultOutputRepairer{llm: llmClient}
|
||||
}
|
||||
@@ -39,33 +50,43 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
|
||||
return nil, errors.New("llm client is required for repair")
|
||||
}
|
||||
|
||||
errs := "(none provided)"
|
||||
if len(req.ValidationErrors) > 0 {
|
||||
errs = strings.Join(req.ValidationErrors, "\n")
|
||||
guidance, err := repairGuidance(req.Mode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
prompt := domain.RenderedPrompt{
|
||||
Messages: []domain.RenderedMessage{
|
||||
{
|
||||
Role: "system",
|
||||
Content: "You repair invalid JSON output. Return only corrected JSON. Do not include explanations or markdown code fences.",
|
||||
},
|
||||
{
|
||||
Role: "user",
|
||||
Content: fmt.Sprintf(
|
||||
"Repair attempt %d of %d for validation mode %s.\n\nValidation errors:\n%s\n\nPrevious output:\n%s\n\nReturn only corrected JSON.",
|
||||
req.Attempt,
|
||||
req.MaxAttempts,
|
||||
req.Mode,
|
||||
errs,
|
||||
req.PreviousOutput,
|
||||
),
|
||||
},
|
||||
},
|
||||
hasPreviousOutput := strings.TrimSpace(req.PreviousOutput) != ""
|
||||
suffix := make([]domain.RenderedMessage, 0, 2)
|
||||
if hasPreviousOutput {
|
||||
suffix = append(suffix, domain.RenderedMessage{
|
||||
Role: domain.RoleAssistant,
|
||||
Content: req.PreviousOutput,
|
||||
})
|
||||
}
|
||||
|
||||
previousResponse := "The previous response was empty."
|
||||
if hasPreviousOutput {
|
||||
previousResponse = "The previous response is included immediately before this instruction."
|
||||
}
|
||||
suffix = append(suffix, domain.RenderedMessage{
|
||||
Role: domain.RoleUser,
|
||||
Content: fmt.Sprintf(
|
||||
"Repair attempt %d of %d for validation mode %s.\n"+
|
||||
"Preserve valid values and change only what is necessary.\n"+
|
||||
"%s\n%s\n"+
|
||||
"Validation diagnostics (data):\n%s",
|
||||
req.Attempt,
|
||||
req.MaxAttempts,
|
||||
req.Mode,
|
||||
previousResponse,
|
||||
guidance,
|
||||
formatRepairDiagnostics(req.ValidationErrors),
|
||||
),
|
||||
})
|
||||
messages := domain.ConcatRenderedMessages(req.OriginalMessages, suffix)
|
||||
|
||||
resp, err := r.llm.Generate(ctx, newGenerationRequest(
|
||||
prompt,
|
||||
domain.RenderedPrompt{Messages: messages},
|
||||
req.SessionID,
|
||||
req.Target,
|
||||
req.TargetPresence,
|
||||
@@ -80,3 +101,90 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func repairGuidance(mode domain.ValidationMode) (string, error) {
|
||||
switch mode {
|
||||
case domain.ValidationBasic:
|
||||
return "Return a nonempty response satisfying the original request.", nil
|
||||
case domain.ValidationJSON, domain.ValidationJSONSchema:
|
||||
return "Return only corrected JSON, with no explanation or Markdown fences.", nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported validation mode for repair: %q", mode)
|
||||
}
|
||||
}
|
||||
|
||||
func formatRepairDiagnostics(errors []string) string {
|
||||
diagnostics := make([]string, len(errors))
|
||||
for index, diagnostic := range errors {
|
||||
diagnostics[index] = strings.ToValidUTF8(diagnostic, "\uFFFD")
|
||||
}
|
||||
|
||||
complete, _ := json.Marshal(diagnostics)
|
||||
if len(complete) <= maxRepairDiagnosticBytes {
|
||||
return string(complete)
|
||||
}
|
||||
|
||||
omission, _ := json.Marshal(omittedRepairDiagnostics)
|
||||
encoded := make([]byte, 0, maxRepairDiagnosticBytes)
|
||||
encoded = append(encoded, '[')
|
||||
for _, diagnostic := range diagnostics {
|
||||
entry, _ := json.Marshal(diagnostic)
|
||||
separator := 0
|
||||
if len(encoded) > 1 {
|
||||
separator = 1
|
||||
}
|
||||
available := maxRepairDiagnosticBytes - len(encoded) - separator - 1 - len(omission) - 1
|
||||
if len(entry) <= available {
|
||||
if separator != 0 {
|
||||
encoded = append(encoded, ',')
|
||||
}
|
||||
encoded = append(encoded, entry...)
|
||||
continue
|
||||
}
|
||||
|
||||
if available < len(`""`) {
|
||||
break
|
||||
}
|
||||
if separator != 0 {
|
||||
encoded = append(encoded, ',')
|
||||
}
|
||||
encoded = append(encoded, truncateDiagnosticJSONValue(diagnostic, available)...)
|
||||
break
|
||||
}
|
||||
if len(encoded) > 1 {
|
||||
encoded = append(encoded, ',')
|
||||
}
|
||||
encoded = append(encoded, omission...)
|
||||
encoded = append(encoded, ']')
|
||||
return string(encoded)
|
||||
}
|
||||
|
||||
func truncateDiagnosticJSONValue(value string, maxBytes int) []byte {
|
||||
if maxBytes < len(`""`) {
|
||||
return nil
|
||||
}
|
||||
|
||||
boundaries := []int{0}
|
||||
for end := 0; end < len(value); {
|
||||
_, size := utf8.DecodeRuneInString(value[end:])
|
||||
end += size
|
||||
if end+len(`""`) > maxBytes {
|
||||
break
|
||||
}
|
||||
boundaries = append(boundaries, end)
|
||||
}
|
||||
|
||||
low, high := 0, len(boundaries)-1
|
||||
best := []byte(`""`)
|
||||
for low <= high {
|
||||
mid := low + (high-low)/2
|
||||
candidate, _ := json.Marshal(value[:boundaries[mid]])
|
||||
if len(candidate) <= maxBytes {
|
||||
best = candidate
|
||||
low = mid + 1
|
||||
continue
|
||||
}
|
||||
high = mid - 1
|
||||
}
|
||||
return best
|
||||
}
|
||||
|
||||
351
internal/usecase/repairer_test.go
Normal file
351
internal/usecase/repairer_test.go
Normal file
@@ -0,0 +1,351 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode/utf8"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
type recordingRepairClient struct {
|
||||
requests []domain.GenerateRequest
|
||||
response *domain.GenerateResponse
|
||||
err error
|
||||
}
|
||||
|
||||
func (c *recordingRepairClient) Generate(_ context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) {
|
||||
c.requests = append(c.requests, req)
|
||||
return c.response, c.err
|
||||
}
|
||||
|
||||
func TestDefaultOutputRepairerBuildsFullContextRequest(t *testing.T) {
|
||||
client := &recordingRepairClient{response: &domain.GenerateResponse{Content: "corrected"}}
|
||||
repairer := NewDefaultOutputRepairer(client)
|
||||
original := []domain.RenderedMessage{
|
||||
{Role: "system", Content: "Follow the task.", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
|
||||
{Role: "user", Content: "Summarize the report."},
|
||||
{Role: "assistant", Content: "Consumer-supplied previous response."},
|
||||
{Role: "user", Content: "Consumer-supplied correction."},
|
||||
}
|
||||
before := append([]domain.RenderedMessage(nil), original...)
|
||||
previous := strings.Repeat("candidate ", 12_000)
|
||||
target := domain.ExecutionTarget{BackendID: "backend", Endpoint: "https://provider.example/v1", Model: "model"}
|
||||
presence := domain.ExecutionTargetPresence{Temperature: true, TopP: true}
|
||||
structured := &domain.StructuredOutputSpec{}
|
||||
|
||||
response, err := repairer.Repair(context.Background(), RepairRequest{
|
||||
OriginalMessages: original,
|
||||
PreviousOutput: previous,
|
||||
ValidationErrors: []string{"invalid JSON"},
|
||||
SessionID: "session",
|
||||
Target: target,
|
||||
TargetPresence: presence,
|
||||
StructuredOutput: structured,
|
||||
Attempt: 1,
|
||||
MaxAttempts: 3,
|
||||
Mode: domain.ValidationJSONSchema,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("repair: %v", err)
|
||||
}
|
||||
if response == nil || response.Content != "corrected" {
|
||||
t.Fatalf("response = %+v", response)
|
||||
}
|
||||
if !reflect.DeepEqual(original, before) {
|
||||
t.Fatalf("original messages changed: got %#v, want %#v", original, before)
|
||||
}
|
||||
if len(client.requests) != 1 {
|
||||
t.Fatalf("generation requests = %d, want 1", len(client.requests))
|
||||
}
|
||||
|
||||
request := client.requests[0]
|
||||
if request.Prompt.SessionID != "session" || !reflect.DeepEqual(request.Target, target) || request.TargetPresence != presence || request.StructuredOutput != structured {
|
||||
t.Fatalf("generation request fields = %+v", request)
|
||||
}
|
||||
if len(request.Prompt.Messages) != len(original)+2 {
|
||||
t.Fatalf("message count = %d, want %d", len(request.Prompt.Messages), len(original)+2)
|
||||
}
|
||||
if !reflect.DeepEqual(request.Prompt.Messages[:len(original)], original) {
|
||||
t.Fatalf("original messages = %#v, want %#v", request.Prompt.Messages[:len(original)], original)
|
||||
}
|
||||
assistant := request.Prompt.Messages[len(original)]
|
||||
if assistant.Role != "assistant" || assistant.Content != previous {
|
||||
t.Fatalf("assistant candidate = %+v", assistant)
|
||||
}
|
||||
correction := request.Prompt.Messages[len(original)+1]
|
||||
if correction.Role != "user" || !strings.Contains(correction.Content, "Repair attempt 1 of 3") ||
|
||||
!strings.Contains(correction.Content, "Preserve valid values") ||
|
||||
!strings.Contains(correction.Content, "Return only corrected JSON") {
|
||||
t.Fatalf("correction message = %q", correction.Content)
|
||||
}
|
||||
if diagnostics := repairDiagnosticsFromMessage(t, correction.Content); !reflect.DeepEqual(diagnostics, []string{"invalid JSON"}) {
|
||||
t.Fatalf("diagnostics = %#v", diagnostics)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultOutputRepairerDoesNotAccumulateCandidates(t *testing.T) {
|
||||
client := &recordingRepairClient{response: &domain.GenerateResponse{Content: "corrected"}}
|
||||
repairer := NewDefaultOutputRepairer(client)
|
||||
backing := make([]domain.RenderedMessage, 1, 4)
|
||||
backing[0] = domain.RenderedMessage{Role: "user", Content: "Original task"}
|
||||
before := append([]domain.RenderedMessage(nil), backing...)
|
||||
|
||||
for _, candidate := range []string{"first invalid", "second invalid"} {
|
||||
_, err := repairer.Repair(context.Background(), RepairRequest{
|
||||
OriginalMessages: backing,
|
||||
PreviousOutput: candidate,
|
||||
Attempt: 1,
|
||||
MaxAttempts: 3,
|
||||
Mode: domain.ValidationJSON,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("repair %q: %v", candidate, err)
|
||||
}
|
||||
}
|
||||
if !reflect.DeepEqual(backing, before) {
|
||||
t.Fatalf("caller messages changed: got %#v, want %#v", backing, before)
|
||||
}
|
||||
if len(client.requests) != 2 {
|
||||
t.Fatalf("generation requests = %d, want 2", len(client.requests))
|
||||
}
|
||||
for index, request := range client.requests {
|
||||
messages := request.Prompt.Messages
|
||||
if len(messages) != 3 || messages[0] != backing[0] || messages[1].Role != "assistant" || messages[1].Content != []string{"first invalid", "second invalid"}[index] {
|
||||
t.Fatalf("request %d messages = %#v", index, messages)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultOutputRepairerOmitsEmptyCandidateMessage(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
candidate string
|
||||
}{
|
||||
{name: "empty", candidate: ""},
|
||||
{name: "whitespace", candidate: " \n\t "},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
client := &recordingRepairClient{response: &domain.GenerateResponse{Content: "corrected"}}
|
||||
repairer := NewDefaultOutputRepairer(client)
|
||||
_, err := repairer.Repair(context.Background(), RepairRequest{
|
||||
OriginalMessages: []domain.RenderedMessage{{Role: "user", Content: "Original task"}},
|
||||
PreviousOutput: tc.candidate,
|
||||
Attempt: 1,
|
||||
MaxAttempts: 1,
|
||||
Mode: domain.ValidationBasic,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("repair: %v", err)
|
||||
}
|
||||
messages := client.requests[0].Prompt.Messages
|
||||
if len(messages) != 2 || messages[1].Role != "user" || !strings.Contains(messages[1].Content, "previous response was empty") {
|
||||
t.Fatalf("messages = %#v", messages)
|
||||
}
|
||||
if strings.Contains(messages[1].Content, tc.candidate) && tc.candidate != "" {
|
||||
t.Fatalf("correction message repeated whitespace candidate: %q", messages[1].Content)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultOutputRepairerUsesModeSpecificGuidance(t *testing.T) {
|
||||
tests := []struct {
|
||||
mode domain.ValidationMode
|
||||
want string
|
||||
}{
|
||||
{mode: domain.ValidationBasic, want: "Return a nonempty response"},
|
||||
{mode: domain.ValidationJSON, want: "Return only corrected JSON"},
|
||||
{mode: domain.ValidationJSONSchema, want: "Return only corrected JSON"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(string(tc.mode), func(t *testing.T) {
|
||||
client := &recordingRepairClient{response: &domain.GenerateResponse{Content: "corrected"}}
|
||||
_, err := NewDefaultOutputRepairer(client).Repair(context.Background(), RepairRequest{Mode: tc.mode, Attempt: 1, MaxAttempts: 1})
|
||||
if err != nil {
|
||||
t.Fatalf("repair: %v", err)
|
||||
}
|
||||
if !strings.Contains(client.requests[0].Prompt.Messages[0].Content, tc.want) {
|
||||
t.Fatalf("correction message = %q", client.requests[0].Prompt.Messages[0].Content)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
client := &recordingRepairClient{response: &domain.GenerateResponse{Content: "unexpected"}}
|
||||
_, err := NewDefaultOutputRepairer(client).Repair(context.Background(), RepairRequest{Mode: domain.ValidationNone})
|
||||
if err == nil || !strings.Contains(err.Error(), "unsupported validation mode") {
|
||||
t.Fatalf("unsupported mode error = %v", err)
|
||||
}
|
||||
if len(client.requests) != 0 {
|
||||
t.Fatalf("generation requests = %d, want 0", len(client.requests))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatRepairDiagnosticsBoundsAndPreservesData(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
errors []string
|
||||
wantOmission bool
|
||||
check func(*testing.T, []string)
|
||||
}{
|
||||
{
|
||||
name: "below limit",
|
||||
errors: []string{"first", "second"},
|
||||
check: func(t *testing.T, got []string) {
|
||||
t.Helper()
|
||||
if !reflect.DeepEqual(got, []string{"first", "second"}) {
|
||||
t.Fatalf("diagnostics = %#v", got)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "at limit",
|
||||
errors: []string{strings.Repeat("x", maxRepairDiagnosticBytes-4)},
|
||||
check: func(t *testing.T, got []string) {
|
||||
t.Helper()
|
||||
if len(got) != 1 || len(got[0]) != maxRepairDiagnosticBytes-4 {
|
||||
t.Fatalf("diagnostics lengths = %#v", got)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "multibyte truncation preserves prior entries",
|
||||
errors: []string{"first", strings.Repeat("界", maxRepairDiagnosticBytes)},
|
||||
wantOmission: true,
|
||||
check: func(t *testing.T, got []string) {
|
||||
t.Helper()
|
||||
if len(got) != 3 || got[0] != "first" || !strings.HasPrefix(strings.Repeat("界", maxRepairDiagnosticBytes), got[1]) {
|
||||
t.Fatalf("diagnostics = %#v", got)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "many diagnostics",
|
||||
errors: manyRepairDiagnostics(),
|
||||
wantOmission: true,
|
||||
check: func(t *testing.T, got []string) {
|
||||
t.Helper()
|
||||
if len(got) < 2 || got[0] != manyRepairDiagnostics()[0] {
|
||||
t.Fatalf("diagnostics = %#v", got)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "no room for partial diagnostic",
|
||||
errors: func() []string {
|
||||
omission, _ := json.Marshal(omittedRepairDiagnostics)
|
||||
return []string{
|
||||
strings.Repeat("x", maxRepairDiagnosticBytes-len(omission)-5),
|
||||
strings.Repeat("y", 128),
|
||||
}
|
||||
}(),
|
||||
wantOmission: true,
|
||||
check: func(t *testing.T, got []string) {
|
||||
t.Helper()
|
||||
if len(got) != 2 || got[1] != omittedRepairDiagnostics {
|
||||
t.Fatalf("diagnostics = %#v", got)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid UTF-8",
|
||||
errors: []string{"broken\xffinput"},
|
||||
check: func(t *testing.T, got []string) {
|
||||
t.Helper()
|
||||
if !reflect.DeepEqual(got, []string{"broken\uFFFDinput"}) {
|
||||
t.Fatalf("diagnostics = %#v", got)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "one huge diagnostic",
|
||||
errors: []string{strings.Repeat("x", maxRepairDiagnosticBytes*2)},
|
||||
wantOmission: true,
|
||||
check: func(t *testing.T, got []string) {
|
||||
t.Helper()
|
||||
if len(got) != 2 || !strings.HasPrefix(strings.Repeat("x", maxRepairDiagnosticBytes*2), got[0]) {
|
||||
t.Fatalf("diagnostics = %#v", got)
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
before := append([]string(nil), tc.errors...)
|
||||
encoded := formatRepairDiagnostics(tc.errors)
|
||||
if len(encoded) > maxRepairDiagnosticBytes || !utf8.ValidString(encoded) {
|
||||
t.Fatalf("encoded diagnostic length/UTF-8 = (%d, %v)", len(encoded), utf8.ValidString(encoded))
|
||||
}
|
||||
var got []string
|
||||
if err := json.Unmarshal([]byte(encoded), &got); err != nil {
|
||||
t.Fatalf("decode diagnostics: %v; encoded=%q", err, encoded)
|
||||
}
|
||||
if !reflect.DeepEqual(tc.errors, before) {
|
||||
t.Fatalf("input diagnostics changed: got %#v, want %#v", tc.errors, before)
|
||||
}
|
||||
if hasOmission := len(got) > 0 && got[len(got)-1] == omittedRepairDiagnostics; hasOmission != tc.wantOmission {
|
||||
t.Fatalf("omission = %v, want %v; diagnostics=%#v", hasOmission, tc.wantOmission, got)
|
||||
}
|
||||
tc.check(t, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func manyRepairDiagnostics() []string {
|
||||
diagnostics := make([]string, 1_000)
|
||||
for index := range diagnostics {
|
||||
diagnostics[index] = fmt.Sprintf("diagnostic %04d %s", index, strings.Repeat("x", 128))
|
||||
}
|
||||
return diagnostics
|
||||
}
|
||||
|
||||
func TestDefaultOutputRepairerPropagatesGenerationFailures(t *testing.T) {
|
||||
expected := errors.New("generation failed")
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
client *recordingRepairClient
|
||||
want error
|
||||
}{
|
||||
{name: "nil client", want: nil},
|
||||
{name: "generation error", client: &recordingRepairClient{err: expected}, want: expected},
|
||||
{name: "nil response", client: &recordingRepairClient{}, want: nil},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
var repairer OutputRepairer
|
||||
if tc.client == nil {
|
||||
repairer = NewDefaultOutputRepairer(nil)
|
||||
} else {
|
||||
repairer = NewDefaultOutputRepairer(tc.client)
|
||||
}
|
||||
response, err := repairer.Repair(context.Background(), RepairRequest{Mode: domain.ValidationJSON})
|
||||
if response != nil || err == nil {
|
||||
t.Fatalf("response/error = (%+v, %v)", response, err)
|
||||
}
|
||||
if tc.want != nil && !errors.Is(err, tc.want) {
|
||||
t.Fatalf("error = %v, want %v", err, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func repairDiagnosticsFromMessage(t *testing.T, message string) []string {
|
||||
t.Helper()
|
||||
const marker = "Validation diagnostics (data):\n"
|
||||
index := strings.Index(message, marker)
|
||||
if index < 0 {
|
||||
t.Fatalf("missing diagnostics marker in %q", message)
|
||||
}
|
||||
encoded := message[index+len(marker):]
|
||||
var diagnostics []string
|
||||
if err := json.Unmarshal([]byte(encoded), &diagnostics); err != nil {
|
||||
t.Fatalf("decode diagnostics: %v", err)
|
||||
}
|
||||
return diagnostics
|
||||
}
|
||||
@@ -4,10 +4,13 @@ import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"hash"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -184,10 +187,10 @@ func (r *Runner) executePreparedRun(
|
||||
prepared.StructuredOutput,
|
||||
))
|
||||
if err != nil {
|
||||
if errors.Is(err, llm.ErrInvalidRequest) {
|
||||
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
|
||||
return nil, wrapGenerationError(err)
|
||||
}
|
||||
if genResp == nil {
|
||||
return nil, fmt.Errorf("%w: model returned nil response", ErrLLMGenerate)
|
||||
}
|
||||
usage := genResp.Usage
|
||||
|
||||
@@ -203,6 +206,7 @@ func (r *Runner) executePreparedRun(
|
||||
attemptsUsed++
|
||||
|
||||
repairResp, repairErr := r.repairer.Repair(ctx, RepairRequest{
|
||||
OriginalMessages: prepared.Messages,
|
||||
PreviousOutput: genResp.Content,
|
||||
ValidationErrors: validationResult.Errors,
|
||||
SessionID: prepared.SessionID,
|
||||
@@ -214,10 +218,10 @@ func (r *Runner) executePreparedRun(
|
||||
Mode: prepared.OutputContract.ValidationMode,
|
||||
})
|
||||
if repairErr != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrValidation, repairErr)
|
||||
return nil, wrapGenerationError(repairErr)
|
||||
}
|
||||
if repairResp == nil {
|
||||
return nil, fmt.Errorf("%w: repairer returned nil response", ErrValidation)
|
||||
return nil, fmt.Errorf("%w: repairer returned nil response", ErrLLMGenerate)
|
||||
}
|
||||
|
||||
genResp = repairResp
|
||||
@@ -257,6 +261,13 @@ func (r *Runner) executePreparedRun(
|
||||
}, nil
|
||||
}
|
||||
|
||||
func wrapGenerationError(err error) error {
|
||||
if errors.Is(err, llm.ErrInvalidRequest) {
|
||||
return fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||
}
|
||||
return fmt.Errorf("%w: %w", ErrLLMGenerate, err)
|
||||
}
|
||||
|
||||
func addTokenUsage(total, next domain.TokenUsage) domain.TokenUsage {
|
||||
return domain.TokenUsage{
|
||||
PromptTokens: total.PromptTokens + next.PromptTokens,
|
||||
@@ -436,9 +447,11 @@ func (r *Runner) completePreparationWithStructuredOutput(
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
|
||||
}
|
||||
effectivePrompt := *renderedPrompt
|
||||
if state.directSessionID != "" {
|
||||
renderedPrompt.SessionID = state.directSessionID
|
||||
effectivePrompt.SessionID = state.directSessionID
|
||||
}
|
||||
effectivePrompt.Messages = domain.ConcatRenderedMessages(renderedPrompt.Messages, req.AppendedMessages)
|
||||
|
||||
end := time.Now().UTC()
|
||||
effectiveModel := state.effectiveModel
|
||||
@@ -454,9 +467,9 @@ func (r *Runner) completePreparationWithStructuredOutput(
|
||||
OutputContract: state.effectiveContract,
|
||||
StructuredOutput: structuredOutput,
|
||||
InputHashes: inputHashes,
|
||||
SessionID: renderedPrompt.SessionID,
|
||||
RenderedPromptHash: hashRenderedPrompt(*renderedPrompt),
|
||||
Messages: renderedPrompt.Messages,
|
||||
SessionID: effectivePrompt.SessionID,
|
||||
RenderedPromptHash: hashRenderedPrompt(effectivePrompt),
|
||||
Messages: effectivePrompt.Messages,
|
||||
StartTime: state.start,
|
||||
EndTime: end,
|
||||
DurationMS: end.Sub(state.start).Milliseconds(),
|
||||
@@ -524,7 +537,12 @@ func (r *Runner) shouldAttemptRepair(contract domain.OutputContract, validationR
|
||||
if validationResult.Status != domain.ValidationFailed {
|
||||
return false
|
||||
}
|
||||
return contract.ValidationMode == domain.ValidationJSON || contract.ValidationMode == domain.ValidationJSONSchema
|
||||
switch contract.ValidationMode {
|
||||
case domain.ValidationBasic, domain.ValidationJSON, domain.ValidationJSONSchema:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func mergeExecutionTarget(base domain.ExecutionTarget, override domain.ExecutionTarget) domain.ExecutionTarget {
|
||||
@@ -624,12 +642,12 @@ func validateAPIKey(apiKeyEnv string, apiKey string, apiKeyRequired bool) error
|
||||
if strings.TrimSpace(apiKey) != "" {
|
||||
return nil
|
||||
}
|
||||
if !apiKeyRequired {
|
||||
return nil
|
||||
}
|
||||
envName := strings.TrimSpace(apiKeyEnv)
|
||||
if envName == "" {
|
||||
if apiKeyRequired {
|
||||
return ErrAPIKeyRequired
|
||||
}
|
||||
return nil
|
||||
return ErrAPIKeyRequired
|
||||
}
|
||||
if strings.TrimSpace(os.Getenv(envName)) == "" {
|
||||
return fmt.Errorf("%w: api key environment variable %q is not set", ErrAPIKeyEnvMissing, envName)
|
||||
@@ -701,29 +719,38 @@ func resolveOutputContract(def *domain.PromptDefinition, override *domain.Output
|
||||
return contract, nil
|
||||
}
|
||||
|
||||
const renderedPromptHashVersion = "promptkit/rendered-prompt/v2\x00"
|
||||
|
||||
func hashRenderedPrompt(p domain.RenderedPrompt) string {
|
||||
var b strings.Builder
|
||||
if p.SessionID != "" {
|
||||
b.WriteString("session_id=")
|
||||
b.WriteString(p.SessionID)
|
||||
b.WriteString("\n---\n")
|
||||
}
|
||||
for _, msg := range p.Messages {
|
||||
b.WriteString(msg.Role)
|
||||
b.WriteByte('\n')
|
||||
b.WriteString(msg.Content)
|
||||
if msg.CacheControl != nil {
|
||||
b.WriteString("\ncache_control.type=")
|
||||
b.WriteString(string(msg.CacheControl.Type))
|
||||
if msg.CacheControl.TTL != "" {
|
||||
b.WriteString("\ncache_control.ttl=")
|
||||
b.WriteString(msg.CacheControl.TTL)
|
||||
}
|
||||
hasher := sha256.New()
|
||||
var lengthBuffer [8]byte
|
||||
_, _ = io.WriteString(hasher, renderedPromptHashVersion)
|
||||
writeRenderedPromptHashString(hasher, &lengthBuffer, p.SessionID)
|
||||
writeRenderedPromptHashLength(hasher, &lengthBuffer, len(p.Messages))
|
||||
for _, message := range p.Messages {
|
||||
writeRenderedPromptHashString(hasher, &lengthBuffer, message.Role)
|
||||
writeRenderedPromptHashString(hasher, &lengthBuffer, message.Content)
|
||||
if message.CacheControl == nil {
|
||||
lengthBuffer[0] = 0
|
||||
_, _ = hasher.Write(lengthBuffer[:1])
|
||||
continue
|
||||
}
|
||||
b.WriteString("\n---\n")
|
||||
lengthBuffer[0] = 1
|
||||
_, _ = hasher.Write(lengthBuffer[:1])
|
||||
writeRenderedPromptHashString(hasher, &lengthBuffer, string(message.CacheControl.Type))
|
||||
writeRenderedPromptHashString(hasher, &lengthBuffer, message.CacheControl.TTL)
|
||||
}
|
||||
h := sha256.Sum256([]byte(b.String()))
|
||||
return hex.EncodeToString(h[:])
|
||||
return hex.EncodeToString(hasher.Sum(nil))
|
||||
}
|
||||
|
||||
func writeRenderedPromptHashLength(hasher hash.Hash, buffer *[8]byte, length int) {
|
||||
binary.BigEndian.PutUint64(buffer[:], uint64(length))
|
||||
_, _ = hasher.Write(buffer[:])
|
||||
}
|
||||
|
||||
func writeRenderedPromptHashString(hasher hash.Hash, buffer *[8]byte, value string) {
|
||||
writeRenderedPromptHashLength(hasher, buffer, len(value))
|
||||
_, _ = io.WriteString(hasher, value)
|
||||
}
|
||||
|
||||
func buildOutputArtifact(content string, format domain.OutputFormat) domain.Artifact {
|
||||
|
||||
@@ -295,8 +295,10 @@ func (c *controlledRepairLLM) Generate(
|
||||
ctx context.Context,
|
||||
req domain.GenerateRequest,
|
||||
) (*domain.GenerateResponse, error) {
|
||||
isRepair := len(req.Prompt.Messages) > 0 &&
|
||||
strings.HasPrefix(req.Prompt.Messages[0].Content, "You repair invalid JSON")
|
||||
messages := req.Prompt.Messages
|
||||
isRepair := len(messages) >= 2 &&
|
||||
messages[len(messages)-2].Role == "assistant" &&
|
||||
messages[len(messages)-1].Role == "user"
|
||||
c.mu.Lock()
|
||||
c.calls++
|
||||
c.active++
|
||||
@@ -1108,14 +1110,18 @@ func TestDeriveStructuredSchemaName(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
||||
uncached := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||
{Role: "system", Content: "sys"},
|
||||
{Role: "user", Content: "usr"},
|
||||
func TestHashRenderedPromptIncludesEveryContractedField(t *testing.T) {
|
||||
base := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||
{Role: domain.RoleSystem, Content: "sys"},
|
||||
{Role: domain.RoleUser, Content: "usr"},
|
||||
}}
|
||||
wantLegacyHash := hashString("system\nsys\n---\nuser\nusr\n---\n")
|
||||
if got := hashRenderedPrompt(uncached); got != wantLegacyHash {
|
||||
t.Fatalf("expected no-cache hash to preserve legacy input, got %q want %q", got, wantLegacyHash)
|
||||
identical := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||
{Role: domain.RoleSystem, Content: "sys"},
|
||||
{Role: domain.RoleUser, Content: "usr"},
|
||||
}}
|
||||
baseHash := hashRenderedPrompt(base)
|
||||
if baseHash != hashRenderedPrompt(identical) {
|
||||
t.Fatal("identical rendered prompts must have the same hash")
|
||||
}
|
||||
|
||||
withCache := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||
@@ -1127,7 +1133,7 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
||||
TTL: "1h",
|
||||
},
|
||||
},
|
||||
{Role: "user", Content: "usr"},
|
||||
{Role: domain.RoleUser, Content: "usr"},
|
||||
}}
|
||||
alsoWithCache := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||
{
|
||||
@@ -1138,7 +1144,7 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
||||
TTL: "1h",
|
||||
},
|
||||
},
|
||||
{Role: "user", Content: "usr"},
|
||||
{Role: domain.RoleUser, Content: "usr"},
|
||||
}}
|
||||
withoutTTL := domain.RenderedPrompt{Messages: []domain.RenderedMessage{
|
||||
{
|
||||
@@ -1148,11 +1154,11 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
||||
Type: domain.CacheControlEphemeral,
|
||||
},
|
||||
},
|
||||
{Role: "user", Content: "usr"},
|
||||
{Role: domain.RoleUser, Content: "usr"},
|
||||
}}
|
||||
|
||||
cachedHash := hashRenderedPrompt(withCache)
|
||||
if cachedHash == hashRenderedPrompt(uncached) {
|
||||
if cachedHash == baseHash {
|
||||
t.Fatal("expected cache control to change rendered prompt hash")
|
||||
}
|
||||
if cachedHash != hashRenderedPrompt(alsoWithCache) {
|
||||
@@ -1161,6 +1167,29 @@ func TestHashRenderedPromptIncludesCacheControlWhenPresent(t *testing.T) {
|
||||
if cachedHash == hashRenderedPrompt(withoutTTL) {
|
||||
t.Fatal("expected ttl changes to affect rendered prompt hash")
|
||||
}
|
||||
|
||||
variants := []struct {
|
||||
name string
|
||||
prompt domain.RenderedPrompt
|
||||
}{
|
||||
{name: "role", prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleDeveloper, Content: "sys"}, {Role: domain.RoleUser, Content: "usr"}}}},
|
||||
{name: "content", prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleSystem, Content: "changed"}, {Role: domain.RoleUser, Content: "usr"}}}},
|
||||
{name: "order", prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleUser, Content: "usr"}, {Role: domain.RoleSystem, Content: "sys"}}}},
|
||||
{name: "session", prompt: domain.RenderedPrompt{SessionID: "session", Messages: base.Messages}},
|
||||
}
|
||||
for _, variant := range variants {
|
||||
t.Run(variant.name, func(t *testing.T) {
|
||||
if hashRenderedPrompt(variant.prompt) == baseHash {
|
||||
t.Fatalf("%s did not change the rendered-prompt hash", variant.name)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
oneMessage := domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleUser, Content: "x\n---\nassistant\ny"}}}
|
||||
twoMessages := domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: domain.RoleUser, Content: "x"}, {Role: domain.RoleAssistant, Content: "y"}}}
|
||||
if hashRenderedPrompt(oneMessage) == hashRenderedPrompt(twoMessages) {
|
||||
t.Fatal("length-framed hashing did not distinguish legacy separator collision")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHashRenderedPromptIncludesSessionIDWhenPresent(t *testing.T) {
|
||||
@@ -1202,6 +1231,75 @@ func TestHashRenderedPromptIncludesSessionIDWhenPresent(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerComposesAppendedMessagesIntoEveryRenderedPrompt(t *testing.T) {
|
||||
renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{
|
||||
SessionID: "rendered-session",
|
||||
Messages: []domain.RenderedMessage{
|
||||
{Role: domain.RoleSystem, Content: "ordinary prefix", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
|
||||
{Role: domain.RoleUser, Content: "ordinary request"},
|
||||
},
|
||||
}}
|
||||
client := &fakeLLM{resp: &domain.GenerateResponse{Content: "output"}}
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
renderer,
|
||||
client,
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
appended := []domain.RenderedMessage{
|
||||
{Role: domain.RoleAssistant, Content: "consumer response"},
|
||||
{Role: domain.RoleUser, Content: "consumer correction", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral}},
|
||||
}
|
||||
request := domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef(), AppendedMessages: appended}
|
||||
|
||||
plain, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare without appended messages: %v", err)
|
||||
}
|
||||
empty, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef(), AppendedMessages: []domain.RenderedMessage{}})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare with empty appended messages: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(empty.Messages, plain.Messages) || empty.RenderedPromptHash != plain.RenderedPromptHash {
|
||||
t.Fatalf("empty appended messages changed preparation: empty=%#v plain=%#v", empty, plain)
|
||||
}
|
||||
prepared, err := runner.Prepare(context.Background(), request)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare with appended messages: %v", err)
|
||||
}
|
||||
if prepared.PromptHash != plain.PromptHash || prepared.RenderedPromptHash == plain.RenderedPromptHash {
|
||||
t.Fatalf("hashes with appended messages = (%q, %q), plain = (%q, %q)", prepared.PromptHash, prepared.RenderedPromptHash, plain.PromptHash, plain.RenderedPromptHash)
|
||||
}
|
||||
|
||||
expected := []domain.RenderedMessage{
|
||||
{Role: domain.RoleSystem, Content: "ordinary prefix", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral, TTL: "1h"}},
|
||||
{Role: domain.RoleUser, Content: "ordinary request"},
|
||||
{Role: domain.RoleAssistant, Content: "consumer response"},
|
||||
{Role: domain.RoleUser, Content: "consumer correction", CacheControl: &domain.CacheControl{Type: domain.CacheControlEphemeral}},
|
||||
}
|
||||
if !reflect.DeepEqual(prepared.Messages, expected) {
|
||||
t.Fatalf("prepared messages = %#v, want %#v", prepared.Messages, expected)
|
||||
}
|
||||
|
||||
result, err := runner.Run(context.Background(), request)
|
||||
if err != nil {
|
||||
t.Fatalf("run with appended messages: %v", err)
|
||||
}
|
||||
if result.RenderedPromptHash != prepared.RenderedPromptHash || !reflect.DeepEqual(client.lastReq.Prompt.Messages, expected) {
|
||||
t.Fatalf("run did not use the prepared effective prompt: result=%#v request=%#v prepared=%#v", result, client.lastReq.Prompt.Messages, prepared)
|
||||
}
|
||||
|
||||
renderer.rendered.Messages[0].Content = "changed source"
|
||||
appended[1].CacheControl.Type = "changed caller value"
|
||||
if !reflect.DeepEqual(prepared.Messages, expected) {
|
||||
t.Fatalf("prepared messages changed after source or caller mutation: %#v", prepared.Messages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunSuccessful(t *testing.T) {
|
||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatMarkdown, domain.ValidationBasic, 0)}
|
||||
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}
|
||||
@@ -1787,22 +1885,26 @@ func TestRunnerRunAPIKeyEnvResolvesFromEnvironment(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunAPIKeyEnvMissingEnvironmentValueFailsClearly(t *testing.T) {
|
||||
func TestRunnerRunOptionalAPIKeyEnvMissingEnvironmentValueReachesLLM(t *testing.T) {
|
||||
const environmentName = "PROMPTKIT_MISSING_KEY"
|
||||
t.Setenv(environmentName, "")
|
||||
|
||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
|
||||
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
|
||||
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_MISSING_KEY"},
|
||||
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: environmentName},
|
||||
}}
|
||||
runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil)
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}
|
||||
runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil)
|
||||
|
||||
_, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
result, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
|
||||
if err != nil || result == nil {
|
||||
t.Fatalf("optional credential run = (%+v, %v), want success", result, err)
|
||||
}
|
||||
if !errors.Is(err, ErrAPIKeyEnvMissing) {
|
||||
t.Fatalf("expected ErrAPIKeyEnvMissing, got %v", err)
|
||||
if llmClient.calls != 1 {
|
||||
t.Fatalf("LLM calls = %d, want 1", llmClient.calls)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "PROMPTKIT_MISSING_KEY") {
|
||||
t.Fatalf("expected missing env name in error, got %v", err)
|
||||
if llmClient.lastReq.Target.APIKeyEnv != environmentName {
|
||||
t.Fatalf("LLM api_key_env = %q, want %q", llmClient.lastReq.Target.APIKeyEnv, environmentName)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2139,6 +2241,8 @@ func TestRunnerRepairStateMachine(t *testing.T) {
|
||||
TopP: true,
|
||||
TimeoutSeconds: true,
|
||||
}
|
||||
emptyThenValid := responses(2)
|
||||
emptyThenValid[0].Content = " \t "
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -2161,12 +2265,13 @@ func TestRunnerRepairStateMachine(t *testing.T) {
|
||||
wantStatus: domain.ValidationPassed,
|
||||
},
|
||||
{
|
||||
name: "basic failure is ineligible despite budget",
|
||||
name: "empty basic output repairs successfully",
|
||||
mode: domain.ValidationBasic,
|
||||
budget: 3,
|
||||
validationResults: []domain.ValidationResult{failed(domain.ValidationBasic, "empty output")},
|
||||
responses: responses(1),
|
||||
wantStatus: domain.ValidationFailed,
|
||||
validationResults: []domain.ValidationResult{failed(domain.ValidationBasic, "empty output"), passed(domain.ValidationBasic)},
|
||||
responses: emptyThenValid,
|
||||
wantRepairs: 1,
|
||||
wantStatus: domain.ValidationPassed,
|
||||
},
|
||||
{
|
||||
name: "inherited numeric values remain absent",
|
||||
@@ -2198,7 +2303,7 @@ func TestRunnerRepairStateMachine(t *testing.T) {
|
||||
{
|
||||
name: "successful repair stops below larger budget",
|
||||
mode: domain.ValidationJSON,
|
||||
budget: 4,
|
||||
budget: 3,
|
||||
validationResults: []domain.ValidationResult{
|
||||
failed(domain.ValidationJSON, "candidate zero"),
|
||||
failed(domain.ValidationJSON, "candidate one"),
|
||||
@@ -2301,6 +2406,9 @@ func TestRunnerRepairStateMachine(t *testing.T) {
|
||||
!reflect.DeepEqual(req.ValidationErrors, tc.validationResults[index].Errors) {
|
||||
t.Fatalf("repair request %d prior state = %+v", index, req)
|
||||
}
|
||||
if !reflect.DeepEqual(req.OriginalMessages, initialRequest.Prompt.Messages) {
|
||||
t.Fatalf("repair request %d original messages drifted: %#v", index, req.OriginalMessages)
|
||||
}
|
||||
if req.TargetPresence != tc.wantPresence || !reflect.DeepEqual(req.Target, initialRequest.Target) ||
|
||||
req.SessionID != initialRequest.Prompt.SessionID ||
|
||||
!reflect.DeepEqual(req.StructuredOutput, initialRequest.StructuredOutput) {
|
||||
@@ -2314,8 +2422,19 @@ func TestRunnerRepairStateMachine(t *testing.T) {
|
||||
!reflect.DeepEqual(generated.StructuredOutput, initialRequest.StructuredOutput) {
|
||||
t.Fatalf("repair generation request %d common fields drifted: %+v", index, generated)
|
||||
}
|
||||
if reflect.DeepEqual(generated.Prompt.Messages, initialRequest.Prompt.Messages) {
|
||||
t.Fatalf("repair generation request %d reused the initial prompt", index)
|
||||
expectedMessages := len(initialRequest.Prompt.Messages) + 1
|
||||
if strings.TrimSpace(tc.responses[index].Content) != "" {
|
||||
expectedMessages++
|
||||
}
|
||||
if len(generated.Prompt.Messages) != expectedMessages ||
|
||||
generated.Prompt.Messages[len(generated.Prompt.Messages)-1].Role != "user" {
|
||||
t.Fatalf("repair generation request %d messages = %#v", index, generated.Prompt.Messages)
|
||||
}
|
||||
if strings.TrimSpace(tc.responses[index].Content) != "" {
|
||||
assistant := generated.Prompt.Messages[len(generated.Prompt.Messages)-2]
|
||||
if assistant.Role != "assistant" || assistant.Content != tc.responses[index].Content {
|
||||
t.Fatalf("repair generation request %d candidate = %+v", index, assistant)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -174,9 +174,8 @@ func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract d
|
||||
}
|
||||
|
||||
res := domain.ValidationResult{
|
||||
Mode: contract.ValidationMode,
|
||||
SchemaPath: contract.SchemaPath,
|
||||
RepairAttempts: contract.RepairAttempts,
|
||||
Mode: contract.ValidationMode,
|
||||
SchemaPath: contract.SchemaPath,
|
||||
}
|
||||
|
||||
if artifact == nil {
|
||||
|
||||
@@ -64,6 +64,43 @@ func TestStandardValidatorBasicFailureEmpty(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorReportsZeroRepairAttempts(t *testing.T) {
|
||||
v := NewStandardValidator("")
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
body string
|
||||
contract domain.OutputContract
|
||||
}{
|
||||
{
|
||||
name: "passed basic validation",
|
||||
body: "answer",
|
||||
contract: domain.OutputContract{
|
||||
ValidationMode: domain.ValidationBasic,
|
||||
RepairAttempts: 3,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "failed JSON validation",
|
||||
body: `{"answer":`,
|
||||
contract: domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSON,
|
||||
RepairAttempts: 3,
|
||||
},
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
res, err := v.Validate(context.Background(), &domain.Artifact{Body: []byte(tc.body)}, tc.contract)
|
||||
if err != nil {
|
||||
t.Fatalf("validate: %v", err)
|
||||
}
|
||||
if res.RepairAttempts != 0 {
|
||||
t.Fatalf("repair attempts = %d, want 0", res.RepairAttempts)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorJSONSuccess(t *testing.T) {
|
||||
v := NewStandardValidator("")
|
||||
|
||||
|
||||
14
message_roles.go
Normal file
14
message_roles.go
Normal file
@@ -0,0 +1,14 @@
|
||||
package promptkit
|
||||
|
||||
import "gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
|
||||
const (
|
||||
// RoleDeveloper identifies a provider-bound developer instruction message.
|
||||
RoleDeveloper = domain.RoleDeveloper
|
||||
// RoleSystem identifies a provider-bound system instruction message.
|
||||
RoleSystem = domain.RoleSystem
|
||||
// RoleUser identifies a provider-bound user message.
|
||||
RoleUser = domain.RoleUser
|
||||
// RoleAssistant identifies a provider-bound assistant message.
|
||||
RoleAssistant = domain.RoleAssistant
|
||||
)
|
||||
@@ -22,27 +22,52 @@ func TestPreparationRejectsInvalidOutputContractWithPublicError(t *testing.T) {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
|
||||
req := promptkit.RunRequest{
|
||||
PromptID: "prompt",
|
||||
Validation: &promptkit.OutputContract{
|
||||
Format: promptkit.OutputFormat("binary"),
|
||||
ValidationMode: promptkit.ValidationNone,
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
contract promptkit.OutputContract
|
||||
}{
|
||||
{
|
||||
name: "unsupported format",
|
||||
contract: promptkit.OutputContract{
|
||||
Format: promptkit.OutputFormat("binary"),
|
||||
ValidationMode: promptkit.ValidationNone,
|
||||
},
|
||||
},
|
||||
}
|
||||
{
|
||||
name: "repair attempts above maximum",
|
||||
contract: promptkit.OutputContract{
|
||||
Format: promptkit.FormatText,
|
||||
ValidationMode: promptkit.ValidationBasic,
|
||||
RepairAttempts: 4,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "none validation with repair attempts",
|
||||
contract: promptkit.OutputContract{
|
||||
Format: promptkit.FormatText,
|
||||
ValidationMode: promptkit.ValidationNone,
|
||||
RepairAttempts: 1,
|
||||
},
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
req := promptkit.RunRequest{PromptID: "prompt", Validation: &tc.contract}
|
||||
|
||||
prepared, err := engine.Prepare(context.Background(), req)
|
||||
if prepared != nil {
|
||||
t.Fatalf("expected no partial prepared run, got %+v", prepared)
|
||||
}
|
||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("prepare error = %v, want ErrInvalidRequest", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), req)
|
||||
if prepared != nil {
|
||||
t.Fatalf("expected no partial prepared run, got %+v", prepared)
|
||||
}
|
||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("prepare error = %v, want ErrInvalidRequest", err)
|
||||
}
|
||||
|
||||
preparedExecution, err := engine.PrepareExecution(context.Background(), req)
|
||||
if preparedExecution != nil {
|
||||
t.Fatalf("expected no partial prepared execution, got %+v", preparedExecution)
|
||||
}
|
||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("prepare execution error = %v, want ErrInvalidRequest", err)
|
||||
preparedExecution, err := engine.PrepareExecution(context.Background(), req)
|
||||
if preparedExecution != nil {
|
||||
t.Fatalf("expected no partial prepared execution, got %+v", preparedExecution)
|
||||
}
|
||||
if !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("prepare execution error = %v, want ErrInvalidRequest", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,7 +17,9 @@ type PreparedExecution struct {
|
||||
|
||||
// Details returns a fresh caller-owned, credential-redacted copy of the
|
||||
// prepared request details. Mutating the result cannot affect execution or a
|
||||
// later Details call. Details remains available after execution or discard.
|
||||
// later Details call. Its complete effective message content, including any
|
||||
// appended request messages, remains subject to the caller's data-handling
|
||||
// policy. Details remains available after execution or discard.
|
||||
//
|
||||
// A nil receiver or zero-value PreparedExecution returns a zero [PreparedRun].
|
||||
func (p *PreparedExecution) Details() PreparedRun {
|
||||
|
||||
@@ -5,7 +5,6 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -257,6 +256,50 @@ func TestPreparedExecutionLifecycleAndEngineBinding(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedExecutionRepairsEmptyBasicOutput(t *testing.T) {
|
||||
client := &preparedRecordingClient{responses: []*promptkit.GenerateResponse{
|
||||
{
|
||||
Content: "",
|
||||
Usage: promptkit.TokenUsage{PromptTokens: 3, CompletionTokens: 5, TotalTokens: 8},
|
||||
},
|
||||
{
|
||||
Content: "Corrected summary.",
|
||||
Usage: promptkit.TokenUsage{PromptTokens: 7, CompletionTokens: 11, TotalTokens: 18},
|
||||
},
|
||||
}}
|
||||
engine := newPreparedContractEngine(t, client, "Summarize the source.")
|
||||
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{
|
||||
PromptID: "prepared",
|
||||
Validation: &promptkit.OutputContract{
|
||||
Format: promptkit.FormatMarkdown,
|
||||
ValidationMode: promptkit.ValidationBasic,
|
||||
RepairAttempts: 1,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
details := prepared.Details()
|
||||
|
||||
result, err := engine.RunPrepared(context.Background(), prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("run prepared: %v", err)
|
||||
}
|
||||
if result.RawOutput != "Corrected summary." || result.Validation.Status != promptkit.ValidationPassed ||
|
||||
result.Validation.RepairAttempts != 1 || result.Usage != (promptkit.TokenUsage{PromptTokens: 10, CompletionTokens: 16, TotalTokens: 26}) {
|
||||
t.Fatalf("repaired result = %+v", result)
|
||||
}
|
||||
requests := client.snapshot()
|
||||
if len(requests) != 2 || len(requests[1].Prompt.Messages) != len(details.Messages)+1 ||
|
||||
!reflect.DeepEqual(requests[1].Prompt.Messages[:len(details.Messages)], details.Messages) ||
|
||||
requests[1].Prompt.Messages[len(requests[1].Prompt.Messages)-1].Role != "user" {
|
||||
t.Fatalf("prepared repair requests = %#v", requests)
|
||||
}
|
||||
if _, err := engine.RunPrepared(context.Background(), prepared); !errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("second RunPrepared error = %v, want ErrInvalidRequest", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedExecutionConcurrentClaimAllowsOneGeneration(t *testing.T) {
|
||||
release := make(chan struct{})
|
||||
client := &preparedRecordingClient{
|
||||
@@ -524,19 +567,25 @@ func TestPreparedExecutionCredentialCapacityAndTimingBoundaries(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prepared", "profile", "content"), "."),
|
||||
promptkit.WithProfileFS(preparedCredentialProfileSource(environmentName), "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "profile",
|
||||
Endpoint: "http://example.test/v1",
|
||||
Model: "model",
|
||||
APIKeyRequired: true,
|
||||
}),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct credential engine: %v", err)
|
||||
}
|
||||
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prepared"})
|
||||
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{
|
||||
PromptID: "prepared",
|
||||
Execution: &promptkit.ExecutionTargetOverride{APIKeyEnv: environmentName},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare credential execution: %v", err)
|
||||
}
|
||||
if err := os.Unsetenv(environmentName); err != nil {
|
||||
t.Fatalf("unset credential environment: %v", err)
|
||||
}
|
||||
t.Setenv(environmentName, "")
|
||||
|
||||
result, err := engine.RunPrepared(context.Background(), prepared)
|
||||
if result != nil ||
|
||||
@@ -666,12 +715,13 @@ func (r *mutablePreparedArtifactReader) callCount() int {
|
||||
}
|
||||
|
||||
type preparedRecordingClient struct {
|
||||
mu sync.Mutex
|
||||
response *promptkit.GenerateResponse
|
||||
err error
|
||||
requests []promptkit.GenerateRequest
|
||||
started chan struct{}
|
||||
release <-chan struct{}
|
||||
mu sync.Mutex
|
||||
response *promptkit.GenerateResponse
|
||||
responses []*promptkit.GenerateResponse
|
||||
err error
|
||||
requests []promptkit.GenerateRequest
|
||||
started chan struct{}
|
||||
release <-chan struct{}
|
||||
}
|
||||
|
||||
func (c *preparedRecordingClient) Generate(
|
||||
@@ -695,6 +745,12 @@ func (c *preparedRecordingClient) Generate(
|
||||
if c.err != nil {
|
||||
return nil, c.err
|
||||
}
|
||||
if len(c.responses) > 0 {
|
||||
if index := len(c.requests) - 1; index < len(c.responses) {
|
||||
return c.responses[index], nil
|
||||
}
|
||||
return nil, fmt.Errorf("no response configured for generation %d", len(c.requests))
|
||||
}
|
||||
return c.response, nil
|
||||
}
|
||||
|
||||
@@ -755,16 +811,6 @@ model: ` + model + `
|
||||
}
|
||||
}
|
||||
|
||||
func preparedCredentialProfileSource(environmentName string) fstest.MapFS {
|
||||
return fstest.MapFS{
|
||||
"profile.yaml": &fstest.MapFile{Data: []byte(`id: profile
|
||||
endpoint: http://example.test/v1
|
||||
model: model
|
||||
api_key_env: ` + environmentName + `
|
||||
`)},
|
||||
}
|
||||
}
|
||||
|
||||
func preparedSchemaSource() fstest.MapFS {
|
||||
return fstest.MapFS{
|
||||
"schema.json": &fstest.MapFile{Data: []byte(`{
|
||||
|
||||
243
profile_inheritance_contract_test.go
Normal file
243
profile_inheritance_contract_test.go
Normal file
@@ -0,0 +1,243 @@
|
||||
package promptkit_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io/fs"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit"
|
||||
)
|
||||
|
||||
func TestProfileInheritanceBuiltInAliasWorkflow(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: frameworkPromptDir,
|
||||
SchemaDir: frameworkSchemaDir,
|
||||
}, promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "weather-light",
|
||||
BaseProfileID: "deepseek-4-flash",
|
||||
ReasoningEffort: "high",
|
||||
TimeoutSeconds: 120,
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("construct alias engine: %v", err)
|
||||
}
|
||||
|
||||
base, err := engine.InspectProfile(context.Background(), "deepseek-4-flash")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect base: %v", err)
|
||||
}
|
||||
child, err := engine.InspectProfile(context.Background(), "weather-light")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect alias: %v", err)
|
||||
}
|
||||
if child.ProfileID != "weather-light" ||
|
||||
child.EffectiveModelParams.BackendID != base.EffectiveModelParams.BackendID ||
|
||||
child.EffectiveModelParams.Model != base.EffectiveModelParams.Model ||
|
||||
child.EffectiveModelParams.ReasoningEffort != "high" ||
|
||||
child.EffectiveModelParams.TimeoutSeconds != 120 {
|
||||
t.Fatalf("alias inspection = %+v, base = %+v", child, base)
|
||||
}
|
||||
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{
|
||||
PromptID: frameworkMarkdownSummaryPromptID,
|
||||
ProfileID: "weather-light",
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"transcript": promptkit.Inline("Rin opens the gate."),
|
||||
"glossary": promptkit.Inline("gate: A guarded passage."),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare alias: %v", err)
|
||||
}
|
||||
if prepared.SelectedProfileID != "weather-light" {
|
||||
t.Fatalf("SelectedProfileID = %q", prepared.SelectedProfileID)
|
||||
}
|
||||
|
||||
timeout := 15
|
||||
reasoning := "low"
|
||||
overridden, err := engine.Prepare(context.Background(), promptkit.RunRequest{
|
||||
PromptID: frameworkMarkdownSummaryPromptID,
|
||||
ProfileID: "weather-light",
|
||||
Execution: &promptkit.ExecutionTargetOverride{
|
||||
TimeoutSeconds: &timeout,
|
||||
ReasoningEffort: &reasoning,
|
||||
},
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"transcript": promptkit.Inline("Rin opens the gate."),
|
||||
"glossary": promptkit.Inline("gate: A guarded passage."),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare override: %v", err)
|
||||
}
|
||||
if overridden.EffectiveModelParams.TimeoutSeconds != timeout ||
|
||||
overridden.EffectiveModelParams.ReasoningEffort != reasoning {
|
||||
t.Fatalf("runtime override target = %+v", overridden.EffectiveModelParams)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileInheritanceYAMLAliasOfBuiltIn(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(fstest.MapFS{}, "."),
|
||||
promptkit.WithProfileFS(fstest.MapFS{
|
||||
"alias.yaml": &fstest.MapFile{Data: []byte("id: yaml-alias\nbase_profile: deepseek-4-flash\n")},
|
||||
}, "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct YAML alias engine: %v", err)
|
||||
}
|
||||
inspection, err := engine.InspectProfile(context.Background(), "yaml-alias")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect YAML alias: %v", err)
|
||||
}
|
||||
if inspection.ProfileID != "yaml-alias" || inspection.EffectiveModelParams.Model == "" {
|
||||
t.Fatalf("YAML alias inspection = %+v", inspection)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileInheritanceRetainsRequiredCredentialBehavior(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(fstest.MapFS{}, "."),
|
||||
promptkit.WithBackend(promptkit.Backend{
|
||||
ID: "credential-backend",
|
||||
Endpoint: "https://credential.example/v1",
|
||||
APIKeyEnv: "OPTIONAL_BACKEND_KEY",
|
||||
}),
|
||||
promptkit.WithProfiles(
|
||||
promptkit.Profile{
|
||||
ID: "credential-base",
|
||||
BackendID: "credential-backend",
|
||||
Model: "model",
|
||||
APIKeyRequired: true,
|
||||
},
|
||||
promptkit.Profile{ID: "credential-child", BaseProfileID: "credential-base"},
|
||||
),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct credential inheritance engine: %v", err)
|
||||
}
|
||||
inspection, err := engine.InspectProfile(context.Background(), "credential-child")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect credential child: %v", err)
|
||||
}
|
||||
if !inspection.APIKeyRequired || inspection.EffectiveModelParams.APIKeyEnv != "" {
|
||||
t.Fatalf("credential inspection = %+v", inspection)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileInheritancePreservesPublicErrorIdentities(t *testing.T) {
|
||||
newEngine := func(t *testing.T, profiles fs.FS) *promptkit.Engine {
|
||||
t.Helper()
|
||||
options := []promptkit.Option{promptkit.WithPromptFS(fstest.MapFS{}, ".")}
|
||||
if profiles != nil {
|
||||
options = append(options, promptkit.WithProfileFS(profiles, "."))
|
||||
}
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{}, options...)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
return engine
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
profiles fs.FS
|
||||
profile string
|
||||
contains []string
|
||||
want error
|
||||
wantNot error
|
||||
}{
|
||||
{
|
||||
name: "missing selected profile",
|
||||
profile: "missing",
|
||||
want: promptkit.ErrProfileNotFound,
|
||||
wantNot: promptkit.ErrProfileLoad,
|
||||
},
|
||||
{
|
||||
name: "missing base",
|
||||
profiles: fstest.MapFS{
|
||||
"child.yaml": &fstest.MapFile{Data: []byte("id: child\nbase_profile: missing\n")},
|
||||
},
|
||||
profile: "child",
|
||||
contains: []string{"child", "missing"},
|
||||
want: promptkit.ErrProfileLoad,
|
||||
wantNot: promptkit.ErrProfileNotFound,
|
||||
},
|
||||
{
|
||||
name: "cycle",
|
||||
profiles: fstest.MapFS{
|
||||
"a.yaml": &fstest.MapFile{Data: []byte("id: a\nbase_profile: b\n")},
|
||||
"b.yaml": &fstest.MapFile{Data: []byte("id: b\nbase_profile: a\n")},
|
||||
},
|
||||
profile: "a",
|
||||
contains: []string{"a", "b"},
|
||||
want: promptkit.ErrProfileLoad,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
result, err := newEngine(t, tc.profiles).InspectProfile(context.Background(), tc.profile)
|
||||
if result != nil || !errors.Is(err, tc.want) || (tc.wantNot != nil && errors.Is(err, tc.wantNot)) {
|
||||
t.Fatalf("inspection=(%+v, %v), want %v without %v", result, err, tc.want, tc.wantNot)
|
||||
}
|
||||
for _, fragment := range tc.contains {
|
||||
if !strings.Contains(err.Error(), fragment) {
|
||||
t.Fatalf("error = %v, want %q", err, fragment)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileInheritanceFreezesPreparedExecution(t *testing.T) {
|
||||
profiles := &mutableInheritanceProfileFS{files: fstest.MapFS{
|
||||
"child.yaml": &fstest.MapFile{Data: []byte("id: child\nbase_profile: base\n")},
|
||||
"base.yaml": &fstest.MapFile{Data: []byte("id: base\nendpoint: https://base.example/v1\nmodel: first-model\n")},
|
||||
}}
|
||||
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "ok"}}
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prepared", "child", "content"), "."),
|
||||
promptkit.WithProfileFS(profiles, "."),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
|
||||
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prepared"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
profiles.set("base.yaml", "id: base\nendpoint: https://base.example/v1\nmodel: second-model\n")
|
||||
|
||||
result, err := engine.RunPrepared(context.Background(), prepared)
|
||||
if err != nil || result == nil || len(client.requests) != 1 || client.requests[0].Target.Model != "first-model" {
|
||||
t.Fatalf("prepared execution=(%+v, %v), requests=%+v", result, err, client.requests)
|
||||
}
|
||||
inspection, err := engine.InspectProfile(context.Background(), "child")
|
||||
if err != nil || inspection.EffectiveModelParams.Model != "second-model" {
|
||||
t.Fatalf("fresh inspection=(%+v, %v)", inspection, err)
|
||||
}
|
||||
}
|
||||
|
||||
type mutableInheritanceProfileFS struct {
|
||||
mu sync.RWMutex
|
||||
files fstest.MapFS
|
||||
}
|
||||
|
||||
func (f *mutableInheritanceProfileFS) Open(name string) (fs.File, error) {
|
||||
f.mu.RLock()
|
||||
defer f.mu.RUnlock()
|
||||
return f.files.Open(name)
|
||||
}
|
||||
|
||||
func (f *mutableInheritanceProfileFS) set(name, content string) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.files[name] = &fstest.MapFile{Data: []byte(content)}
|
||||
}
|
||||
41
profiles.go
41
profiles.go
@@ -2,17 +2,16 @@ package promptkit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
)
|
||||
|
||||
// OpenAICompatibleProfile returns an ordinary in-memory Profile for an
|
||||
// OpenAI-compatible chat-completions endpoint.
|
||||
// OpenAICompatibleProfile returns an in-memory Profile for an OpenAI-compatible
|
||||
// chat-completions endpoint. A non-blank BaseProfileID permits its target
|
||||
// fields to be inherited when the profile is selected or inspected.
|
||||
//
|
||||
// It does not register global state, maintain a model catalog, or resolve
|
||||
// credentials. If APIKeyRequired is true, callers satisfy it with
|
||||
@@ -25,6 +24,7 @@ import (
|
||||
func OpenAICompatibleProfile(cfg OpenAICompatibleProfileConfig) Profile {
|
||||
return Profile{
|
||||
ID: cfg.ID,
|
||||
BaseProfileID: cfg.BaseProfileID,
|
||||
BackendID: cfg.BackendID,
|
||||
Endpoint: cfg.Endpoint,
|
||||
Model: cfg.Model,
|
||||
@@ -87,8 +87,9 @@ func toDomainProfile(publicProfile Profile) (domain.ExecutionProfile, error) {
|
||||
return domain.ExecutionProfile{}, err
|
||||
}
|
||||
prof := domain.ExecutionProfile{
|
||||
ID: strings.TrimSpace(publicProfile.ID),
|
||||
BackendID: strings.TrimSpace(publicProfile.BackendID),
|
||||
ID: publicProfile.ID,
|
||||
BaseProfileID: publicProfile.BaseProfileID,
|
||||
BackendID: publicProfile.BackendID,
|
||||
Endpoint: publicProfile.Endpoint,
|
||||
Model: publicProfile.Model,
|
||||
Temperature: publicProfile.Temperature,
|
||||
@@ -100,34 +101,8 @@ func toDomainProfile(publicProfile Profile) (domain.ExecutionProfile, error) {
|
||||
APIKeyRequired: publicProfile.APIKeyRequired,
|
||||
ExtraParams: extraParams,
|
||||
}
|
||||
if err := normalizeAndValidatePublicProfile(&prof); err != nil {
|
||||
if err := profile.NormalizeAndValidateDefinition(&prof); err != nil {
|
||||
return domain.ExecutionProfile{}, err
|
||||
}
|
||||
return prof, nil
|
||||
}
|
||||
|
||||
func normalizeAndValidatePublicProfile(prof *domain.ExecutionProfile) error {
|
||||
if strings.TrimSpace(prof.ID) == "" {
|
||||
return errors.New("id is required")
|
||||
}
|
||||
prof.Endpoint = strings.TrimSpace(prof.Endpoint)
|
||||
if strings.TrimSpace(prof.BackendID) == "" && prof.Endpoint == "" {
|
||||
return errors.New("backend or endpoint is required")
|
||||
}
|
||||
if prof.Endpoint != "" {
|
||||
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(prof.Endpoint)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
prof.Endpoint = endpoint
|
||||
}
|
||||
if strings.TrimSpace(prof.Model) == "" {
|
||||
return errors.New("model is required")
|
||||
}
|
||||
return domain.ValidateExecutionTargetSettings(domain.ExecutionTarget{
|
||||
Temperature: prof.Temperature,
|
||||
MaxTokens: prof.MaxTokens,
|
||||
TopP: prof.TopP,
|
||||
TimeoutSeconds: prof.TimeoutSeconds,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -784,9 +784,12 @@ func TestBackendRegistrationRejectsInvalidAndDuplicateDefinitions(t *testing.T)
|
||||
{ID: " custom ", Endpoint: "http://one.example/v1"},
|
||||
{ID: "custom", Endpoint: "http://two.example/v1"},
|
||||
}},
|
||||
{name: "reserved built-in id", backends: []promptkit.Backend{{
|
||||
{name: "reserved OpenRouter ID", backends: []promptkit.Backend{{
|
||||
ID: promptkit.BackendOpenRouter, Endpoint: "http://replacement.example/v1",
|
||||
}}},
|
||||
{name: "reserved Rakestrawhome ID", backends: []promptkit.Backend{{
|
||||
ID: promptkit.BackendRakestrawHome, Endpoint: "http://replacement.example/v1",
|
||||
}}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -841,7 +844,7 @@ func TestBackendExtraParamsAreDeeplyCopiedAtConstructionAndLookup(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEngineValidationIsSinglePass(t *testing.T) {
|
||||
func TestEngineValidationWithZeroRepairBudgetIsSinglePass(t *testing.T) {
|
||||
client := &fakeLLMClient{
|
||||
response: &promptkit.GenerateResponse{Content: "not-json"},
|
||||
}
|
||||
@@ -856,7 +859,7 @@ func TestEngineValidationIsSinglePass(t *testing.T) {
|
||||
Validation: &promptkit.OutputContract{
|
||||
Format: promptkit.FormatJSON,
|
||||
ValidationMode: promptkit.ValidationJSON,
|
||||
RepairAttempts: 3,
|
||||
RepairAttempts: 0,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
@@ -871,6 +874,87 @@ func TestEngineValidationIsSinglePass(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEngineRunRepairsJSONSchemaOutput(t *testing.T) {
|
||||
client := &fakeLLMClient{responses: []*promptkit.GenerateResponse{
|
||||
{
|
||||
Content: "{}",
|
||||
Usage: promptkit.TokenUsage{PromptTokens: 2, CompletionTokens: 3, TotalTokens: 5, CachedTokens: 7, CacheWriteTokens: 11},
|
||||
},
|
||||
{
|
||||
Content: `{"events":[{"title":"Repaired event"}]}`,
|
||||
Usage: promptkit.TokenUsage{PromptTokens: 13, CompletionTokens: 17, TotalTokens: 19, CachedTokens: 23, CacheWriteTokens: 29},
|
||||
},
|
||||
}}
|
||||
engine := newContractEngineWithOptions(t, frameworkSchemaDir, promptkit.WithLLMClient(client))
|
||||
|
||||
result, err := engine.Run(context.Background(), promptkit.RunRequest{
|
||||
PromptID: frameworkStructuredEventsPromptID,
|
||||
SessionID: " repair-session ",
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"transcript": promptkit.Inline("Rin opens the gate."),
|
||||
"glossary": promptkit.Inline("gate: A guarded passage."),
|
||||
},
|
||||
Validation: &promptkit.OutputContract{
|
||||
Format: promptkit.FormatJSON,
|
||||
ValidationMode: promptkit.ValidationJSONSchema,
|
||||
SchemaPath: "structured_events.schema.json",
|
||||
RepairAttempts: 1,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("run: %v", err)
|
||||
}
|
||||
if result.RawOutput != client.responses[1].Content || result.Validation.Status != promptkit.ValidationPassed ||
|
||||
result.Validation.RepairAttempts != 1 {
|
||||
t.Fatalf("repaired result = %+v", result)
|
||||
}
|
||||
wantUsage := promptkit.TokenUsage{PromptTokens: 15, CompletionTokens: 20, TotalTokens: 24, CachedTokens: 30, CacheWriteTokens: 40}
|
||||
if result.Usage != wantUsage {
|
||||
t.Fatalf("usage = %+v, want %+v", result.Usage, wantUsage)
|
||||
}
|
||||
if len(client.requests) != 2 {
|
||||
t.Fatalf("generation calls = %d, want 2", len(client.requests))
|
||||
}
|
||||
initial, repaired := client.requests[0], client.requests[1]
|
||||
if initial.Prompt.SessionID != "repair-session" || repaired.Prompt.SessionID != initial.Prompt.SessionID ||
|
||||
!reflect.DeepEqual(repaired.Target, initial.Target) || repaired.TargetPresence != initial.TargetPresence ||
|
||||
!reflect.DeepEqual(repaired.StructuredOutput, initial.StructuredOutput) {
|
||||
t.Fatalf("generation request state drifted: initial=%+v repaired=%+v", initial, repaired)
|
||||
}
|
||||
if initial.StructuredOutput == nil || initial.StructuredOutput.JSONSchema == nil {
|
||||
t.Fatalf("expected structured output on initial request: %+v", initial)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEngineRunReturnsFinalResultAfterRepairExhaustion(t *testing.T) {
|
||||
client := &fakeLLMClient{responses: []*promptkit.GenerateResponse{
|
||||
{Content: "not-json", Usage: promptkit.TokenUsage{PromptTokens: 2, CompletionTokens: 3, TotalTokens: 5}},
|
||||
{Content: "still-not-json", Usage: promptkit.TokenUsage{PromptTokens: 7, CompletionTokens: 11, TotalTokens: 18}},
|
||||
}}
|
||||
engine := newContractEngineWithOptions(t, frameworkSchemaDir, promptkit.WithLLMClient(client))
|
||||
|
||||
result, err := engine.Run(context.Background(), promptkit.RunRequest{
|
||||
PromptID: frameworkMarkdownSummaryPromptID,
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"transcript": promptkit.Inline("Rin opens the gate."),
|
||||
"glossary": promptkit.Inline("gate: A guarded passage."),
|
||||
},
|
||||
Validation: &promptkit.OutputContract{
|
||||
Format: promptkit.FormatJSON,
|
||||
ValidationMode: promptkit.ValidationJSON,
|
||||
RepairAttempts: 1,
|
||||
},
|
||||
})
|
||||
if err != nil || result == nil {
|
||||
t.Fatalf("run = (%+v, %v), want exhausted result", result, err)
|
||||
}
|
||||
if result.RawOutput != "still-not-json" || result.Validation.Status != promptkit.ValidationFailed ||
|
||||
result.Validation.RepairAttempts != 1 || len(result.Validation.Errors) == 0 ||
|
||||
result.Usage != (promptkit.TokenUsage{PromptTokens: 9, CompletionTokens: 14, TotalTokens: 23}) {
|
||||
t.Fatalf("exhausted result = %+v", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRepeatedOptionsUseLastValueInEachCategory(t *testing.T) {
|
||||
profile := promptkit.Profile{ID: "profile", Endpoint: "http://example.test/v1", Model: "model"}
|
||||
|
||||
|
||||
424
testdata/builtin-catalog-v1.json
vendored
Normal file
424
testdata/builtin-catalog-v1.json
vendored
Normal file
@@ -0,0 +1,424 @@
|
||||
{
|
||||
"backends": [
|
||||
{
|
||||
"id": "openrouter",
|
||||
"endpoint": "https://openrouter.ai/api/v1",
|
||||
"api_key_env": "OPENROUTER_API_KEY",
|
||||
"extra_params": null,
|
||||
"concurrency_limit": 16,
|
||||
"queue_capacity": 1024,
|
||||
"queue_capacity_set": true
|
||||
},
|
||||
{
|
||||
"id": "rakestrawhome",
|
||||
"endpoint": "https://inference.ai.rakestrawhome.com/v1",
|
||||
"api_key_env": "RAKESTRAWHOME_INFERENCE_API_KEY",
|
||||
"extra_params": null,
|
||||
"concurrency_limit": 4,
|
||||
"queue_capacity": 1024,
|
||||
"queue_capacity_set": true
|
||||
}
|
||||
],
|
||||
"profiles": [
|
||||
{
|
||||
"id": "aion-2",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "aion-labs/aion-2.0",
|
||||
"temperature": 0.72,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0.95,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "claude-fable-latest",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "~anthropic/claude-fable-latest",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 600,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "claude-haiku-latest",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "~anthropic/claude-haiku-latest",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "medium",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "claude-opus-latest",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "~anthropic/claude-opus-latest",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "claude-sonnet-latest",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "~anthropic/claude-sonnet-latest",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "deepseek-3-2",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "deepseek/deepseek-v3.2",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "deepseek-4-flash",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "deepseek/deepseek-v4-flash",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "deepseek-4-pro",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "deepseek/deepseek-v4-pro",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gemini-2-flash",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "google/gemini-2.5-flash",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gemini-2-flash-lite",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "google/gemini-2.5-flash-lite",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gemini-2-pro",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "google/gemini-2.5-pro",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gemini-3-flash-lite",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "google/gemini-3.1-flash-lite",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gemini-flash-latest",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "~google/gemini-flash-latest",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gemini-pro-latest",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "~google/gemini-pro-latest",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gemma-4-31b",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "google/gemma-4-31b-it:exacto",
|
||||
"temperature": 0.15,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0.98,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gpt-5-mini",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "openai/gpt-5.4-mini",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "gpt-5-nano",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "openai/gpt-5.4-nano",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 240,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "minimax-m2",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "minimax/minimax-m2.5",
|
||||
"temperature": 0.5,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0.95,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "minimax-m3",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "minimax/minimax-m3",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "mistral-large-2512",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "mistralai/mistral-large-2512",
|
||||
"temperature": 0.15,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0.98,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "",
|
||||
"reasoning_effort": "",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "mistral-medium-3-5",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "mistralai/mistral-medium-3-5",
|
||||
"temperature": 0.15,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0.98,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "mistral-small-3",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "mistralai/mistral-small-3.2-24b-instruct",
|
||||
"temperature": 0.05,
|
||||
"max_tokens": 0,
|
||||
"top_p": 1,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "",
|
||||
"reasoning_effort": "",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "mistral-small-4",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "mistralai/mistral-small-2603",
|
||||
"temperature": 0.1,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0.98,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "nemotron-3-ultra",
|
||||
"base_profile": "",
|
||||
"backend": "openrouter",
|
||||
"endpoint": "",
|
||||
"model": "nvidia/nemotron-3-ultra-550b-a55b",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 180,
|
||||
"service_tier": "flex",
|
||||
"reasoning_effort": "high",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
},
|
||||
{
|
||||
"id": "rakestrawhome-gemma-4-31b",
|
||||
"base_profile": "",
|
||||
"backend": "rakestrawhome",
|
||||
"endpoint": "",
|
||||
"model": "google/gemma-4-31b-it",
|
||||
"temperature": 0,
|
||||
"max_tokens": 0,
|
||||
"top_p": 0,
|
||||
"timeout_seconds": 0,
|
||||
"service_tier": "",
|
||||
"reasoning_effort": "",
|
||||
"api_key_env": "",
|
||||
"api_key_required": false,
|
||||
"extra_params": null
|
||||
}
|
||||
]
|
||||
}
|
||||
127
types.go
127
types.go
@@ -123,6 +123,21 @@ type RunRequest struct {
|
||||
// Validation optionally replaces the prompt's complete output contract. It
|
||||
// does not merge individual fields. Nil uses the prompt contract.
|
||||
Validation *OutputContract
|
||||
// AppendedMessages are already-rendered messages appended in caller order
|
||||
// after every prompt definition message. Promptkit neither templates nor
|
||||
// resolves files in them. Content must be valid UTF-8 and is preserved
|
||||
// exactly, including empty or whitespace-only content. Roles are trimmed and
|
||||
// lowercased, then must be [RoleDeveloper], [RoleSystem], [RoleUser], or
|
||||
// [RoleAssistant]. Any CacheControl is normalized as documented on that
|
||||
// type.
|
||||
//
|
||||
// Nil and empty slices are equivalent. Promptkit imposes no message-count,
|
||||
// byte-size, token, or context-window limit and does not truncate content;
|
||||
// an upstream rejection follows the ordinary generation-error contract.
|
||||
// Prepare, PrepareExecution, and Run validate and copy the messages before
|
||||
// source or model work. Malformed values return an error matching
|
||||
// ErrInvalidRequest.
|
||||
AppendedMessages []RenderedMessage
|
||||
}
|
||||
|
||||
// PreparedRun contains prepared prompt execution state returned by
|
||||
@@ -162,8 +177,9 @@ type PreparedRun struct {
|
||||
SessionID string `json:"session_id,omitempty"`
|
||||
// RenderedPromptHash is an opaque equality value for SessionID and Messages.
|
||||
RenderedPromptHash string `json:"rendered_prompt_hash"`
|
||||
// Messages are the rendered messages that Run or RunPrepared passes to the
|
||||
// LLM client.
|
||||
// Messages are the complete effective messages that Run or RunPrepared passes
|
||||
// to the LLM client: definition messages in rendered order followed by any
|
||||
// RunRequest.AppendedMessages in caller order.
|
||||
Messages []RenderedMessage `json:"messages"`
|
||||
// StartTime is the UTC time at which preparation began.
|
||||
StartTime time.Time `json:"start_time,omitempty"`
|
||||
@@ -322,7 +338,10 @@ type ExecutionTarget struct {
|
||||
// ReasoningEffort is the effective opaque provider-specific reasoning
|
||||
// setting. An empty value instructs model clients to omit reasoning.
|
||||
ReasoningEffort string `json:"reasoning_effort"`
|
||||
// APIKeyEnv is an environment-variable name, not its credential value.
|
||||
// APIKeyEnv is the resolved name of an optional environment lookup source,
|
||||
// not its credential value. The built-in client omits Authorization when no
|
||||
// usable direct or environment credential is available; injected clients may
|
||||
// resolve this metadata differently.
|
||||
APIKeyEnv string `json:"api_key_env"`
|
||||
// ExtraParams contains copied JSON-compatible provider parameters.
|
||||
ExtraParams map[string]any `json:"extra_params"`
|
||||
@@ -426,9 +445,10 @@ type ExecutionTargetOverride struct {
|
||||
// inherited value and disables reasoning for this run. Non-blank values
|
||||
// are opaque and are not validated against a fixed vocabulary.
|
||||
ReasoningEffort *string
|
||||
// APIKeyEnv replaces the profile or backend environment-variable name when
|
||||
// non-blank. A direct RunRequest.APIKey still takes precedence over
|
||||
// environment lookup.
|
||||
// APIKeyEnv replaces the profile or backend optional environment lookup
|
||||
// source when non-blank. A direct RunRequest.APIKey still takes precedence.
|
||||
// The built-in client omits Authorization when neither source has a usable
|
||||
// value; injected clients may resolve this metadata differently.
|
||||
APIKeyEnv string
|
||||
// ExtraParams, when non-empty, replaces the complete profile or backend map.
|
||||
// Values must be JSON-compatible: nil, booleans, finite numbers, strings,
|
||||
@@ -439,31 +459,43 @@ type ExecutionTargetOverride struct {
|
||||
|
||||
// Profile is an in-memory execution profile for library consumers.
|
||||
//
|
||||
// It is equivalent to a loaded profile file after validation. Raw API keys do
|
||||
// not belong in profiles; use APIKeyRequired to require callers to provide a
|
||||
// RunRequest.APIKey or explicit request ExecutionTargetOverride.APIKeyEnv, or
|
||||
// use profile YAML api_key_env with file and FS profile sources. Profile has no
|
||||
// stable JSON representation.
|
||||
// A standalone Profile is equivalent to a loaded profile file after local
|
||||
// validation. A derived profile names BaseProfileID and can inherit target
|
||||
// fields when selected or inspected. Raw API keys do not belong in profiles;
|
||||
// use APIKeyRequired to require callers to provide a RunRequest.APIKey or
|
||||
// explicit request ExecutionTargetOverride.APIKeyEnv, or use profile YAML
|
||||
// api_key_env with file and FS profile sources. Profile has no stable JSON
|
||||
// representation.
|
||||
//
|
||||
// WithProfiles validates and copies Profile values during NewEngine. Zero
|
||||
// Temperature, MaxTokens, and TopP values and blank ServiceTier and
|
||||
// ReasoningEffort values leave those provider controls unspecified. A zero
|
||||
// TimeoutSeconds retains the framework deadline, while an empty ExtraParams map
|
||||
// inherits backend request defaults. Use ExecutionTargetOverride pointer fields
|
||||
// to request an explicit numeric zero.
|
||||
// WithProfiles locally validates and copies Profile values during NewEngine.
|
||||
// It checks base-reference existence and resolved target completeness when a
|
||||
// derived profile is selected or inspected. Zero Temperature, MaxTokens, and
|
||||
// TopP values and blank ServiceTier and ReasoningEffort values leave those
|
||||
// provider controls unspecified. A zero TimeoutSeconds retains the framework
|
||||
// deadline, while an empty ExtraParams map inherits backend request defaults.
|
||||
// Use ExecutionTargetOverride pointer fields to request an explicit numeric
|
||||
// zero.
|
||||
type Profile struct {
|
||||
// ID is the required non-blank profile identifier. WithProfiles trims it.
|
||||
ID string
|
||||
// BaseProfileID optionally names one base profile. WithProfiles trims it. A
|
||||
// non-blank value permits required target fields to be inherited when the
|
||||
// profile is selected or inspected, which is also when reference existence
|
||||
// and resolved completeness are checked. A blank value leaves this as a
|
||||
// standalone profile.
|
||||
BaseProfileID string
|
||||
// BackendID optionally selects an engine backend. WithProfiles trims it.
|
||||
// Backend membership is checked when a request selects the profile; an
|
||||
// unknown ID makes preparation fail with ErrProfileLoad.
|
||||
BackendID string
|
||||
// Endpoint is the model-provider base URL. It is required only when
|
||||
// BackendID is blank and otherwise overrides the backend endpoint when
|
||||
// Endpoint is the model-provider base URL. A standalone Profile requires an
|
||||
// endpoint when BackendID is blank; a derived Profile may inherit either
|
||||
// field. A non-blank endpoint overrides the backend endpoint when
|
||||
// non-blank. WithProfiles trims it and requires an absolute HTTP or HTTPS URL
|
||||
// with a host and no user information, query, or fragment.
|
||||
Endpoint string
|
||||
// Model is the required non-blank provider model identifier.
|
||||
// Model is the provider model identifier. It is required for a standalone
|
||||
// Profile and may be inherited by a derived Profile.
|
||||
Model string
|
||||
// Temperature is from 0 through 2. Zero leaves the provider control
|
||||
// unspecified.
|
||||
@@ -481,7 +513,8 @@ type Profile struct {
|
||||
ReasoningEffort string
|
||||
// APIKeyRequired clears a backend's inherited API-key environment name and
|
||||
// requires a non-blank RunRequest.APIKey unless the request explicitly
|
||||
// supplies ExecutionTargetOverride.APIKeyEnv. It does not store a credential.
|
||||
// supplies ExecutionTargetOverride.APIKeyEnv. When false, a named
|
||||
// environment source remains optional. It does not store a credential.
|
||||
APIKeyRequired bool
|
||||
// ExtraParams contains provider-specific JSON-compatible values. An empty
|
||||
// map inherits backend request defaults, when any. WithProfiles validates
|
||||
@@ -494,13 +527,16 @@ type Profile struct {
|
||||
// profile.
|
||||
//
|
||||
// It contains ordinary profile fields for OpenAI-compatible chat-completions
|
||||
// endpoints. APIKeyRequired follows Profile.APIKeyRequired. Raw API keys do not
|
||||
// belong in this config. OpenAICompatibleProfileConfig has no stable JSON
|
||||
// endpoints. BaseProfileID and APIKeyRequired follow Profile. Raw API keys do
|
||||
// not belong in this config. OpenAICompatibleProfileConfig has no stable JSON
|
||||
// representation and is not validated until its resulting Profile is supplied
|
||||
// through WithProfiles to NewEngine.
|
||||
type OpenAICompatibleProfileConfig struct {
|
||||
// ID becomes Profile.ID.
|
||||
ID string
|
||||
// BaseProfileID becomes Profile.BaseProfileID. A non-blank value permits the
|
||||
// resulting Profile to inherit target fields when it is selected or inspected.
|
||||
BaseProfileID string
|
||||
// BackendID becomes Profile.BackendID.
|
||||
BackendID string
|
||||
// Endpoint becomes Profile.Endpoint.
|
||||
@@ -545,8 +581,8 @@ type ExecutionTargetPresence struct {
|
||||
// JSON representation.
|
||||
//
|
||||
// A non-nil RunRequest.Validation replaces the complete prompt contract. It
|
||||
// does not merge fields. The public Engine validates generated output once and
|
||||
// does not install an output repairer.
|
||||
// does not merge fields. The public Engine performs bounded correction after a
|
||||
// failed eligible validation when RepairAttempts is positive.
|
||||
type OutputContract struct {
|
||||
// Format selects generated artifact metadata. An empty value in a non-nil
|
||||
// request replacement defaults to FormatText.
|
||||
@@ -557,9 +593,9 @@ type OutputContract struct {
|
||||
// SchemaPath is required when ValidationMode is ValidationJSONSchema and is
|
||||
// ignored by other modes.
|
||||
SchemaPath string `json:"schema_path"`
|
||||
// RepairAttempts is a non-negative requested repair limit. Zero requests no
|
||||
// repairs. The public Engine performs no repairs even when this value is
|
||||
// positive, so its runs report zero attempts used.
|
||||
// RepairAttempts is an additional generation-call budget from zero through
|
||||
// three. Zero is single-pass. A positive value is valid only with basic,
|
||||
// json, or json_schema validation.
|
||||
RepairAttempts int `json:"repair_attempts"`
|
||||
}
|
||||
|
||||
@@ -575,8 +611,8 @@ type ValidationResult struct {
|
||||
Errors []string `json:"errors,omitempty"`
|
||||
// SchemaPath is the effective schema path for JSON Schema validation.
|
||||
SchemaPath string `json:"schema_path,omitempty"`
|
||||
// RepairAttempts is the number of repairs actually attempted. It is always
|
||||
// zero for the public Engine.
|
||||
// RepairAttempts is the number of corrective generation calls actually
|
||||
// started for this result.
|
||||
RepairAttempts int `json:"repair_attempts"`
|
||||
// IsValid is true for ValidationPassed and ValidationSkipped and false for
|
||||
// ValidationFailed.
|
||||
@@ -605,23 +641,29 @@ type RenderedPrompt struct {
|
||||
// SessionID is the optional effective direct or rendered session
|
||||
// identifier supplied to the model client.
|
||||
SessionID string `json:"session_id,omitempty"`
|
||||
// Messages contains rendered messages in definition order.
|
||||
// Messages contains the effective messages in provider order: definition
|
||||
// messages in rendered order followed by any RunRequest.AppendedMessages.
|
||||
Messages []RenderedMessage `json:"messages"`
|
||||
}
|
||||
|
||||
// RenderedMessage is a rendered chat message and has a stable JSON
|
||||
// representation.
|
||||
// RenderedMessage is a prepared, provider-bound text chat message and has a
|
||||
// stable JSON representation. At a request boundary its role is trimmed and
|
||||
// lowercased, then must be one of [RoleDeveloper], [RoleSystem], [RoleUser],
|
||||
// or [RoleAssistant].
|
||||
type RenderedMessage struct {
|
||||
// Role is the definition-supplied chat role.
|
||||
// Role is the provider-bound chat role.
|
||||
Role string `json:"role"`
|
||||
// Content is the rendered message text.
|
||||
// Content is the provider-bound message text. Request-supplied content must
|
||||
// be valid UTF-8 and may be empty or whitespace-only.
|
||||
Content string `json:"content"`
|
||||
// CacheControl is optional provider cache metadata.
|
||||
CacheControl *CacheControl `json:"cache_control,omitempty"`
|
||||
}
|
||||
|
||||
// CacheControl describes provider cache metadata attached to prompt content
|
||||
// and has a stable JSON representation.
|
||||
// and has a stable JSON representation. At a request boundary Type and TTL
|
||||
// must be valid UTF-8 and are trimmed. Type must be [CacheControlEphemeral],
|
||||
// and TTL must be empty or "1h"; other values make the request invalid.
|
||||
type CacheControl struct {
|
||||
// Type identifies the cache behavior.
|
||||
Type CacheControlType `json:"type"`
|
||||
@@ -665,10 +707,10 @@ type StructuredOutputJSONSpec struct {
|
||||
// and retained copies. It is responsible for the cancellation behavior of any
|
||||
// work it starts and for synchronizing access to retained or shared data.
|
||||
//
|
||||
// A returned error makes Run or RunPrepared return ErrLLMGenerate while
|
||||
// preserving the client error through errors.Is. A nil response with a nil
|
||||
// error also produces ErrLLMGenerate. Promptkit copies the non-nil response
|
||||
// before returning from either method.
|
||||
// An arbitrary returned error makes Run or RunPrepared return ErrLLMGenerate
|
||||
// while preserving the client error through errors.Is rather than translating
|
||||
// it. A nil response with a nil error also produces ErrLLMGenerate. Promptkit
|
||||
// copies the non-nil response before returning from either method.
|
||||
type LLMClient interface {
|
||||
Generate(context.Context, GenerateRequest) (*GenerateResponse, error)
|
||||
}
|
||||
@@ -694,9 +736,8 @@ type GenerateRequest struct {
|
||||
// GenerateResponse is returned by an injected LLM client and has a stable JSON
|
||||
// representation.
|
||||
type GenerateResponse struct {
|
||||
// Content is the generated output. It must be non-empty when using the
|
||||
// built-in client; injected clients may return empty content for Promptkit
|
||||
// validation to classify.
|
||||
// Content is the generated output. It may be explicitly empty; Promptkit
|
||||
// applies the effective output contract to classify it.
|
||||
Content string `json:"content"`
|
||||
// Usage is the client's token accounting.
|
||||
Usage TokenUsage `json:"usage"`
|
||||
|
||||
Reference in New Issue
Block a user