Compare commits
73 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4ca3be2c14 | |||
| 227fb35f99 | |||
| e291b8bfe9 | |||
| 2b6a7f83c4 | |||
| 3a43550f70 | |||
| c281f721bc | |||
| 350b0e76d9 | |||
| e43350fd0d | |||
| 20d3e3b5ee | |||
| a93b799236 | |||
| e83a3ce179 | |||
| a04a3bbc5f | |||
| 731b66cff5 | |||
| 70e0ea0cf0 | |||
| d45c474c1e | |||
| 25f1ba0b30 | |||
| a718762da1 | |||
| 58ac3ce298 | |||
| 57f2ce1ce4 | |||
| c8b6d5c490 | |||
| abeb50b525 | |||
| 1cb07c7d91 | |||
| 8cfc71c351 | |||
| 5ccfa4a345 | |||
| 14e03f19d0 | |||
| 7c562a9374 | |||
| ef97d85ac9 | |||
| 805e48c873 | |||
| e9e126dcba | |||
| 32e7a3557c | |||
| 1d1b04e2e0 | |||
| 9748897751 | |||
| 9d020039d5 | |||
| 5247ce0b73 | |||
| c434aa1dae | |||
| ac9b3f3d80 | |||
| 4f12a89a1b | |||
| df31e7f58e | |||
| 0678d242b9 | |||
| 3b4ea21208 | |||
| 1430e85147 | |||
| 34d7a19da5 | |||
| ebf1602635 | |||
| 31f2ce3a09 | |||
| fd06e4ca6b | |||
| e63b8de1e9 | |||
| 9354d2b373 | |||
| 01ca5430bd | |||
| ae2179d103 | |||
| a248433d0f | |||
| bd6cffc9d0 | |||
| e40c4f182b | |||
| 7428e50c2c | |||
| 63c67a4520 | |||
| 25a7052a3d | |||
| fc3255967e | |||
| e920168b30 | |||
| 272b6a4bc1 | |||
| dde48a31fc | |||
| 242eace4a7 | |||
| 0bf5f88136 | |||
| 369ab5392d | |||
| 2ba0146e5d | |||
| 6112c2af0c | |||
| f5e12c00f5 | |||
| 49fe402dd2 | |||
| c301eb8d55 | |||
| c13e9710d9 | |||
| 87b5ec3d75 | |||
| cb4028a637 | |||
| 5a1bff4529 | |||
| 805a7f965d | |||
| 147f5e5ff5 |
12
README.md
12
README.md
@@ -33,6 +33,18 @@ 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).
|
||||
|
||||
Earlier adopters 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
|
||||
[v0.4.0 changelog and adoption guide](docs/releases/v0.4.0.md).
|
||||
|
||||
Consumers upgrading from `v0.2.0` to `v0.3.0` should read the
|
||||
[v0.3.0 changelog](docs/releases/v0.3.0.md).
|
||||
|
||||
Consumers moving from `v0.1.0` to `v0.2.0` should read the
|
||||
[v0.2.0 changelog and migration guide](docs/releases/v0.2.0.md).
|
||||
|
||||
|
||||
40
backends.go
40
backends.go
@@ -9,6 +9,11 @@ import (
|
||||
// backend.
|
||||
const BackendOpenRouter = backend.OpenRouterID
|
||||
|
||||
// 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].
|
||||
const BackendLocal = "local"
|
||||
|
||||
// Backend configures one engine-scoped OpenAI-compatible backend.
|
||||
//
|
||||
// Backend has no stable JSON representation. Use keyed literals so additions
|
||||
@@ -29,21 +34,42 @@ type Backend struct {
|
||||
// JSON-compatible, finite, acyclic, and keyed by non-empty strings. Keys
|
||||
// must not be model, session_id, messages, temperature, max_tokens, top_p,
|
||||
// service_tier, reasoning_effort, or response_format. An empty map supplies
|
||||
// no defaults. NewEngine deeply copies the map.
|
||||
// no defaults. NewEngine deeply copies the map and rejects excessively deep
|
||||
// or large values for safety.
|
||||
ExtraParams map[string]any
|
||||
// ConcurrencyLimit is the maximum number of simultaneous model-generation
|
||||
// calls allowed for this backend within one Engine. Zero leaves the backend
|
||||
// unlimited. A negative value makes NewEngine fail with ErrInvalidConfig.
|
||||
ConcurrencyLimit int
|
||||
// QueueCapacity controls how many additional Run calls may be admitted
|
||||
// beyond ConcurrencyLimit. Nil uses 1024 when ConcurrencyLimit is positive;
|
||||
// a pointer uses its exact value, including zero. The pointed-to value must
|
||||
// be non-negative, and QueueCapacity must be nil when ConcurrencyLimit is
|
||||
// zero. Their sum must fit in an int. WithBackend copies the value and does
|
||||
// not retain the pointer.
|
||||
// QueueCapacity controls how many additional Run or RunPrepared calls may
|
||||
// be admitted beyond ConcurrencyLimit. Nil uses 1024 when ConcurrencyLimit
|
||||
// is positive; a pointer uses its exact value, including zero. The pointed-to
|
||||
// value must be non-negative, and QueueCapacity must be nil when
|
||||
// ConcurrencyLimit is zero. Their sum must fit in an int. WithBackend copies
|
||||
// the value and does not retain the pointer.
|
||||
QueueCapacity *int
|
||||
}
|
||||
|
||||
// LocalBackend returns a caller-owned Backend for a conventional local
|
||||
// OpenAI-compatible endpoint. It sets ID to BackendLocal and copies endpoint
|
||||
// and concurrencyLimit into Endpoint and ConcurrencyLimit without
|
||||
// normalization or validation. APIKeyEnv, ExtraParams, and QueueCapacity keep
|
||||
// their zero values.
|
||||
//
|
||||
// LocalBackend does not read environment variables, register the value, or
|
||||
// mutate engine or package state. Supply the returned value through
|
||||
// [WithBackend]; [NewEngine] then applies the ordinary backend validation and
|
||||
// concurrency semantics, including default queue capacity for a positive
|
||||
// limit, unlimited behavior for zero, and ErrInvalidConfig for a negative
|
||||
// limit.
|
||||
func LocalBackend(endpoint string, concurrencyLimit int) Backend {
|
||||
return Backend{
|
||||
ID: BackendLocal,
|
||||
Endpoint: endpoint,
|
||||
ConcurrencyLimit: concurrencyLimit,
|
||||
}
|
||||
}
|
||||
|
||||
// WithBackend adds one Backend registration to the constructed Engine.
|
||||
//
|
||||
// Registrations accumulate in option order. Every normalized ID must be unique
|
||||
|
||||
@@ -60,7 +60,18 @@ func TestEngineRejectsRunBeforeCompletionWhenAdmissionIsFull(t *testing.T) {
|
||||
|
||||
awaitCapacitySignal(t, reader.entered, "first artifact read")
|
||||
|
||||
result, err := engine.Run(context.Background(), capacityInputRequest("http://second.example/v1"))
|
||||
canceledContext, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
result, err := engine.Run(canceledContext, capacityInputRequest("http://canceled.example/v1"))
|
||||
if result != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("canceled capacity admission=(%+v, %v), want context cancellation", result, err)
|
||||
}
|
||||
var canceledCapacityErr *promptkit.CapacityError
|
||||
if errors.Is(err, promptkit.ErrCapacityExceeded) || errors.As(err, &canceledCapacityErr) {
|
||||
t.Fatalf("canceled admission exposed capacity rejection: %v", err)
|
||||
}
|
||||
|
||||
result, err = engine.Run(context.Background(), capacityInputRequest("http://second.example/v1"))
|
||||
if result != nil {
|
||||
t.Fatalf("capacity rejection returned partial result: %+v", result)
|
||||
}
|
||||
@@ -70,6 +81,21 @@ func TestEngineRejectsRunBeforeCompletionWhenAdmissionIsFull(t *testing.T) {
|
||||
if errors.Is(err, promptkit.ErrInvalidRequest) || errors.Is(err, promptkit.ErrLLMGenerate) {
|
||||
t.Fatalf("capacity rejection had an unrelated category: %v", err)
|
||||
}
|
||||
var capacityErr *promptkit.CapacityError
|
||||
if !errors.As(err, &capacityErr) || capacityErr == nil {
|
||||
t.Fatalf("capacity rejection=%v, want CapacityError", err)
|
||||
}
|
||||
if capacityErr.BackendID != "limited" {
|
||||
t.Fatalf("capacity backend ID=%q, want limited", capacityErr.BackendID)
|
||||
}
|
||||
capacityErr.BackendID = "changed"
|
||||
|
||||
result, err = engine.Run(context.Background(), capacityInputRequest("http://third.example/v1"))
|
||||
var subsequentCapacityErr *promptkit.CapacityError
|
||||
if result != nil || !errors.As(err, &subsequentCapacityErr) ||
|
||||
subsequentCapacityErr == nil || subsequentCapacityErr.BackendID != "limited" {
|
||||
t.Fatalf("subsequent capacity rejection=(%+v, %v), want independent limited CapacityError", result, err)
|
||||
}
|
||||
if calls := reader.callCount(); calls != 1 {
|
||||
t.Fatalf("artifact calls=%d, want only the admitted run", calls)
|
||||
}
|
||||
@@ -174,6 +200,19 @@ func TestCapacityExceededSentinelContract(t *testing.T) {
|
||||
if promptkit.ErrCapacityExceeded == nil {
|
||||
t.Fatal("ErrCapacityExceeded is nil")
|
||||
}
|
||||
var nilCapacityErr *promptkit.CapacityError
|
||||
zeroCapacityErr := &promptkit.CapacityError{}
|
||||
populatedCapacityErr := &promptkit.CapacityError{BackendID: "limited"}
|
||||
for _, capacityErr := range []error{nilCapacityErr, zeroCapacityErr} {
|
||||
if !errors.Is(capacityErr, promptkit.ErrCapacityExceeded) {
|
||||
t.Fatalf("capacity error=%v, want ErrCapacityExceeded", capacityErr)
|
||||
}
|
||||
}
|
||||
var discoveredCapacityErr *promptkit.CapacityError
|
||||
if !errors.As(populatedCapacityErr, &discoveredCapacityErr) || discoveredCapacityErr != populatedCapacityErr {
|
||||
t.Fatalf("populated capacity error is not discoverable: %v", populatedCapacityErr)
|
||||
}
|
||||
|
||||
for _, unrelated := range []error{
|
||||
promptkit.ErrInvalidConfig,
|
||||
promptkit.ErrInvalidRequest,
|
||||
@@ -181,7 +220,8 @@ func TestCapacityExceededSentinelContract(t *testing.T) {
|
||||
promptkit.ErrValidation,
|
||||
} {
|
||||
if errors.Is(promptkit.ErrCapacityExceeded, unrelated) ||
|
||||
errors.Is(unrelated, promptkit.ErrCapacityExceeded) {
|
||||
errors.Is(unrelated, promptkit.ErrCapacityExceeded) ||
|
||||
errors.Is(populatedCapacityErr, unrelated) {
|
||||
t.Fatalf("ErrCapacityExceeded aliases unrelated sentinel %v", unrelated)
|
||||
}
|
||||
}
|
||||
|
||||
40
capacity_error.go
Normal file
40
capacity_error.go
Normal file
@@ -0,0 +1,40 @@
|
||||
package promptkit
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// CapacityError reports bounded admission rejected for a selected backend.
|
||||
//
|
||||
// Engine-produced values identify only rejection at Promptkit's bounded
|
||||
// [Engine.Run] or [Engine.RunPrepared] admission boundary. BackendID is the
|
||||
// normalized registered backend ID used for routing and capacity; endpoint
|
||||
// overrides do not change it. Every engine-produced value is nonnil and has a
|
||||
// nonblank BackendID. Provider errors, active-generation waiting, and caller
|
||||
// cancellation are not represented by this type.
|
||||
//
|
||||
// Callers own returned values and may mutate BackendID without affecting engine
|
||||
// state or another error. CapacityError and its default Go encoding have no
|
||||
// stable JSON contract. Consumer-constructed values do not establish that an
|
||||
// engine rejected work.
|
||||
type CapacityError struct {
|
||||
// BackendID is the normalized registered backend ID whose admission was
|
||||
// rejected.
|
||||
BackendID string
|
||||
}
|
||||
|
||||
// Error returns diagnostic wording that is not a parsing contract. It is safe
|
||||
// to call on a nil receiver or a value with a blank BackendID.
|
||||
func (e *CapacityError) Error() string {
|
||||
if e == nil || strings.TrimSpace(e.BackendID) == "" {
|
||||
return ErrCapacityExceeded.Error()
|
||||
}
|
||||
return fmt.Sprintf("backend %q admission: %v", e.BackendID, ErrCapacityExceeded)
|
||||
}
|
||||
|
||||
// Unwrap returns ErrCapacityExceeded so errors.Is and errors.As can be used
|
||||
// together. It is safe to call on a nil receiver or a zero value.
|
||||
func (e *CapacityError) Unwrap() error {
|
||||
return ErrCapacityExceeded
|
||||
}
|
||||
34
convert.go
34
convert.go
@@ -170,6 +170,40 @@ func fromDomainExecutionTarget(target domain.ExecutionTarget) ExecutionTarget {
|
||||
}
|
||||
}
|
||||
|
||||
func fromDomainProfileInspection(inspection *domain.ProfileInspection) *ProfileInspection {
|
||||
if inspection == nil {
|
||||
return nil
|
||||
}
|
||||
return &ProfileInspection{
|
||||
ProfileID: inspection.ProfileID,
|
||||
EffectiveModelParams: fromDomainExecutionTarget(inspection.EffectiveModelParams),
|
||||
APIKeyRequired: inspection.APIKeyRequired,
|
||||
}
|
||||
}
|
||||
|
||||
func fromDomainPromptInspection(inspection *domain.PromptInspection) *PromptInspection {
|
||||
if inspection == nil {
|
||||
return nil
|
||||
}
|
||||
inputs := make([]PromptInputDefinition, len(inspection.Inputs))
|
||||
for i, input := range inspection.Inputs {
|
||||
inputs[i] = PromptInputDefinition{
|
||||
Name: input.Name,
|
||||
Required: input.Required,
|
||||
ContentType: input.ContentType,
|
||||
Description: input.Description,
|
||||
}
|
||||
}
|
||||
return &PromptInspection{
|
||||
PromptID: inspection.PromptID,
|
||||
PromptVersion: inspection.PromptVersion,
|
||||
PromptHash: inspection.PromptHash,
|
||||
DefaultProfileID: inspection.DefaultProfileID,
|
||||
Inputs: inputs,
|
||||
OutputContract: fromDomainOutputContract(inspection.OutputContract),
|
||||
}
|
||||
}
|
||||
|
||||
func fromDomainExecutionTargetPresence(presence domain.ExecutionTargetPresence) ExecutionTargetPresence {
|
||||
return ExecutionTargetPresence{
|
||||
Temperature: presence.Temperature,
|
||||
|
||||
41
doc.go
41
doc.go
@@ -3,23 +3,28 @@
|
||||
//
|
||||
// Applications construct an [Engine] with [NewEngine], select filesystem or
|
||||
// in-memory sources and optional engine-scoped [Backend] registrations, and
|
||||
// call [Engine.Prepare] or [Engine.Run]. Concrete registries, repositories,
|
||||
// validators, and the built-in OpenAI-compatible client remain internal
|
||||
// implementation details.
|
||||
// call [Engine.InspectPrompt], [Engine.InspectProfile], [Engine.Prepare],
|
||||
// [Engine.PrepareExecution], [Engine.Run], or [Engine.RunPrepared]. Concrete
|
||||
// registries, repositories, validators, and the built-in OpenAI-compatible
|
||||
// client remain internal implementation details.
|
||||
//
|
||||
// # Concurrency and ownership
|
||||
//
|
||||
// An Engine supports concurrent Prepare and Run calls. Engine-local backend
|
||||
// policies bound admitted Run calls and model generations where configured,
|
||||
// while different backend pools and unlimited backends continue independently.
|
||||
// An injected [LLMClient] or [ArtifactReader] can therefore still receive
|
||||
// concurrent calls and must be safe for that use.
|
||||
// An Engine supports concurrent InspectPrompt, InspectProfile, Prepare,
|
||||
// PrepareExecution, Run, and RunPrepared calls. Engine-local backend policies
|
||||
// bound admitted Run and RunPrepared calls and model generations where
|
||||
// configured, while different backend pools and unlimited backends continue
|
||||
// independently. An injected [LLMClient] or [ArtifactReader] can therefore
|
||||
// still receive concurrent calls and must be safe for that use.
|
||||
//
|
||||
// NewEngine copies in-memory profiles and backend definitions. Prepare and Run
|
||||
// copy request maps, slices, pointer values, and JSON-compatible extra
|
||||
// parameters before using them. 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.
|
||||
// NewEngine copies in-memory profiles and backend definitions. Prepare,
|
||||
// PrepareExecution, and Run copy request maps, slices, pointer values, and
|
||||
// JSON-compatible extra parameters before using them. InspectPrompt and
|
||||
// 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.
|
||||
//
|
||||
// # Security and sensitive data
|
||||
//
|
||||
@@ -45,10 +50,12 @@
|
||||
// [GenerateResponse], [ExecutionTargetPresence], and the string value types
|
||||
// used by those values.
|
||||
//
|
||||
// Construction values, including [Config], [Backend], [RunRequest],
|
||||
// [ArtifactRef], [ExecutionTargetOverride], [Profile], and
|
||||
// [OpenAICompatibleProfileConfig], do not have stable JSON representations.
|
||||
// Direct API keys are nevertheless excluded from JSON for every public value.
|
||||
// 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.
|
||||
//
|
||||
// JSON timestamps use time.Time's RFC 3339 encoding and are omitted when zero.
|
||||
// PreparedRun and RunResult durations are encoded as integer milliseconds in
|
||||
|
||||
@@ -40,11 +40,67 @@ validation, and default transport behavior. Source discovery, format
|
||||
validation, and profile precedence are defined by the
|
||||
[framework format reference](../formats.md).
|
||||
|
||||
## Supply Embedded Application Defaults
|
||||
|
||||
Use `WithFallbackProfileFS` when an application packages profile definitions
|
||||
that should apply unless an operator provides an ordinary configured profile
|
||||
with the same ID. For example, an application can embed its defaults while
|
||||
continuing to use `ProfileDir` for operator overrides:
|
||||
|
||||
```go
|
||||
//go:embed profiles/*.yaml
|
||||
var applicationProfiles embed.FS
|
||||
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: "prompts",
|
||||
ProfileDir: operatorProfileDir,
|
||||
},
|
||||
promptkit.WithFallbackProfileFS(applicationProfiles, "profiles"),
|
||||
)
|
||||
```
|
||||
|
||||
Keep application-owned profile IDs and definitions in the embedded source.
|
||||
Use the ordinary configured profile source for operator overrides. Leave
|
||||
`operatorProfileDir` empty when the operator did not configure an override
|
||||
directory; a non-empty path names an authoritative higher-precedence source,
|
||||
so an unavailable or unreadable directory is an error rather than a reason to
|
||||
fall back. The
|
||||
[framework format reference](../formats.md#source-and-profile-precedence)
|
||||
owns the exact profile format and lookup order; the
|
||||
[`WithFallbackProfileFS` GoDoc](../../engine.go) owns its option contract and
|
||||
validation rules.
|
||||
|
||||
## Inspect A Prompt Before Preparation
|
||||
|
||||
Use [`Engine.InspectPrompt`](../../engine.go) to check one configured prompt's
|
||||
declared inputs and output workflow without creating placeholder inputs or
|
||||
resolving a profile:
|
||||
|
||||
```go
|
||||
inspection, err := engine.InspectPrompt(ctx, "meeting.summary", "")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, input := range inspection.Inputs {
|
||||
// Compare the declared input with application configuration.
|
||||
}
|
||||
```
|
||||
|
||||
Use this configuration-time boundary when the application needs only the
|
||||
declared prompt interface. Use `InspectProfile` separately when it must also
|
||||
check a configured profile. Use `Prepare` when it needs inputs, schemas, or
|
||||
rendered messages, and use prepared execution when that work must remain tied
|
||||
to later execution. The method's [GoDoc](../../engine.go) owns exact fields,
|
||||
hash, ownership, and error semantics.
|
||||
|
||||
## Prepare Without Model Execution
|
||||
|
||||
[`Engine.Prepare`](../../engine.go) resolves the selected prompt and profile,
|
||||
loads inputs and any structured-output schema, and renders messages without
|
||||
calling a model client:
|
||||
calling a model client. Choose it when the prepared value is the final
|
||||
inspection or persistence result and no later execution must be tied to that
|
||||
exact snapshot:
|
||||
|
||||
```go
|
||||
prepared, err := engine.Prepare(ctx, promptkit.RunRequest{
|
||||
@@ -61,12 +117,48 @@ shows a complete runnable setup with a prompt file, in-memory profile, and
|
||||
inline input. Exact request requirements and prepared-result fields belong to
|
||||
the [`RunRequest` and `PreparedRun` GoDoc](../../types.go).
|
||||
|
||||
## Prepare Now And Execute The Same Snapshot Later
|
||||
|
||||
Use [`Engine.PrepareExecution`](../../engine.go) when an application must
|
||||
inspect or persist preflight details before deciding whether to start model
|
||||
work, while ensuring that later execution uses those exact rendered messages,
|
||||
inputs, target settings, and validation resources:
|
||||
|
||||
```go
|
||||
preparedExecution, err := engine.PrepareExecution(ctx, promptkit.RunRequest{
|
||||
PromptID: "meeting.summary",
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"note": promptkit.Inline("Synthetic meeting notes"),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer preparedExecution.Discard()
|
||||
|
||||
details := preparedExecution.Details()
|
||||
// Inspect or persist an application-selected safe subset of details.
|
||||
|
||||
result, err := engine.RunPrepared(ctx, preparedExecution)
|
||||
```
|
||||
|
||||
Preparation does not call the model or reserve backend capacity.
|
||||
`RunPrepared` executes from the retained snapshot rather than reloading
|
||||
consumer sources. The handle is opaque in-process state, while `Details`
|
||||
contains rendered content and remains subject to the application's data
|
||||
handling policy. The
|
||||
[`PreparedExecution` and method GoDoc](../../prepared_execution.go) and
|
||||
[engine operation GoDoc](../../engine.go) own exact lifecycle, engine-binding,
|
||||
credential, cancellation, timing, and error semantics.
|
||||
|
||||
## Execute And Validate
|
||||
|
||||
[`Engine.Run`](../../engine.go) performs the same preparation, invokes the
|
||||
configured model client, classifies the generated artifact, and validates the
|
||||
content. A completed content check may return `ValidationFailed` in the result;
|
||||
an operational inability to validate returns an error.
|
||||
content in one call. Choose it when the application does not need a preflight
|
||||
boundary tied to the eventual execution. A completed content check may return
|
||||
`ValidationFailed` in the result; an operational inability to validate returns
|
||||
an error.
|
||||
|
||||
The maintained
|
||||
[offline execution example](../../examples/go-library/run/main.go) injects a
|
||||
@@ -90,13 +182,41 @@ replace execution settings or the complete output contract.
|
||||
The [public value GoDoc](../../types.go) defines nil, empty, zero, replacement,
|
||||
copy, and credential behavior. The
|
||||
[framework format reference](../formats.md) defines how those request values
|
||||
interact with prompt definitions, file-backed profiles, built-ins, schemas,
|
||||
and framework defaults.
|
||||
interact with prompt definitions, file-backed and application fallback
|
||||
profiles, built-ins, schemas, and framework defaults.
|
||||
|
||||
For programmatic profiles,
|
||||
[`OpenAICompatibleProfile`](../../profiles.go) converts ordinary
|
||||
OpenAI-compatible settings into a value accepted by `WithProfiles`.
|
||||
|
||||
### Inspect A Profile Before Prompt Work
|
||||
|
||||
Use [`Engine.InspectProfile`](../../engine.go) to validate one configured
|
||||
profile without constructing a synthetic prompt or placeholder inputs. It
|
||||
resolves the profile's effective target but does not prepare or execute a
|
||||
prompt:
|
||||
|
||||
```go
|
||||
inspection, err := engine.InspectProfile(ctx, profileID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
target := inspection.EffectiveModelParams
|
||||
if target.APIKeyEnv != "" {
|
||||
// Apply application policy for the named environment variable.
|
||||
} else if inspection.APIKeyRequired {
|
||||
// Arrange a direct credential before later execution.
|
||||
}
|
||||
```
|
||||
|
||||
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.
|
||||
|
||||
### Set A Per-Run Session And Reasoning
|
||||
|
||||
Supply a direct session ID when one prompt should be correlated with a
|
||||
@@ -124,35 +244,86 @@ providers. The
|
||||
[`RunRequest` and `ExecutionTargetOverride` GoDoc](../../types.go) owns the
|
||||
exact normalization, precedence, error, copying, and exposure contract.
|
||||
|
||||
### Register A Custom Backend
|
||||
### Configure A Local OpenAI-Compatible Endpoint
|
||||
|
||||
Register a reusable OpenAI-compatible connection once, then select it from a
|
||||
profile. This local backend limits model generation to two simultaneous calls;
|
||||
because `QueueCapacity` is omitted, the engine admits up to 1024 additional
|
||||
calls waiting behind them:
|
||||
Choose the smallest configuration that fits how the endpoint will be reused.
|
||||
|
||||
#### Use An Endpoint-Only Profile
|
||||
|
||||
Put the endpoint directly on an in-memory profile when only that profile needs
|
||||
it and shared backend identity or capacity policy is unnecessary:
|
||||
|
||||
```go
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: "prompts",
|
||||
},
|
||||
promptkit.WithBackend(promptkit.Backend{
|
||||
ID: "local",
|
||||
Endpoint: "http://localhost:8000/v1",
|
||||
APIKeyEnv: "LOCAL_LLM_API_KEY",
|
||||
ConcurrencyLimit: 2,
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "local-summary",
|
||||
Endpoint: "http://localhost:8000/v1",
|
||||
Model: "example-model",
|
||||
}),
|
||||
)
|
||||
```
|
||||
|
||||
Endpoint-only profiles have an empty backend ID and remain unrestricted by
|
||||
backend capacity policy.
|
||||
|
||||
#### Use The Conventional Local Backend
|
||||
|
||||
Use `LocalBackend` when profiles should share the conventional `local`
|
||||
identity, endpoint, and concurrency limit:
|
||||
|
||||
```go
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: "prompts",
|
||||
},
|
||||
promptkit.WithBackend(
|
||||
promptkit.LocalBackend("http://localhost:8000/v1", 2),
|
||||
),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "local-summary",
|
||||
BackendID: "local",
|
||||
BackendID: promptkit.BackendLocal,
|
||||
Model: "example-model",
|
||||
}),
|
||||
)
|
||||
```
|
||||
|
||||
The helper is explicit: it does not pre-register a backend or read environment
|
||||
variables. Supplying a positive limit leaves queue capacity omitted, so normal
|
||||
backend registration selects the existing default waiting capacity of 1024.
|
||||
The returned value still enters the engine through `WithBackend`.
|
||||
|
||||
#### Configure A Complete Backend
|
||||
|
||||
Use a keyed `Backend` value for authentication, extra request parameters, an
|
||||
explicit queue capacity, a custom ID, or multiple local endpoints:
|
||||
|
||||
```go
|
||||
noWaiting := 0
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: "prompts",
|
||||
},
|
||||
promptkit.WithBackend(promptkit.Backend{
|
||||
ID: "local-gpu",
|
||||
Endpoint: "http://gpu-host:8000/v1",
|
||||
APIKeyEnv: "LOCAL_GPU_API_KEY",
|
||||
ExtraParams: map[string]any{"provider_option": "enabled"},
|
||||
ConcurrencyLimit: 2,
|
||||
QueueCapacity: &noWaiting,
|
||||
}),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "gpu-summary",
|
||||
BackendID: "local-gpu",
|
||||
Model: "example-model",
|
||||
}),
|
||||
)
|
||||
```
|
||||
|
||||
Use distinct custom IDs when registering multiple local endpoints.
|
||||
Registrations belong to one engine and custom IDs cannot replace built-ins.
|
||||
The [`Backend` and `WithBackend` GoDoc](../../backends.go) defines validation,
|
||||
copying, uniqueness, exact concurrency-field semantics, and request-default
|
||||
behavior.
|
||||
The [`Backend`, `LocalBackend`, and `WithBackend` GoDoc](../../backends.go)
|
||||
defines exact construction, validation, copying, uniqueness, concurrency, and
|
||||
request-default behavior.
|
||||
|
||||
Both file-backed and in-memory profiles select a registration through
|
||||
`backend` or `Profile.BackendID`. Profile and request endpoint overrides retain
|
||||
@@ -173,8 +344,8 @@ zero:
|
||||
```go
|
||||
noWaiting := 0
|
||||
backend := promptkit.Backend{
|
||||
ID: "local",
|
||||
Endpoint: "http://localhost:8000/v1",
|
||||
ID: "local-gpu",
|
||||
Endpoint: "http://gpu-host:8000/v1",
|
||||
ConcurrencyLimit: 2,
|
||||
QueueCapacity: &noWaiting,
|
||||
}
|
||||
@@ -236,6 +407,11 @@ When a limited backend has admitted all active and waiting calls, handle
|
||||
```go
|
||||
result, err := engine.Run(ctx, request)
|
||||
if errors.Is(err, promptkit.ErrCapacityExceeded) {
|
||||
var capacityErr *promptkit.CapacityError
|
||||
if errors.As(err, &capacityErr) {
|
||||
// Record capacityErr.BackendID using application-owned diagnostics.
|
||||
}
|
||||
|
||||
// Apply application policy: shed work, report overload, or retry later.
|
||||
}
|
||||
```
|
||||
@@ -243,8 +419,9 @@ if errors.Is(err, promptkit.ErrCapacityExceeded) {
|
||||
A rejected call returns no partial result and does not invoke the model
|
||||
client. Promptkit does not prescribe retries or map this error to an HTTP
|
||||
status; those choices remain with the consuming application. The
|
||||
[`Engine.Run` and error GoDoc](../../engine.go) owns exact error and
|
||||
cancellation identities.
|
||||
[`CapacityError` GoDoc](../../capacity_error.go) owns the exact typed-error
|
||||
contract, while the [`Engine.Run` and error GoDoc](../../engine.go) owns broad
|
||||
error and cancellation identities.
|
||||
|
||||
## Application Boundary
|
||||
|
||||
|
||||
@@ -45,11 +45,17 @@ For cross-cutting changes, follow every applicable row. Do not create
|
||||
placeholder documents for packages, APIs, or integrations that do not yet
|
||||
exist.
|
||||
|
||||
## Maintainer-Run Validation
|
||||
## Maintainer Validation
|
||||
|
||||
Promptkit does not currently use hosted CI. Maintainers are responsible for
|
||||
running the documented checks before accepting changes. Run the default Go
|
||||
validation from the Promptkit repository root:
|
||||
This section is the canonical local validation workflow for Promptkit. Run
|
||||
every command from the repository root before accepting a change. The test
|
||||
suite and maintained examples are deterministic, offline, and require no real
|
||||
provider credentials.
|
||||
|
||||
### Tests, Analysis, Build, And Examples
|
||||
|
||||
Run the ordinary and race-enabled suites, static analysis, the build, and both
|
||||
maintained consumer examples:
|
||||
|
||||
```sh
|
||||
go test ./...
|
||||
@@ -57,85 +63,171 @@ go test -race ./...
|
||||
go vet ./...
|
||||
go build ./...
|
||||
go run ./examples/go-library/prepare
|
||||
go run ./examples/go-library/run
|
||||
```
|
||||
|
||||
Check formatting across every tracked Go file:
|
||||
Both examples must exit successfully. Review their JSON output: preparation
|
||||
must report the selected offline prompt, profile, model, and message count;
|
||||
execution must report the deterministic generated output, successful
|
||||
validation, selected offline model, and usage. Neither command may contact a
|
||||
provider or require credentials.
|
||||
|
||||
### Go Formatting
|
||||
|
||||
Check every tracked Go file. The final command must succeed and the captured
|
||||
list must be empty:
|
||||
|
||||
```sh
|
||||
gofmt -l $(git ls-files '*.go')
|
||||
unformatted=$(
|
||||
git ls-files '*.go' |
|
||||
while IFS= read -r go_file
|
||||
do
|
||||
gofmt -l "$go_file"
|
||||
done
|
||||
)
|
||||
test -z "$unformatted"
|
||||
```
|
||||
|
||||
The formatting command must produce no paths. Follow every added or changed
|
||||
Markdown link and confirm its target exists. Finally, check whitespace:
|
||||
### Local Markdown Links
|
||||
|
||||
Use the Python standard library to verify every repository-relative Markdown
|
||||
target and local heading fragment. The check is offline and prints nothing on
|
||||
success:
|
||||
|
||||
```sh
|
||||
python3 - <<'PY'
|
||||
from pathlib import Path
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
from urllib.parse import unquote
|
||||
|
||||
root = Path.cwd().resolve()
|
||||
markdown_files = [
|
||||
root / name
|
||||
for name in subprocess.check_output(
|
||||
["git", "ls-files", "*.md"], text=True
|
||||
).splitlines()
|
||||
]
|
||||
link_pattern = re.compile(r"!?\[[^]]*\]\(([^)]+)\)")
|
||||
heading_pattern = re.compile(r"^#{1,6}\s+(.+?)\s*#*\s*$")
|
||||
scheme_pattern = re.compile(r"^[a-z][a-z0-9+.-]*:", re.IGNORECASE)
|
||||
|
||||
|
||||
def markdown_lines(path):
|
||||
in_fence = False
|
||||
fence = ""
|
||||
for line in path.read_text(encoding="utf-8").splitlines():
|
||||
stripped = line.lstrip()
|
||||
marker = stripped[:3]
|
||||
if marker in {"```", "~~~"}:
|
||||
if not in_fence:
|
||||
in_fence = True
|
||||
fence = marker
|
||||
elif marker == fence:
|
||||
in_fence = False
|
||||
fence = ""
|
||||
continue
|
||||
if not in_fence:
|
||||
yield line
|
||||
|
||||
|
||||
anchor_cache = {}
|
||||
|
||||
|
||||
def anchors(path):
|
||||
if path in anchor_cache:
|
||||
return anchor_cache[path]
|
||||
found = set()
|
||||
counts = {}
|
||||
for line in markdown_lines(path):
|
||||
match = heading_pattern.match(line)
|
||||
if not match:
|
||||
continue
|
||||
heading = re.sub(r"<[^>]+>", "", match.group(1)).replace("`", "")
|
||||
base = re.sub(r"[^\w\- ]", "", heading.lower()).replace(" ", "-")
|
||||
count = counts.get(base, 0)
|
||||
counts[base] = count + 1
|
||||
found.add(base if count == 0 else f"{base}-{count}")
|
||||
anchor_cache[path] = found
|
||||
return found
|
||||
|
||||
|
||||
failures = []
|
||||
for source in markdown_files:
|
||||
text = "\n".join(markdown_lines(source))
|
||||
for match in link_pattern.finditer(text):
|
||||
target = match.group(1).strip()
|
||||
if target.startswith("<") and target.endswith(">"):
|
||||
target = target[1:-1]
|
||||
if scheme_pattern.match(target) or target.startswith("//"):
|
||||
continue
|
||||
path_text, separator, fragment = target.partition("#")
|
||||
destination = source if not path_text else source.parent / unquote(path_text)
|
||||
try:
|
||||
destination = destination.resolve()
|
||||
destination.relative_to(root)
|
||||
except ValueError:
|
||||
failures.append(f"{source.relative_to(root)}: escapes repository: {target}")
|
||||
continue
|
||||
if not destination.exists():
|
||||
failures.append(f"{source.relative_to(root)}: missing target: {target}")
|
||||
continue
|
||||
if separator and destination.suffix.lower() == ".md":
|
||||
fragment = unquote(fragment).lower()
|
||||
if fragment not in anchors(destination):
|
||||
failures.append(f"{source.relative_to(root)}: missing anchor: {target}")
|
||||
|
||||
if failures:
|
||||
print("\n".join(failures), file=sys.stderr)
|
||||
raise SystemExit(1)
|
||||
PY
|
||||
```
|
||||
|
||||
### Repository Hygiene And Review
|
||||
|
||||
Reject an active Go workspace, tracked workspace files, a vendor tree, or a
|
||||
module replacement:
|
||||
|
||||
```sh
|
||||
case "$(go env GOWORK)" in
|
||||
''|off) ;;
|
||||
*) printf '%s\n' 'an active Go workspace is not allowed' >&2; exit 1 ;;
|
||||
esac
|
||||
test -z "$(git ls-files go.work go.work.sum)"
|
||||
test ! -e vendor
|
||||
if grep -Eq '^[[:space:]]*replace([[:space:]]|\()' go.mod
|
||||
then
|
||||
printf '%s\n' 'go.mod contains a replacement' >&2
|
||||
exit 1
|
||||
fi
|
||||
```
|
||||
|
||||
Check whitespace in both unstaged and staged changes. List ignored files and
|
||||
scan tracked content for common credential forms:
|
||||
|
||||
```sh
|
||||
git diff --check
|
||||
git diff --cached --check
|
||||
test -z "$(git ls-files --others --ignored --exclude-standard)"
|
||||
credential_pattern='-----BEGIN ([A-Z0-9]+ )?PRIV''ATE KEY-----|AKI''A[0-9A-Z]{16}|gh[pousr]_[A-Za-z0-9]{36,}|sk-[A-Za-z0-9]{32,}'
|
||||
if git grep -nEI -e "$credential_pattern" -- .
|
||||
then
|
||||
printf '%s\n' 'possible credential found' >&2
|
||||
exit 1
|
||||
fi
|
||||
```
|
||||
|
||||
Documentation-only work does not require unrelated new tests, but it still
|
||||
requires link validation and `git diff --check`. Run the Go validation whenever
|
||||
documentation changes commands, examples, generated output, or another
|
||||
behavior checked by the module.
|
||||
Inspect `git status --short --untracked-files=all` and the complete diff before
|
||||
accepting a change. The status may contain only the intended source changes
|
||||
during development. Reject credentials, private keys, environment files,
|
||||
generated binaries, test or coverage output, downloaded assets, template
|
||||
residue, and any other artifact that does not belong in source control. The
|
||||
credential scan catches common forms but does not replace inspection of the
|
||||
actual change.
|
||||
|
||||
## Focused Validation
|
||||
|
||||
Use focused checks while iterating, then run the complete validation sequence
|
||||
before accepting the change. The root package supports:
|
||||
After committing the accepted change, require a clean candidate:
|
||||
|
||||
```sh
|
||||
go test .
|
||||
go vet .
|
||||
go build .
|
||||
test -z "$(git status --porcelain)"
|
||||
```
|
||||
|
||||
Filter tests by name without assuming a fixed internal package layout:
|
||||
|
||||
```sh
|
||||
go test ./... -run 'TestName'
|
||||
```
|
||||
|
||||
Replace `TestName` with a useful regular expression. Target only paths that
|
||||
exist, and consult the internal component overview for their owning
|
||||
documentation. A filtered or package-specific run does not replace the
|
||||
complete repository validation.
|
||||
|
||||
## Coordinated Work With Scriptorium
|
||||
|
||||
Promptkit and Scriptorium must remain independently valid. For temporary local
|
||||
integration, use either a Go workspace outside both repositories or an
|
||||
uncommitted replacement in the consuming module.
|
||||
|
||||
If the repositories are sibling directories, run the workspace commands from
|
||||
their parent directory:
|
||||
|
||||
```sh
|
||||
go work init ./promptkit ./scriptorium
|
||||
go work sync
|
||||
```
|
||||
|
||||
Use the workspace only for coordinated local checks. From the same parent
|
||||
directory, remove it when finished:
|
||||
|
||||
```sh
|
||||
rm -f go.work go.work.sum
|
||||
```
|
||||
|
||||
Alternatively, from the Scriptorium repository root, temporarily point its
|
||||
Promptkit dependency at the sibling checkout:
|
||||
|
||||
```sh
|
||||
go mod edit -replace gitea.maximumdirect.net/eric/promptkit=../promptkit
|
||||
```
|
||||
|
||||
After coordinated checks, remove the replacement and reconcile module
|
||||
metadata:
|
||||
|
||||
```sh
|
||||
go mod edit -dropreplace gitea.maximumdirect.net/eric/promptkit
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Never commit `go.work`, `go.work.sum`, or a local filesystem `replace`
|
||||
directive. Before committing in either repository, inspect its module files and
|
||||
working tree independently. Published consumer versions must depend on a tagged
|
||||
Promptkit version, not a workspace, local replacement, or unpublished commit.
|
||||
|
||||
@@ -9,9 +9,10 @@ explains how to select these sources and invoke the engine. The
|
||||
owns the resulting outbound wire behavior.
|
||||
|
||||
Prompt and profile sources recursively discover files ending in `.yaml` or
|
||||
`.yml`. YAML decoding is strict: unknown fields are errors for the selected
|
||||
definition. Definitions are selected by their YAML `id`, not their file name
|
||||
or directory.
|
||||
`.yml`. Each prompt-definition and profile file contains exactly one YAML
|
||||
document; comments and trailing whitespace are allowed. YAML decoding is
|
||||
strict: unknown fields are errors for the selected definition. Definitions are
|
||||
selected by their YAML `id`, not their file name or directory.
|
||||
|
||||
## Prompt Definitions
|
||||
|
||||
@@ -58,6 +59,11 @@ When a request omits a version, the selected prompt ID must identify exactly
|
||||
one definition. When it supplies a version, the ID and version pair must be
|
||||
unique.
|
||||
|
||||
Exact prompt inspection uses this same configured source, strict decoding,
|
||||
referenced content-file resolution, and ID/version selection. It reports the
|
||||
selected definition's declared metadata without changing the prompt format or
|
||||
executing the definition.
|
||||
|
||||
### Inputs
|
||||
|
||||
Each `inputs` item has these fields:
|
||||
@@ -81,9 +87,16 @@ Each message has a non-empty `role` and exactly one of:
|
||||
- `content`, containing an inline Go template; or
|
||||
- `content_file`, naming a file whose contents are the Go template.
|
||||
|
||||
For directory and `fs.FS` prompt sources, `content_file` resolves relative to
|
||||
the prompt file and remains within the source root. `WithPromptFile` also
|
||||
resolves it relative to that file.
|
||||
`content_file` must be a relative path. It resolves from the directory that
|
||||
contains the prompt file and must remain within the configured prompt source
|
||||
root; parent components are allowed only when the resolved target remains
|
||||
inside that root. Absolute paths and paths that escape the root are rejected.
|
||||
Operating-system directory and single-file sources also reject symlink targets
|
||||
outside the root, while injected `fs.FS` sources apply containment in that
|
||||
filesystem's relative path namespace. For `WithPromptFile`, the source root is
|
||||
the directory containing the selected prompt file. Promptkit uses the parsed
|
||||
path text exactly after checking separately that it is not blank, so leading
|
||||
and trailing whitespace can name real filesystem entries.
|
||||
|
||||
Request variables are the template data, so a variable named `audience` is
|
||||
referenced as `{{.audience}}`. The `{{input "note"}}` helper renders the body
|
||||
@@ -151,9 +164,9 @@ extra_params:
|
||||
|
||||
| Field | Required | Meaning |
|
||||
| --- | --- | --- |
|
||||
| `id` | yes | Non-empty profile identifier. IDs must be unique within one source. |
|
||||
| `backend` | unless `endpoint` is present | Backend registry ID. It is trimmed and registry membership is checked when the profile is prepared. |
|
||||
| `endpoint` | unless `backend` is present | Non-empty OpenAI-compatible base URL, including an API version path when required. When both connection fields are present, this overrides the backend endpoint without changing backend identity. |
|
||||
| `id` | yes | Profile identifier, trimmed before selection and publication. It must be non-empty after trimming and unique within one source after normalization. |
|
||||
| `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. |
|
||||
| `temperature` | no | Number from 0 through 2. |
|
||||
| `max_tokens` | no | Integer zero or greater. |
|
||||
@@ -178,30 +191,33 @@ GoDoc.
|
||||
objects with string keys. Keys must be non-empty. With the built-in client,
|
||||
they also cannot collide with the standard fields listed in the
|
||||
[outbound request contract](integrations/openai-compatible-chat.md#request-body).
|
||||
Excessively deep or large JSON-shaped values are rejected for safety.
|
||||
|
||||
### Defaults And Overrides
|
||||
|
||||
Execution settings resolve in this order:
|
||||
|
||||
1. framework defaults;
|
||||
1. the framework timeout baseline;
|
||||
2. the selected backend, when the profile names one;
|
||||
3. the selected profile; and
|
||||
4. request `ExecutionTargetOverride` values.
|
||||
|
||||
The framework defaults are:
|
||||
The framework baseline is:
|
||||
|
||||
| Setting | Default |
|
||||
| --- | --- |
|
||||
| `temperature` | `0` |
|
||||
| `max_tokens` | `0` |
|
||||
| `top_p` | `1` |
|
||||
| `temperature` | Unspecified and omitted from compatible provider requests unless a profile or runtime override selects it. |
|
||||
| `max_tokens` | Unspecified and omitted from compatible provider requests unless a profile or runtime override selects it. |
|
||||
| `top_p` | Unspecified and omitted from compatible provider requests unless a profile or runtime override selects it. |
|
||||
| `timeout_seconds` | `600` |
|
||||
|
||||
Numeric zero in a file or in-memory profile means that the profile does not
|
||||
replace the framework default. Numeric request overrides use pointers, so an
|
||||
explicit zero is preserved. In particular, an explicit request
|
||||
`timeout_seconds` of zero disables the per-generation deadline while leaving
|
||||
the caller context and transport timeout intact.
|
||||
Numeric zero in a file or in-memory profile does not select a numeric value.
|
||||
For `temperature`, `max_tokens`, and `top_p`, it leaves the provider control
|
||||
unspecified. For `timeout_seconds`, it retains the framework deadline. Numeric
|
||||
request overrides use pointers, so an explicit zero is retained and sent to
|
||||
compatible providers. In particular, an explicit request `timeout_seconds` of
|
||||
zero disables the per-generation deadline while leaving the caller context and
|
||||
transport timeout intact.
|
||||
|
||||
Non-empty profile strings replace backend defaults, and non-empty request
|
||||
strings replace both. Request reasoning is the exception: a nil
|
||||
@@ -219,18 +235,25 @@ defines how the effective settings are serialized.
|
||||
### Source And Profile Precedence
|
||||
|
||||
An explicit request profile ID takes precedence over the prompt's
|
||||
`default_profile`. If neither is present, preparation fails.
|
||||
`default_profile`. If neither is present, preparation fails. Exact profile
|
||||
inspection instead takes one explicit profile ID and does not use a prompt
|
||||
default.
|
||||
|
||||
Profile sources resolve matching IDs in this order:
|
||||
|
||||
1. in-memory profiles supplied with `WithProfiles`;
|
||||
2. a profile file, `fs.FS`, or configured profile directory; and
|
||||
3. embedded built-in profiles.
|
||||
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.
|
||||
|
||||
A higher-precedence source falls back only when the profile is absent. An
|
||||
invalid matching profile is an error and does not fall back. In-memory
|
||||
`Profile` values follow the same ranges as YAML profiles. They use
|
||||
`APIKeyRequired` for request-scoped credentials instead of `api_key_env`.
|
||||
A profile source supplies a complete definition; definitions and their fields
|
||||
are not merged across sources. A higher-precedence source falls back only when
|
||||
the requested profile ID is absent. An invalid matching profile is an error and
|
||||
does not fall back. In-memory `Profile` values follow the same ranges as YAML
|
||||
profiles. They use `APIKeyRequired` for request-scoped credentials instead of
|
||||
`api_key_env`. Preparation and exact profile inspection use this same source
|
||||
precedence.
|
||||
|
||||
## Built-In Profile Catalog
|
||||
|
||||
@@ -238,8 +261,8 @@ 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 custom or in-memory profile with the same profile ID takes
|
||||
precedence.
|
||||
values. A configured, application fallback, or in-memory profile with the same
|
||||
profile ID takes precedence.
|
||||
|
||||
| Provider | ID | Model |
|
||||
| --- | --- | --- |
|
||||
|
||||
@@ -14,10 +14,16 @@ that produce these outbound settings.
|
||||
|
||||
Generation sends an HTTP `POST` with `Content-Type: application/json`.
|
||||
Before the client is called, the engine resolves framework, backend, profile,
|
||||
and request values into one execution target. A non-empty endpoint from that
|
||||
target overrides the client's configured base URL. After trailing slashes are
|
||||
removed, `/chat/completions` is appended. Generation fails before sending when
|
||||
neither source supplies an endpoint.
|
||||
and request values into one execution target. Endpoint configuration is trimmed
|
||||
and must be an absolute HTTP or HTTPS URL with a host and without user
|
||||
information, a query, or a fragment. A non-empty endpoint from the target
|
||||
overrides the client's configured base URL. The final selected endpoint is
|
||||
validated again before transport.
|
||||
|
||||
The completion URL is composed through parsed URL path operations. Nested base
|
||||
paths are retained, repeated trailing slashes are normalized, and the result
|
||||
has exactly one appended `/chat/completions` suffix. Generation fails before
|
||||
sending when neither source supplies a valid endpoint.
|
||||
|
||||
The target's backend ID is routing metadata for prepared values, results, and
|
||||
injected clients. The built-in client does not derive the URL from that ID and
|
||||
@@ -53,8 +59,9 @@ never also sent as a session header.
|
||||
|
||||
The client conditionally includes:
|
||||
|
||||
- `temperature`, `max_tokens`, and `top_p` when non-zero or explicitly
|
||||
present;
|
||||
- `temperature`, `max_tokens`, and `top_p` only when selected by a profile or
|
||||
runtime override, including an explicit runtime zero; they are absent when
|
||||
unspecified;
|
||||
- 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,
|
||||
@@ -81,13 +88,29 @@ request fields.
|
||||
|
||||
## Response Handling
|
||||
|
||||
Any 2xx response is decoded as an OpenAI-compatible chat response. The client
|
||||
returns the first choice's non-empty message content and maps prompt,
|
||||
completion, total, cached, and cache-write token counts.
|
||||
Any 2xx response body is limited to 16 MiB (16,777,216 bytes). A larger
|
||||
declared `Content-Length` is rejected before the body is read, and streamed,
|
||||
chunked, or underreported bodies are read through the same bound with at most
|
||||
one additional byte used to detect overflow. A body exactly at the limit is
|
||||
allowed. The body is closed on every outcome and an oversized stream is not
|
||||
drained.
|
||||
|
||||
Invalid JSON, absent choices, and empty first-choice content are malformed
|
||||
responses. For a non-2xx status, the error includes the status code but never
|
||||
the provider response body.
|
||||
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.
|
||||
|
||||
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).
|
||||
|
||||
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.
|
||||
The rendered error does not include the selected endpoint, request headers,
|
||||
request content, credentials, or provider body.
|
||||
|
||||
## Timeout And Cancellation
|
||||
|
||||
@@ -102,5 +125,7 @@ Timeouts are layered:
|
||||
timeout when the supplied value is not positive.
|
||||
|
||||
The earliest applicable caller, generation, or transport deadline controls the
|
||||
request. Constructing the internal client does not mutate a supplied
|
||||
`http.Client`.
|
||||
request. Caller cancellation retains `context.Canceled`; caller, generation,
|
||||
and whole-request timeout failures retain `context.DeadlineExceeded`, together
|
||||
with the request-failure identity. Constructing the internal client does not
|
||||
mutate a supplied `http.Client`.
|
||||
|
||||
@@ -29,27 +29,37 @@ One pool owns immutable active and total limits plus mutex-protected admission
|
||||
count, active count, and ordered waiter list. Pool state exists only for the
|
||||
lifetime of its engine.
|
||||
|
||||
## Bounded Run Admission
|
||||
## Bounded Execution Admission
|
||||
|
||||
The runner asks the manager to admit a run after resolving the prompt, profile,
|
||||
selected backend, effective execution target, credentials, and output contract,
|
||||
but before schema loading, artifact loading, or rendering. Admission is
|
||||
immediate: a limited pool either reserves a slot or returns the internal
|
||||
`ErrCapacityExceeded` identity. The root facade maps that identity to the
|
||||
public error without treating it as an invalid request or generation failure.
|
||||
For ordinary `Run`, the runner asks the manager to admit after resolving the
|
||||
prompt, profile, selected backend, effective execution target, credentials, and
|
||||
output contract, but before schema loading, artifact loading, or rendering.
|
||||
`PrepareExecution` performs no admission. `RunPrepared` claims its handle,
|
||||
rechecks credential availability, and then asks the manager to admit the
|
||||
frozen backend before generation.
|
||||
|
||||
Admission is immediate: a limited pool either reserves a slot or returns only
|
||||
the internal `ErrCapacityExceeded` identity. The runner attaches the selected
|
||||
backend identity at its use-case boundary, and the root facade translates that
|
||||
typed value without treating it as an invalid request or generation failure.
|
||||
|
||||
The total admitted bound is the active-generation limit plus its configured
|
||||
waiting capacity. The returned release function is idempotent. The runner
|
||||
defers it as soon as admission succeeds and holds the lease across remaining
|
||||
preparation, initial generation, validation, every repair attempt, and all
|
||||
failure or cancellation exits. A repair is part of its original admission and
|
||||
does not reserve another bounded slot.
|
||||
defers it as soon as admission succeeds. An ordinary run holds the lease across
|
||||
remaining preparation, initial generation, validation, every repair attempt,
|
||||
and all failure or cancellation exits. Prepared execution holds the normal
|
||||
lease across generation, validation, every internal repair attempt, and all
|
||||
execution exits. A repair is part of its original admission and does not
|
||||
reserve another bounded slot.
|
||||
|
||||
## FIFO Generation Permits
|
||||
|
||||
`NewClient` wraps the engine's selected internal model client after public
|
||||
client adaptation or built-in client construction. Initial generation and the
|
||||
default repairer receive the same wrapper.
|
||||
default repairer receive the same wrapper. Their requests retain the same
|
||||
effective backend, credential, numeric-presence metadata, and structured-output
|
||||
settings, so scheduling does not change provider omission semantics between
|
||||
calls.
|
||||
|
||||
For each `Generate` call, the wrapper selects a pool from the request's
|
||||
effective backend ID. An unlimited call passes directly to the next client. A
|
||||
@@ -64,7 +74,9 @@ other backend IDs.
|
||||
|
||||
The wrapper passes generation requests, responses, and collaborator errors
|
||||
through unchanged. It owns scheduling only; the concrete model client remains
|
||||
responsible for provider transport behavior.
|
||||
responsible for provider transport behavior. The runner, rather than the
|
||||
capacity layer, sums all five usage fields from the initial response and every
|
||||
completed repair response into the successful run result.
|
||||
|
||||
## Cancellation And Release
|
||||
|
||||
@@ -90,11 +102,17 @@ unlimited admission. The
|
||||
FIFO transfer, canceled-waiter removal, grant/cancel races, independent pools,
|
||||
unlimited calls, passthrough behavior, and panic release.
|
||||
|
||||
The [runner tests](../../internal/usecase/runner_test.go) own early admission,
|
||||
lease lifetime, failure release, and shared initial/repair scheduling. The
|
||||
The [runner tests](../../internal/usecase/runner_test.go) own ordinary early
|
||||
admission, lease lifetime, failure release, and shared initial/repair
|
||||
scheduling. The
|
||||
[prepared-execution use-case tests](../../internal/usecase/prepared_execution_test.go)
|
||||
own deferred admission, credential ordering, and prepared-execution lease
|
||||
release. The
|
||||
[external package capacity tests](../../capacity_contract_test.go) own the
|
||||
assembled public-engine behavior for configured limits, capacity errors,
|
||||
endpoint identity, engine independence, and injected clients. The
|
||||
[prepared-execution contract tests](../../prepared_execution_contract_test.go)
|
||||
own the public prepared-capacity boundary. The
|
||||
[root error-boundary tests](../../errors_internal_test.go) own preservation of
|
||||
the public generation category and context identity when generation is
|
||||
canceled.
|
||||
|
||||
@@ -25,16 +25,20 @@ and request precedence. 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 validates the configured base URL and clones any supplied
|
||||
`http.Client` so Promptkit can apply its timeout default without mutating the
|
||||
caller's client. Generation then:
|
||||
Construction trims and validates a nonempty configured base URL and clones any
|
||||
supplied `http.Client` so Promptkit can apply its timeout default without
|
||||
mutating the caller's client. An empty configured base remains valid because a
|
||||
resolved request target may supply the endpoint. Generation then:
|
||||
|
||||
1. validates request-level timeout and endpoint requirements;
|
||||
1. validates shared execution-setting invariants and the final selected base
|
||||
endpoint;
|
||||
2. maps the internal request into the OpenAI-compatible chat payload;
|
||||
3. validates and merges extra parameters;
|
||||
4. resolves authentication;
|
||||
5. performs the outbound request under the applicable deadlines; and
|
||||
6. decodes the first response choice and token usage.
|
||||
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.
|
||||
|
||||
`internal/llm` owns the set of reserved OpenAI-compatible request fields used
|
||||
when validating extra parameters. Backend registration consumes the same rule
|
||||
@@ -43,6 +47,21 @@ without making the model client depend on registry configuration.
|
||||
The implementation has no retry loop, tool-call support, provider catalog,
|
||||
inbound HTTP behavior, or durable session store.
|
||||
|
||||
## Prepared Generation
|
||||
|
||||
For [`RunPrepared`](../../engine.go), the runner supplies the model client with
|
||||
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
|
||||
[`PreparedExecution` GoDoc](../../prepared_execution.go).
|
||||
|
||||
## Failure Categories
|
||||
|
||||
The package preserves distinct error identities for invalid client
|
||||
@@ -50,9 +69,28 @@ 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.
|
||||
|
||||
Caller cancellation and deadline failures during the outbound request are
|
||||
reported as request execution failures. The runner classifies these identities
|
||||
without depending on HTTP status mapping.
|
||||
Invalid nonempty configured endpoints are configuration failures. A missing or
|
||||
invalid final selected endpoint is an invalid generation request and is
|
||||
rejected before transport.
|
||||
|
||||
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).
|
||||
|
||||
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
|
||||
both available through `errors.Is` and `errors.As`, while the rendered text
|
||||
does not expose the endpoint, headers, request content, credential, transport
|
||||
detail, or provider body. Caller cancellation retains `context.Canceled`;
|
||||
caller deadlines, generation deadlines, and whole-request client timeouts
|
||||
retain `context.DeadlineExceeded`. The runner adds its generation category
|
||||
without discarding those identities or depending on HTTP status mapping.
|
||||
|
||||
## Test Ownership
|
||||
|
||||
@@ -60,7 +98,11 @@ The
|
||||
[OpenAI-compatible client tests](../../internal/llm/openai_compatible_client_test.go)
|
||||
own configuration, client cloning, deterministic deadline precedence,
|
||||
authentication, request and response mapping, malformed data, error identity,
|
||||
cancellation, and response-body suppression. The root transport contract test
|
||||
also verifies that resolved backend settings reach this client without
|
||||
serializing backend identity. All use local test servers or test transports;
|
||||
the default suite makes no live or paid provider requests.
|
||||
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.
|
||||
|
||||
@@ -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 and result values, profile construction, extension interfaces, value conversion, redacted formatting, public error mapping, and engine-local assembly. | [Package GoDoc](../../doc.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 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) |
|
||||
| `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/capacity` | Owns engine-local bounded run 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. | [Domain declarations](../../internal/domain/domain.go) |
|
||||
| `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 JSON-compatible extra-parameter trees while preserving supported concrete value types. | [JSON values](../../internal/jsonvalue/jsonvalue.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, and combines it with an optional primary repository. | [Built-in catalog](../formats.md#built-in-profile-catalog), [repository](../../internal/profile/builtin/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/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. | [Framework formats](../formats.md#schemas), [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 backend, profile, and request settings and coordinates preparation and execution across internal sources, rendering, artifact loading, generation, validation, and optional repair. | [Internal runner](runner.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) |
|
||||
|
||||
The root package assembles these internal components without exposing their
|
||||
representations. Consumers depend only on the root facade.
|
||||
|
||||
@@ -20,14 +20,43 @@ and override semantics consumed by the runner.
|
||||
profiles, backend resolution, artifacts, rendering, model generation, and
|
||||
validation. The root engine supplies one immutable registry containing the
|
||||
built-in backend and validated consumer additions, one engine-local run
|
||||
admitter, and a model client wrapped by the same capacity manager. Schema
|
||||
documents are loaded through the validator's optional schema-loader interface.
|
||||
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.
|
||||
|
||||
Each invocation carries its state in request, prepared-run, and result values.
|
||||
The runner has no durable run or session store.
|
||||
|
||||
## Shared Prompt Selection
|
||||
|
||||
The runner uses one prompt-selection and hashing boundary for ordinary
|
||||
preparation and exact prompt inspection. Preparation retains its early
|
||||
request-ID check before direct-session normalization; both operations then use
|
||||
the configured prompt repository to select one definition, load referenced
|
||||
message content, and calculate the same prompt hash.
|
||||
|
||||
Inspection stops after that structural lookup. It does not parse templates or
|
||||
touch profile, artifact, schema, renderer, validator, admission, or model
|
||||
collaborators. The root [`Engine.InspectPrompt`](../../engine.go) GoDoc owns
|
||||
the public operation's exact contract.
|
||||
|
||||
## Shared Profile Selection
|
||||
|
||||
The runner uses one profile-selection and target-resolution boundary for
|
||||
ordinary preparation and exact profile inspection. Preparation first selects a
|
||||
request profile or a prompt default; inspection begins with its required
|
||||
explicit profile ID. Both then apply the ordinary source precedence, resolve a
|
||||
named backend, and construct the effective target from framework, backend, and
|
||||
profile values.
|
||||
|
||||
Inspection stops after the resulting endpoint and model are structurally
|
||||
validated. It does not check credential availability or perform prompt,
|
||||
artifact, schema, rendering, admission, or model-client work. The root
|
||||
[`Engine.InspectProfile`](../../engine.go) GoDoc owns the public operation's
|
||||
exact contract.
|
||||
|
||||
## Shared Preparation Pipeline
|
||||
|
||||
`Prepare` and `Run` share one private preparation pipeline split at the point
|
||||
@@ -41,14 +70,16 @@ performs only the work needed to validate routing and admission:
|
||||
5. resolve application-neutral defaults, backend defaults, profile values,
|
||||
and explicit request overrides in that order;
|
||||
6. validate endpoint, model, numeric overrides, and credential requirements;
|
||||
7. resolve the effective output contract without loading its schema; and
|
||||
7. resolve and validate the effective output contract without loading its
|
||||
schema; and
|
||||
8. retain the definition, source identities, effective settings, output
|
||||
contract, and preparation start time in invocation-local state.
|
||||
|
||||
The completion phase consumes that state without reloading the prompt,
|
||||
profile, or backend:
|
||||
|
||||
1. load structured-output schema metadata when required;
|
||||
1. create one operation-local validation plan and derive structured-output
|
||||
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;
|
||||
@@ -59,7 +90,10 @@ profile, or backend:
|
||||
`Run` performs backend admission between the phases. This structure preserves
|
||||
one execution-precedence and error-ordering implementation while allowing a
|
||||
full backend pool to reject work before expensive schema, artifact, and
|
||||
rendering operations.
|
||||
rendering operations. `Prepare` discards the plan after returning its public
|
||||
metadata. `Run` retains the plan through initial and repaired-output validation
|
||||
and discards it when the operation ends. Prepared execution stores the same
|
||||
kind of plan only in its private payload.
|
||||
|
||||
Pointer-based numeric overrides preserve an explicit zero. Invalid negative or
|
||||
out-of-range values fail as invalid requests. Endpoint overrides do not change
|
||||
@@ -90,8 +124,17 @@ its `RunAdmitter` to reserve capacity for the effective backend ID. A nil
|
||||
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. Invalid generated content remains a validation
|
||||
result; an inability to perform validation is an operational error.
|
||||
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.
|
||||
|
||||
Validation preparation and execution honor cancellation at every
|
||||
Promptkit-controlled boundary and do not publish a partial plan or result.
|
||||
Schema reads are bounded and context-checked between chunks; JSON decoding,
|
||||
schema compilation, and schema execution are checked immediately before and
|
||||
after their synchronous calls. Promptkit does not move arbitrary filesystem or
|
||||
JSON Schema work to background goroutines, so an already-blocked dependency
|
||||
method must return before cancellation can take precedence over its outcome.
|
||||
|
||||
The admission lease covers completion-phase preparation, initial generation,
|
||||
validation, every repair, and every exit. It bounds accepted work without
|
||||
@@ -101,20 +144,27 @@ 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 and session ID, validation errors, prior output, and structured-output
|
||||
specification. 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. This capability remains internal and is not a public
|
||||
option.
|
||||
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.
|
||||
|
||||
A successful result includes the output artifact and raw output, validation
|
||||
state, effective session ID, prompt and rendered-prompt hashes, selected
|
||||
profile and backend, effective settings, input hashes, token usage, a generated
|
||||
run identifier, and UTC timing. The same effective session reaches initial
|
||||
generation and any repair attempt through the rendered prompt. The same
|
||||
effective target, including backend identity, reaches generation and any
|
||||
repair attempt.
|
||||
effective target and presence metadata, including backend identity and direct
|
||||
credential during execution, reaches generation and every repair attempt.
|
||||
Result usage is the field-wise sum of all five usage values from the initial
|
||||
response and every completed repair response. Final raw output, artifact, and
|
||||
validation state still come from the last candidate. A repair error returns no
|
||||
partial run result or partial usage.
|
||||
|
||||
## Failure Categories
|
||||
|
||||
@@ -124,18 +174,21 @@ validation failures. Wrapping preserves the package identities mapped by the
|
||||
public facade and retains collaborator identities where they are part of the
|
||||
internal contract.
|
||||
|
||||
Admission capacity exhaustion retains the internal capacity identity and adds
|
||||
the selected backend ID as context. It is not recategorized as an invalid
|
||||
request or generation failure, and no partial result is returned. A context
|
||||
already done at admission retains its context identity directly. Cancellation
|
||||
while waiting for an active generation permit prevents client invocation when
|
||||
it wins the grant race; the model-client boundary then preserves the context
|
||||
error through the generation-failure category. Deferred release restores the
|
||||
admission lease on preparation, generation, validation, repair, and
|
||||
cancellation failures.
|
||||
Admission capacity exhaustion retains the internal capacity identity. At the
|
||||
use-case boundary, the runner attaches the selected backend ID in an internal
|
||||
typed error, and the root facade copies that value into the public
|
||||
[`CapacityError`](../../capacity_error.go) without parsing diagnostic text. It
|
||||
is not recategorized as an invalid request or generation failure, and no
|
||||
partial result is returned. A context already done at admission retains its
|
||||
context identity directly. Cancellation while waiting for an active generation
|
||||
permit prevents client invocation when it wins the grant race; the model-client
|
||||
boundary then preserves the context error through the generation-failure
|
||||
category. Deferred release restores the admission lease on preparation,
|
||||
generation, validation, repair, and cancellation failures.
|
||||
|
||||
Other context cancellation propagates through the invoked collaborator and is
|
||||
classified by the owning operation.
|
||||
classified by the owning operation. In particular, cancellation observed by
|
||||
validation retains the context identity through the validation error category.
|
||||
An overlong direct session is an invalid request before source loading, while
|
||||
an invalid or overlong prompt session template remains a prompt-render failure.
|
||||
An unknown selected backend, or a selected backend with no configured resolver,
|
||||
@@ -147,8 +200,9 @@ The [runner tests](../../internal/usecase/runner_test.go) own preparation order,
|
||||
selection and override precedence, the two-phase boundary, early admission,
|
||||
lease lifetime and release, direct-session resolution, schema-before-generation
|
||||
behavior, hashing, generation and validation outcomes, backend propagation,
|
||||
bounded repair, shared initial/repair capacity, credentials and redaction,
|
||||
error categories, artifact metadata, usage, and timing. The
|
||||
bounded repair progression, initial/repair request parity, cumulative usage,
|
||||
shared initial/repair capacity, credentials and redaction, error categories,
|
||||
artifact metadata, and timing. The
|
||||
[capacity subsystem document](capacity.md) identifies the focused pool,
|
||||
waiter, and wrapped-client tests.
|
||||
|
||||
|
||||
@@ -12,9 +12,26 @@ validation modes, built-in catalog, and source precedence.
|
||||
|
||||
## Prompt Definitions
|
||||
|
||||
`internal/promptdef` discovers YAML deterministically, decodes and validates
|
||||
definitions, selects an ID and optional version, and resolves file-backed
|
||||
message content within the selected operating-system or `fs.FS` source.
|
||||
`internal/promptdef` uses one source-neutral flow for prompt selection and
|
||||
normalization. That flow scans normalized YAML ID and version metadata,
|
||||
requires one strictly decoded document per file, classifies errors for the
|
||||
selected definition, detects duplicates, and normalizes the exact match.
|
||||
Small operating-system and `fs.FS` adapters own discovery, byte reads, display
|
||||
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.
|
||||
|
||||
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
|
||||
prompt file's containing directory as its root. Every content path must be
|
||||
relative and is opened from its exact parsed text after a separate blank check;
|
||||
contained parent components and whitespace-bearing names remain valid.
|
||||
|
||||
Exact prompt inspection performs one point-in-time lookup through that same
|
||||
repository and validates referenced message content before returning declared
|
||||
metadata. It does not parse templates or read profile, input, or schema
|
||||
sources, and it does not retain the definition for a later execution.
|
||||
|
||||
Its package tests own prompt selection, strict decoding, definition validation,
|
||||
duplicate detection, and source containment:
|
||||
@@ -23,29 +40,56 @@ 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`. It supports a primary repository
|
||||
with fallback only when the primary reports that a profile is absent. Strict
|
||||
YAML decoding recognizes the optional `backend` field, trims its value, and
|
||||
requires a model plus at least one non-blank backend or endpoint. Loading does
|
||||
not check registry membership because the available registry belongs to the
|
||||
assembled engine; the runner checks membership during preparation.
|
||||
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/builtin` embeds the maintained built-in profile catalog and
|
||||
can place a caller-selected repository ahead of that 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 behavior is owned by the
|
||||
[profile repository tests](../../internal/profile/repository_test.go), while
|
||||
catalog completeness, the backend-selection invariant, duplicate IDs, and
|
||||
overlay behavior are owned by the
|
||||
The overlay repository consults the next repository only when the
|
||||
higher-precedence repository reports that a profile is absent. A reliably
|
||||
selected malformed profile stops fallback, while an unrelated malformed file
|
||||
does not become authoritative through its filename. Loading does not check
|
||||
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.
|
||||
|
||||
Exact profile inspection performs one point-in-time lookup through those
|
||||
profile sources and checks the resolved target without reading prompt, input,
|
||||
or schema sources. It does not retain that lookup for a later execution.
|
||||
|
||||
`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).
|
||||
|
||||
## Ordinary Artifacts
|
||||
|
||||
`internal/artifact` resolves inline references and unrestricted,
|
||||
caller-selected file paths. It copies content into an artifact, records
|
||||
metadata and a content hash, applies a content-type fallback, and honors
|
||||
context cancellation.
|
||||
`internal/artifact` accepts explicitly typed inline references even when their
|
||||
body is empty. It also resolves unrestricted, caller-selected paths only when
|
||||
they identify regular operating-system files, checking that condition before
|
||||
and after opening the file. It copies content into an artifact, records
|
||||
metadata and an opaque content-equality value, and applies a content-type
|
||||
fallback.
|
||||
|
||||
Regular files are read synchronously in bounded chunks. Cancellation is
|
||||
checked before opening, before and after every read, and before publishing the
|
||||
artifact, so a canceled read never publishes partial content. The ordinary
|
||||
reader does not detach file reads into background goroutines.
|
||||
|
||||
This ordinary reader does not implement an inbound HTTP security boundary. In
|
||||
particular, it does not constrain files to an application root or impose an
|
||||
@@ -57,8 +101,17 @@ implemented reader behavior and failures.
|
||||
## Rendering
|
||||
|
||||
`internal/prompt` renders definition messages as Go templates using named
|
||||
artifacts and variables. It carries message roles, session IDs, and cache
|
||||
control into the rendered prompt. The
|
||||
artifacts and variables. Within one render, each referenced artifact body is
|
||||
converted to text lazily and cached by input name for reuse across the session
|
||||
and every message; the cache is not shared across renders. Conversion uses
|
||||
bounded chunks and preserves the artifact bytes exactly.
|
||||
|
||||
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
|
||||
[renderer tests](../../internal/prompt/renderer_test.go) own rendering behavior.
|
||||
|
||||
## Schemas And Output Validation
|
||||
@@ -68,6 +121,35 @@ filesystem or an `fs.FS`. Invalid generated content is returned as a validation
|
||||
result; inability to load, register, or compile a schema is an operational
|
||||
error.
|
||||
|
||||
Every preparation operation creates one operation-local validation plan. None,
|
||||
basic, and JSON modes retain the effective output contract without source
|
||||
access. JSON Schema mode loads the root document once, resolves and compiles
|
||||
each transitive reference, and retains the compiled validator. Schema compiler
|
||||
resources use canonical escaped file or private-scheme URLs; loaders decode
|
||||
their paths once and enforce the configured source boundary. The
|
||||
provider-facing structured-output metadata uses the root document captured by
|
||||
the same plan.
|
||||
|
||||
Schema preparation and execution remain synchronous. Promptkit checks
|
||||
cancellation before and after source resolution, JSON decoding, compilation,
|
||||
and validation, and between bounded schema-read chunks. Once cancellation is
|
||||
observed it returns the context error without publishing a partial plan or
|
||||
validation result, even when a compiler or validator has just returned a
|
||||
different error or a successful result. An `fs.FS` method or JSON Schema
|
||||
dependency call already in progress cannot be preempted; Promptkit waits for
|
||||
that call to return and then gives cancellation precedence. Validation does
|
||||
not detach dependency work into background goroutines.
|
||||
|
||||
`Prepare` discards its validation plan after returning metadata. `Run` retains
|
||||
its plan for initial and repaired-output validation, then discards it with the
|
||||
operation. `PrepareExecution` retains the plan in its private frozen payload;
|
||||
`RunPrepared` uses that plan without reopening prompt, profile, input, or
|
||||
schema sources or rerendering the request. A later ordinary `Run` always
|
||||
performs fresh source resolution and preparation.
|
||||
|
||||
The [validator tests](../../internal/validate/standard_validator_test.go) own
|
||||
basic, JSON, JSON Schema, source resolution, schema loading, compilation, and
|
||||
content-failure behavior.
|
||||
basic, JSON, JSON Schema, source resolution, schema loading, compilation,
|
||||
frozen-reference behavior, content-failure behavior, and the synchronous
|
||||
cancellation boundary. Prepared execution
|
||||
orchestration is owned by the
|
||||
[use-case tests](../../internal/usecase/prepared_execution_test.go).
|
||||
|
||||
@@ -19,8 +19,8 @@ results, public values, extension interfaces, profiles, and error sentinels.
|
||||
|
||||
The implemented internal components consist of:
|
||||
|
||||
- `internal/domain`, which owns framework data values shared by later internal
|
||||
components;
|
||||
- `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;
|
||||
- `internal/capacity`, which owns engine-local bounded run admission and
|
||||
@@ -92,6 +92,13 @@ coordinates internal components and adapts the supported public extension
|
||||
interfaces to narrow internal abstractions. Internal components must not depend
|
||||
on consumers or on Scriptorium.
|
||||
|
||||
`internal/domain` owns source-neutral invariants for values shared across
|
||||
multiple input and execution boundaries, including execution-setting bounds,
|
||||
OpenAI-compatible base endpoints, session identifiers, and output-contract
|
||||
legality. Callers retain source parsing, required-field rules, other
|
||||
source-specific normalization, defaulting, error classification, and policy
|
||||
specific to their own boundary.
|
||||
|
||||
## Repository And Consumer Boundary
|
||||
|
||||
Scriptorium is a downstream application that consumes Promptkit through
|
||||
|
||||
@@ -50,19 +50,18 @@ Examples of appropriate seams include clocks, randomness, subprocesses, remote A
|
||||
## Test execution requirements
|
||||
|
||||
Promptkit currently uses maintainer-run validation rather than hosted CI.
|
||||
Maintainers run the repository-documented test, vet, build, formatting,
|
||||
documentation-link, and repository-hygiene checks before accepting changes.
|
||||
Maintainers run the complete local workflow in the
|
||||
[development guide](../development.md#maintainer-validation) before accepting
|
||||
changes. That guide is the canonical owner of exact commands, formatting,
|
||||
documentation-link validation, and repository-hygiene checks.
|
||||
Introducing hosted CI later would supplement, not silently redefine, this
|
||||
documented validation model.
|
||||
|
||||
The complete test sequence includes ordinary and race-enabled package tests.
|
||||
The maintained offline consumer workflow is also run from the repository root:
|
||||
|
||||
```sh
|
||||
go test ./...
|
||||
go test -race ./...
|
||||
go run ./examples/go-library/prepare
|
||||
```
|
||||
Maintainer validation must include ordinary and race-enabled package tests,
|
||||
static analysis, a complete build, and execution of both maintained offline
|
||||
consumer examples. The preparation example protects assembled preparation and
|
||||
inspection behavior. The execution example separately protects assembled
|
||||
`Run`, injected-client, validation, usage, and result behavior.
|
||||
|
||||
Tests in the default suite must be deterministic, offline, and independent of
|
||||
real credentials. They must not invoke paid APIs, use live network
|
||||
@@ -88,8 +87,9 @@ Use each test type where it protects a distinct risk:
|
||||
interaction, while replacing live or nondeterministic external boundaries.
|
||||
- External-package root tests exercise the public facade as a Go consumer,
|
||||
while internal package tests own focused implementation behavior.
|
||||
- The maintained offline preparation example protects one representative
|
||||
assembled consumer workflow without contacting a model provider.
|
||||
- The maintained offline preparation and execution examples protect distinct
|
||||
representative assembled consumer workflows without contacting a model
|
||||
provider.
|
||||
- Fixtures should be minimal, synthetic, versioned with the behavior they
|
||||
exercise, and free of credentials or private data.
|
||||
- Golden files are appropriate only when the complete output is intentionally
|
||||
|
||||
@@ -106,48 +106,12 @@ gitea.maximumdirect.net/eric/promptkit 1.25.5
|
||||
promptkit gitea.maximumdirect.net/eric/promptkit
|
||||
```
|
||||
|
||||
Run the complete maintainer validation required by the
|
||||
[development guide](development.md):
|
||||
|
||||
```sh
|
||||
go test ./...
|
||||
go test -race ./...
|
||||
go vet ./...
|
||||
go build ./...
|
||||
go run ./examples/go-library/prepare
|
||||
```
|
||||
|
||||
Check every tracked Go file. This command must produce no output:
|
||||
|
||||
```sh
|
||||
unformatted=$(
|
||||
git ls-files '*.go' |
|
||||
while IFS= read -r go_file
|
||||
do
|
||||
gofmt -l "$go_file"
|
||||
done
|
||||
)
|
||||
test -z "$unformatted"
|
||||
```
|
||||
|
||||
Follow every maintained Markdown link and confirm that its local or published
|
||||
target exists. Review the repository for generated binaries, test or coverage
|
||||
output, credentials, template residue, downloaded assets, and other files that
|
||||
do not belong in source control.
|
||||
|
||||
Recheck module and repository hygiene, whitespace, and the clean checkout:
|
||||
|
||||
```sh
|
||||
test -z "$(git ls-files go.work go.work.sum)"
|
||||
test ! -e vendor
|
||||
if grep -Eq '^[[:space:]]*replace([[:space:]]|\()' go.mod
|
||||
then
|
||||
printf '%s\n' 'go.mod contains a replacement' >&2
|
||||
exit 1
|
||||
fi
|
||||
git diff --check
|
||||
test -z "$(git status --porcelain)"
|
||||
```
|
||||
As a release prerequisite, run the complete
|
||||
[maintainer validation workflow](development.md#maintainer-validation) against
|
||||
the clean candidate. Do not substitute a partial command list: the development
|
||||
guide owns the tests, race checks, analysis, build, both offline examples,
|
||||
formatting, Markdown links, generated-output and credential review, and
|
||||
repository hygiene. Record the successful workflow result with the candidate.
|
||||
|
||||
## Write The Release Note
|
||||
|
||||
|
||||
@@ -90,7 +90,7 @@ engine, err := promptkit.NewEngine(
|
||||
```
|
||||
|
||||
See the
|
||||
[custom-backend consumer guide](../consumers/pkg-promptkit.md#register-a-custom-backend)
|
||||
[local-endpoint consumer guide](../consumers/pkg-promptkit.md#configure-a-local-openai-compatible-endpoint)
|
||||
for task-oriented usage. The
|
||||
[`Backend` and `WithBackend` GoDoc](../../backends.go) owns exact registration,
|
||||
validation, copying, defaulting, and uniqueness semantics. The
|
||||
|
||||
74
docs/releases/v0.3.0.md
Normal file
74
docs/releases/v0.3.0.md
Normal file
@@ -0,0 +1,74 @@
|
||||
# Promptkit v0.3.0
|
||||
|
||||
This supplemental changelog summarizes the consumer-facing changes from
|
||||
`v0.2.0` to `v0.3.0`. The annotated `v0.3.0` tag is the authoritative release
|
||||
record. Exact current contracts belong to the linked GoDoc and durable
|
||||
documentation.
|
||||
|
||||
## Summary
|
||||
|
||||
`v0.3.0` adds a concise way to register the common local OpenAI-compatible
|
||||
backend configuration:
|
||||
|
||||
- `BackendLocal` provides the conventional, non-reserved backend ID `"local"`;
|
||||
and
|
||||
- `LocalBackend` constructs an ordinary `Backend` from an endpoint and
|
||||
concurrency limit.
|
||||
|
||||
The helper is explicit and additive. It does not pre-register a backend, read
|
||||
environment variables, select a model, or replace the complete `Backend`
|
||||
configuration interface.
|
||||
|
||||
## Compatibility
|
||||
|
||||
Existing `v0.2.0` consumers require no migration. Endpoint-only profiles,
|
||||
complete custom `Backend` values, the built-in OpenRouter backend, and existing
|
||||
registrations using the literal ID `"local"` continue to work unchanged.
|
||||
|
||||
## Upgrade
|
||||
|
||||
Update the module dependency with:
|
||||
|
||||
```sh
|
||||
go get gitea.maximumdirect.net/eric/promptkit@v0.3.0
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Run the consuming project's ordinary tests and race-enabled tests after the
|
||||
upgrade.
|
||||
|
||||
## Configure A Local Backend
|
||||
|
||||
Register the convenience value through the existing `WithBackend` option and
|
||||
select it from one or more profiles:
|
||||
|
||||
```go
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{PromptDir: "prompts"},
|
||||
promptkit.WithBackend(
|
||||
promptkit.LocalBackend("http://localhost:8000/v1", 2),
|
||||
),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "local-summary",
|
||||
BackendID: promptkit.BackendLocal,
|
||||
Model: "example-model",
|
||||
}),
|
||||
)
|
||||
```
|
||||
|
||||
Use an endpoint-only profile when shared backend identity and capacity policy
|
||||
are unnecessary. Continue to use a complete keyed `Backend` value for custom
|
||||
IDs, authentication, extra request parameters, explicit queue capacity, or
|
||||
multiple local endpoints.
|
||||
|
||||
See the
|
||||
[local-endpoint consumer guide](../consumers/pkg-promptkit.md#configure-a-local-openai-compatible-endpoint)
|
||||
for task-oriented configuration choices. The
|
||||
[`BackendLocal`, `LocalBackend`, and `WithBackend` GoDoc](../../backends.go)
|
||||
owns their exact construction, registration, validation, and concurrency
|
||||
semantics.
|
||||
|
||||
## Consumer Action
|
||||
|
||||
None. Adopt the convenience constructor when it simplifies local endpoint
|
||||
configuration.
|
||||
189
docs/releases/v0.4.0.md
Normal file
189
docs/releases/v0.4.0.md
Normal file
@@ -0,0 +1,189 @@
|
||||
# Promptkit v0.4.0
|
||||
|
||||
This supplemental changelog and adoption guide summarizes the consumer-facing
|
||||
changes from `v0.3.0` to `v0.4.0`. The annotated `v0.4.0` tag is the
|
||||
authoritative release record. Exact current contracts belong to the linked
|
||||
GoDoc and durable documentation.
|
||||
|
||||
## Summary
|
||||
|
||||
`v0.4.0` adds four complementary capabilities:
|
||||
|
||||
- opaque prepared-execution handles for preparing once, inspecting safe
|
||||
details, and executing the same frozen snapshot;
|
||||
- exact prompt-definition inspection without profile resolution or execution;
|
||||
- exact profile inspection without selecting a prompt or checking credential
|
||||
availability; and
|
||||
- structured backend identity on engine admission-capacity rejection.
|
||||
|
||||
These APIs let consumers perform more precise preflight work and retain useful
|
||||
operational context without reproducing Promptkit's internal resolution logic.
|
||||
|
||||
## Compatibility
|
||||
|
||||
The release is additive for `v0.3.0` consumers. Existing uses of `Prepare`,
|
||||
`Run`, backend registration, endpoint-only profiles, local-backend helpers,
|
||||
runtime overrides, public JSON values, and error sentinels continue to work
|
||||
without migration.
|
||||
|
||||
Capacity rejection now returns a structured error while continuing to match
|
||||
`ErrCapacityExceeded` through `errors.Is`. Error-string wording and direct
|
||||
sentinel equality were not public contracts.
|
||||
|
||||
The new inspection values, capacity error, and prepared-execution handle do not
|
||||
have stable JSON representations. `PreparedExecution.Details` returns the
|
||||
existing stable `PreparedRun` value.
|
||||
|
||||
## Upgrade
|
||||
|
||||
Update the module dependency with:
|
||||
|
||||
```sh
|
||||
go get gitea.maximumdirect.net/eric/promptkit@v0.4.0
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Run the consuming project's ordinary and race-enabled tests after upgrading.
|
||||
No source migration is required.
|
||||
|
||||
## Prepare Once And Execute The Same Snapshot
|
||||
|
||||
Consumers that need to persist preparation details before generation can now
|
||||
prepare an opaque, engine-bound execution:
|
||||
|
||||
```go
|
||||
prepared, err := engine.PrepareExecution(ctx, request)
|
||||
if err != nil {
|
||||
// Handle preparation failure.
|
||||
}
|
||||
defer prepared.Discard()
|
||||
|
||||
details := prepared.Details()
|
||||
// Persist a consumer-selected, appropriately protected preparation record.
|
||||
|
||||
result, err := engine.RunPrepared(ctx, prepared)
|
||||
```
|
||||
|
||||
Preparation freezes the selected sources, rendered messages, effective
|
||||
settings, input content, structured-output metadata, and validation resources
|
||||
needed by execution. `Details` returns a fresh, caller-owned,
|
||||
credential-redacted `PreparedRun`.
|
||||
|
||||
A handle belongs to its creating engine and permits one execution attempt.
|
||||
`RunPrepared` consumes that attempt on success and on operational failure.
|
||||
`Discard` is idempotent and releases an unclaimed handle's execution-only
|
||||
state. Consumers should discard handles they will not execute, particularly
|
||||
when a direct request API key may be retained privately until claim or
|
||||
discard.
|
||||
|
||||
Prepared execution does not reserve backend admission during preparation.
|
||||
Credential availability and backend admission are checked when execution
|
||||
begins. The execution context is independent of the preparation context.
|
||||
|
||||
See the
|
||||
[prepared-execution consumer guide](../consumers/pkg-promptkit.md#prepare-now-and-execute-the-same-snapshot-later),
|
||||
the [`PreparedExecution` GoDoc](../../prepared_execution.go), and the
|
||||
[`Engine.PrepareExecution` and `Engine.RunPrepared` GoDoc](../../engine.go)
|
||||
for the exact lifecycle, ownership, cancellation, capacity, timing, and
|
||||
failure contracts.
|
||||
|
||||
## Inspect A Prompt
|
||||
|
||||
`Engine.InspectPrompt` resolves one prompt ID and optional version through the
|
||||
engine's configured prompt source:
|
||||
|
||||
```go
|
||||
inspection, err := engine.InspectPrompt(ctx, "report.summary", "")
|
||||
```
|
||||
|
||||
The result includes prompt identity, the opaque prompt hash, declared default
|
||||
profile ID, declared input metadata, and normalized output contract. It
|
||||
structurally loads the selected definition and referenced message content but
|
||||
does not resolve a profile, load schemas or artifacts, render templates,
|
||||
reserve capacity, or contact a model.
|
||||
|
||||
Use inspection for exact configuration checks and metadata discovery. Use
|
||||
`PrepareExecution` rather than relying on a prior inspection when later
|
||||
execution must freeze one exact source state, because filesystem-backed
|
||||
inspection is only a point-in-time lookup.
|
||||
|
||||
See the
|
||||
[prompt-inspection consumer guide](../consumers/pkg-promptkit.md#inspect-a-prompt-before-preparation)
|
||||
and [`Engine.InspectPrompt` GoDoc](../../engine.go) for exact selection,
|
||||
ownership, and error behavior.
|
||||
|
||||
## Inspect A Profile
|
||||
|
||||
`Engine.InspectProfile` resolves one explicit profile independently of a
|
||||
prompt:
|
||||
|
||||
```go
|
||||
inspection, err := engine.InspectProfile(ctx, "report-production")
|
||||
```
|
||||
|
||||
The result includes the resolved effective execution target and whether a
|
||||
later request must provide a direct credential. Environment-variable names may
|
||||
be reported, but inspection does not read credential values or require the
|
||||
named variable to be populated.
|
||||
|
||||
Inspection applies the engine's profile source precedence and resolves any
|
||||
selected backend. It does not load a prompt, render content, reserve capacity,
|
||||
or contact a model.
|
||||
|
||||
See the
|
||||
[profile-inspection consumer guide](../consumers/pkg-promptkit.md#inspect-a-profile-before-prompt-work)
|
||||
and [`Engine.InspectProfile` GoDoc](../../engine.go) for the exact resolution,
|
||||
credential, ownership, and error contracts.
|
||||
|
||||
## Identify Capacity-Rejected Backends
|
||||
|
||||
Calls rejected at Promptkit's bounded engine admission boundary continue to
|
||||
match `ErrCapacityExceeded`. Consumers can additionally obtain the selected
|
||||
registered backend ID without parsing diagnostic text:
|
||||
|
||||
```go
|
||||
result, err := engine.Run(ctx, request)
|
||||
if errors.Is(err, promptkit.ErrCapacityExceeded) {
|
||||
var capacityErr *promptkit.CapacityError
|
||||
if errors.As(err, &capacityErr) {
|
||||
// Record capacityErr.BackendID using application-owned diagnostics.
|
||||
}
|
||||
|
||||
// Apply application-owned overload or retry policy.
|
||||
}
|
||||
```
|
||||
|
||||
The structured error applies to `Run` and `RunPrepared` admission rejection.
|
||||
It does not represent provider throttling, quota exhaustion, cancellation
|
||||
while waiting for generation capacity, or another model-client failure.
|
||||
Promptkit does not prescribe retry timing or transport status mapping.
|
||||
|
||||
See the
|
||||
[error-handling consumer guide](../consumers/pkg-promptkit.md#handle-errors),
|
||||
the [`CapacityError` GoDoc](../../capacity_error.go), and the
|
||||
[`ErrCapacityExceeded` GoDoc](../../engine.go) for the canonical contracts.
|
||||
|
||||
## Public API Additions
|
||||
|
||||
The release adds:
|
||||
|
||||
- `Engine.PrepareExecution`;
|
||||
- `Engine.RunPrepared`;
|
||||
- `PreparedExecution`, including `Details`, `Discard`, `String`, and
|
||||
`GoString`;
|
||||
- `Engine.InspectPrompt`;
|
||||
- `PromptInspection`;
|
||||
- `PromptInputDefinition`;
|
||||
- `Engine.InspectProfile`;
|
||||
- `ProfileInspection`; and
|
||||
- `CapacityError`.
|
||||
|
||||
No public API was removed.
|
||||
|
||||
## Consumer Action
|
||||
|
||||
None. Existing `v0.3.0` workflows may upgrade without adopting the new APIs.
|
||||
|
||||
Consumers that adopt prepared execution should discard unused handles.
|
||||
Consumers that need backend-specific capacity diagnostics may add an
|
||||
`errors.As` check while retaining their existing `errors.Is` classification.
|
||||
125
docs/releases/v0.5.0.md
Normal file
125
docs/releases/v0.5.0.md
Normal file
@@ -0,0 +1,125 @@
|
||||
# Promptkit v0.5.0
|
||||
|
||||
This supplemental changelog and migration guide summarizes the consumer-facing
|
||||
changes from `v0.4.0` to `v0.5.0`. The annotated `v0.5.0` tag is the
|
||||
authoritative release record. Exact current contracts belong to the linked
|
||||
GoDoc and durable documentation.
|
||||
|
||||
## Summary
|
||||
|
||||
`v0.5.0` makes provider requests less prescriptive and adds an application
|
||||
fallback layer for profile definitions:
|
||||
|
||||
- unset optional provider controls are omitted from OpenAI-compatible request
|
||||
bodies instead of being populated with framework values; and
|
||||
- `WithFallbackProfileFS` lets an application package profile defaults that
|
||||
operators can override through the existing ordinary profile sources.
|
||||
|
||||
These changes let compatible providers apply their own model defaults while
|
||||
giving applications stable embedded profile IDs without weakening operator
|
||||
configuration precedence.
|
||||
|
||||
## Compatibility
|
||||
|
||||
The release adds one public function and removes no public declaration.
|
||||
Existing source code should continue to compile.
|
||||
|
||||
There is one intentional behavior change: when no profile or runtime override
|
||||
selects `top_p`, Promptkit no longer sends the former framework value of `1`.
|
||||
It omits `top_p` and lets the provider choose its behavior. Unset
|
||||
`temperature` and `max_tokens` are likewise omitted. Explicit nonzero profile
|
||||
values and runtime values—including explicit runtime zero values—retain their
|
||||
precedence and wire effect.
|
||||
|
||||
Consumers that relied on Promptkit always sending `top_p: 1` should add that
|
||||
value to the relevant profile or runtime override before upgrading. Consumers
|
||||
that did not rely on the implicit sampling value require no migration.
|
||||
|
||||
Application fallback profiles are opt-in. Engines that do not call
|
||||
`WithFallbackProfileFS` retain the previous profile-source behavior.
|
||||
|
||||
## Upgrade
|
||||
|
||||
Update the module dependency with:
|
||||
|
||||
```sh
|
||||
go get gitea.maximumdirect.net/eric/promptkit@v0.5.0
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Run the consuming project's ordinary and race-enabled tests after upgrading.
|
||||
If request payloads or model behavior are asserted in fixtures, review them for
|
||||
the optional-parameter omission described below.
|
||||
|
||||
## Omitted Optional Provider Controls
|
||||
|
||||
The built-in OpenAI-compatible client now includes `temperature`,
|
||||
`max_tokens`, and `top_p` only when a profile or runtime override selects the
|
||||
value. An explicit runtime zero remains present because runtime override
|
||||
pointers distinguish zero from an unspecified value.
|
||||
|
||||
Promptkit's positive generation deadline remains a framework concern and is
|
||||
not a provider request-body default. Required request fields, session IDs,
|
||||
structured output, reasoning selection, credentials, and explicit extra
|
||||
parameters retain their existing behavior.
|
||||
|
||||
See the [framework default and precedence reference](../formats.md#defaults-and-overrides),
|
||||
the [`ExecutionTargetOverride` GoDoc](../../types.go), and the
|
||||
[OpenAI-compatible request-body contract](../integrations/openai-compatible-chat.md#request-body)
|
||||
for current details.
|
||||
|
||||
## Embedded Application Fallback Profiles
|
||||
|
||||
Applications can package ordinary profile YAML in an `fs.FS` and register it
|
||||
as a fallback source:
|
||||
|
||||
```go
|
||||
//go:embed profiles/*.yaml
|
||||
var applicationProfiles embed.FS
|
||||
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||
PromptDir: "prompts",
|
||||
ProfileDir: operatorProfileDir,
|
||||
},
|
||||
promptkit.WithFallbackProfileFS(applicationProfiles, "profiles"),
|
||||
)
|
||||
```
|
||||
|
||||
Leave `operatorProfileDir` empty when no operator source is configured. A
|
||||
configured ordinary source is authoritative: a matching definition overrides
|
||||
the application fallback, while a read or validation failure remains an error
|
||||
instead of silently reaching a lower layer.
|
||||
|
||||
Profile definitions resolve in this order:
|
||||
|
||||
1. in-memory profiles supplied with `WithProfiles`;
|
||||
2. the ordinary configured source selected by `WithProfileFile`,
|
||||
`WithProfileFS`, or `Config.ProfileDir`;
|
||||
3. the application source supplied with `WithFallbackProfileFS`; and
|
||||
4. Promptkit's embedded built-in profiles.
|
||||
|
||||
Only an absent profile ID falls through. Sources provide complete profiles and
|
||||
do not merge fields. Loading remains lazy, and the new source uses the existing
|
||||
strict profile YAML and credential rules.
|
||||
|
||||
See the
|
||||
[embedded-default consumer guidance](../consumers/pkg-promptkit.md#supply-embedded-application-defaults),
|
||||
the [`WithFallbackProfileFS` GoDoc](../../engine.go), and the
|
||||
[profile source reference](../formats.md#source-and-profile-precedence) for
|
||||
current details.
|
||||
|
||||
## Public API Changes
|
||||
|
||||
The release adds:
|
||||
|
||||
- `WithFallbackProfileFS`.
|
||||
|
||||
No public declaration was removed or changed.
|
||||
|
||||
## Consumer Action
|
||||
|
||||
- Review any workflow that depended on Promptkit's implicit `top_p: 1` and
|
||||
configure the value explicitly when required.
|
||||
- Optionally adopt `WithFallbackProfileFS` when an application should package
|
||||
overridable profile defaults.
|
||||
- Run consumer tests after updating the module dependency.
|
||||
131
docs/releases/v0.6.0.md
Normal file
131
docs/releases/v0.6.0.md
Normal file
@@ -0,0 +1,131 @@
|
||||
# Promptkit v0.6.0
|
||||
|
||||
This supplemental changelog and migration guide summarizes the consumer-facing
|
||||
changes from `v0.5.0` to `v0.6.0`. The annotated `v0.6.0` tag is the
|
||||
authoritative release record. Exact current contracts belong to the linked
|
||||
GoDoc and durable documentation.
|
||||
|
||||
## Summary
|
||||
|
||||
`v0.6.0` is a broad correctness, safety, efficiency, and maintainability
|
||||
release. It does not add or remove public declarations. The release:
|
||||
|
||||
- centralizes shared execution-setting, output-contract, endpoint, and
|
||||
JSON-compatible-value rules;
|
||||
- unifies prompt repository behavior and avoids unnecessary prompt and profile
|
||||
decoding;
|
||||
- bounds consumer-controlled JSON trees and successful provider responses;
|
||||
- hardens prompt content paths, artifact files, provider URLs, JSON framing,
|
||||
and error propagation;
|
||||
- reuses compiled schema plans and rendered artifact text within an operation;
|
||||
and
|
||||
- improves cancellation behavior, prepared-value ownership, deterministic
|
||||
transport testing, and maintainer validation.
|
||||
|
||||
## Compatibility
|
||||
|
||||
No public declaration was added, removed, or changed. Ordinary valid `v0.5.0`
|
||||
configurations and requests should continue to compile and behave as before.
|
||||
|
||||
The release intentionally rejects or reports several inputs that were
|
||||
previously accepted, altered, or misclassified:
|
||||
|
||||
- execution settings must be finite, within their documented ranges, and safe
|
||||
to convert to Go durations;
|
||||
- output formats, validation modes, repair counts, and JSON Schema dependencies
|
||||
are validated consistently;
|
||||
- file-backed prompt and profile identity comes from normalized YAML metadata,
|
||||
not filenames;
|
||||
- prompt `content_file` values must be exact relative paths contained by their
|
||||
configured source root;
|
||||
- built-in file artifacts must resolve to regular files;
|
||||
- selected provider endpoints must be absolute HTTP or HTTPS URLs without user
|
||||
information, query strings, or fragments;
|
||||
- JSON documents and successful provider responses must contain exactly one
|
||||
value, and successful provider bodies are limited to 16 MiB; and
|
||||
- excessively deep or expansive JSON-compatible values fail with ordinary
|
||||
validation errors.
|
||||
|
||||
These are compatibility corrections and safety boundaries rather than new
|
||||
consumer configuration requirements. Consumers relying on an invalid or
|
||||
ambiguous input should correct that input before upgrading.
|
||||
|
||||
## Upgrade
|
||||
|
||||
Update the module dependency with:
|
||||
|
||||
```sh
|
||||
go get gitea.maximumdirect.net/eric/promptkit@v0.6.0
|
||||
go mod tidy
|
||||
```
|
||||
|
||||
Run the consuming project's ordinary and race-enabled tests after upgrading.
|
||||
Applications with custom prompt/profile sources, local provider endpoints,
|
||||
unusual artifact paths, or assertions over provider error identities should
|
||||
pay particular attention to the compatibility notes below.
|
||||
|
||||
## Source Loading And Identity
|
||||
|
||||
Prompt definitions now share one source-neutral selection and normalization
|
||||
flow across operating-system and `fs.FS` sources. YAML `id` and `version`
|
||||
metadata are authoritative; filenames do not create a second identity system.
|
||||
Only selected content bodies are loaded, malformed unrelated definitions do
|
||||
not shadow valid exact matches, and per-file read failures are reported as
|
||||
prompt-load failures rather than false absence.
|
||||
|
||||
File-backed profiles likewise use normalized YAML IDs, reuse their metadata
|
||||
read for selected strict decoding, and avoid fully decoding unrelated files.
|
||||
Selected malformed definitions remain authoritative and do not silently fall
|
||||
through to a lower-precedence source.
|
||||
|
||||
Prompt `content_file` paths are opened exactly as declared after a separate
|
||||
blank check. They must remain relative to and contained by the configured
|
||||
prompt source root, including across operating-system symlinks.
|
||||
|
||||
See the [framework source and identity reference](../formats.md) and
|
||||
[internal source overview](../internal/sources.md) for the current contracts.
|
||||
|
||||
## Validation, Cancellation, And Efficiency
|
||||
|
||||
JSON Schema documents preserve exact JSON-number representations. Schema
|
||||
resource URLs safely escape legal filesystem names, and each operation loads
|
||||
and compiles its schema graph once. `Run` and prepared execution reuse that
|
||||
operation-local plan; Promptkit does not introduce a cross-operation cache.
|
||||
|
||||
Artifact reading, rendering, schema loading, compilation, and validation now
|
||||
check cancellation at the synchronous boundaries Promptkit controls. Rendering
|
||||
memoizes each artifact's text within one render operation, while plain JSON
|
||||
validation avoids materializing an unnecessary generic tree.
|
||||
|
||||
The shared JSON-compatible-value owner now limits nesting and produced work so
|
||||
unsafe consumer-controlled structures return errors instead of risking
|
||||
unbounded recursion or allocation. See the
|
||||
[architecture policy](../policy/architecture.md) for invariant ownership and
|
||||
the [format reference](../formats.md) for validation behavior.
|
||||
|
||||
## Provider Transport Hardening
|
||||
|
||||
OpenAI-compatible endpoints are parsed and composed structurally, including
|
||||
nested base paths. Underlying transport cancellation and deadline errors remain
|
||||
discoverable with `errors.Is` through Promptkit's generation error category.
|
||||
|
||||
Successful provider bodies are read with a fixed 16 MiB bound and must contain
|
||||
exactly one JSON response object followed only by whitespace. Oversized,
|
||||
truncated, malformed, or multiply framed responses fail without returning a
|
||||
partial result. See the
|
||||
[OpenAI-compatible integration contract](../integrations/openai-compatible-chat.md)
|
||||
for the canonical request, endpoint, error, and response behavior.
|
||||
|
||||
## Public API Changes
|
||||
|
||||
None.
|
||||
|
||||
## Consumer Action
|
||||
|
||||
- Correct any configuration or request that depends on the formerly permissive
|
||||
cases described under Compatibility.
|
||||
- Confirm custom local provider endpoints are absolute HTTP or HTTPS base URLs
|
||||
without credentials, queries, or fragments.
|
||||
- Confirm prompt content paths remain within their configured source root and
|
||||
file artifacts resolve to regular files.
|
||||
- Run ordinary and race-enabled consumer tests after updating the dependency.
|
||||
@@ -1,236 +0,0 @@
|
||||
# Backend-Specific Concurrency Management
|
||||
|
||||
**Status:** Complete.
|
||||
|
||||
## Purpose
|
||||
|
||||
This roadmap defines the scope and target end state for engine-local,
|
||||
backend-specific concurrency management. It records the intended capability,
|
||||
consumer value, and important policy choices.
|
||||
|
||||
This document is planning material, not a description of current behavior.
|
||||
Current exported contracts remain owned by Go declarations and GoDoc, backend
|
||||
registration guidance by the
|
||||
[consumer guide](../consumers/pkg-promptkit.md#register-a-custom-backend), and
|
||||
implemented orchestration by the
|
||||
[internal runner document](../internal/runner.md).
|
||||
|
||||
## Motivation
|
||||
|
||||
Different model backends can sustain very different request loads. A local
|
||||
network endpoint may need a small concurrency limit, while OpenRouter can
|
||||
usually accept substantially more simultaneous work. Requiring every consumer
|
||||
to build its own semaphores and queues would duplicate routing knowledge,
|
||||
create inconsistent cancellation behavior, and make it easy for one caller to
|
||||
bypass the intended backend limit.
|
||||
|
||||
Promptkit should own this coordination because it already resolves each run to
|
||||
an engine-scoped backend identity and owns every model-generation call made by
|
||||
the runner. Consumers should continue submitting ready-to-run requests through
|
||||
the synchronous API, including concurrently from multiple goroutines, without
|
||||
implementing their own backend scheduler.
|
||||
|
||||
The buffered queue is a safety boundary, not an ordinary throughput
|
||||
restriction. Its primary purpose is to prevent a bug or unintended submission
|
||||
loop from creating an unbounded in-memory backlog.
|
||||
|
||||
## Scope
|
||||
|
||||
The feature will add optional concurrency policy to registered backends and
|
||||
coordinate `Run` calls against independent per-backend capacity pools.
|
||||
|
||||
Each policy has two distinct controls:
|
||||
|
||||
- an active-generation limit, which protects the backend from too many
|
||||
simultaneous model requests; and
|
||||
- a bounded waiting capacity, which protects the process from admitting an
|
||||
unbounded backlog.
|
||||
|
||||
Concurrency policy belongs to a backend registration. It is not a profile
|
||||
model parameter and cannot be overridden per run. Profiles select the policy
|
||||
through their backend ID, while a profile or request endpoint override remains
|
||||
in the selected backend's pool.
|
||||
|
||||
`Prepare` does not call a model and will remain outside concurrency admission.
|
||||
|
||||
## Defaults And Configuration
|
||||
|
||||
The built-in OpenRouter backend will use:
|
||||
|
||||
- an active-generation limit of 16; and
|
||||
- a waiting capacity of 1024.
|
||||
|
||||
The waiting default is intentionally generous. Reaching it should indicate
|
||||
abnormal submission pressure rather than normal application behavior.
|
||||
|
||||
Consumer-registered backends will remain unlimited unless the consumer
|
||||
configures an active-generation limit. When a consumer enables a limit and
|
||||
does not specify waiting capacity, the waiting capacity will default to 1024.
|
||||
Consumers may configure a different bounded capacity, including zero when
|
||||
they want no admitted backlog beyond the active-limit-sized run set.
|
||||
|
||||
The public representation must distinguish an omitted waiting capacity from
|
||||
an explicit zero.
|
||||
|
||||
Endpoint-only profiles have no backend registration from which to obtain
|
||||
policy and will remain unlimited. A future engine-wide or endpoint-keyed
|
||||
policy can be considered separately if consumers demonstrate that need.
|
||||
|
||||
Invalid limits or capacities will fail engine construction as invalid
|
||||
configuration. Policy values will be copied into engine-owned immutable state
|
||||
along with the rest of the backend registration.
|
||||
|
||||
## Admission And Execution Behavior
|
||||
|
||||
`Run` remains a synchronous, wait-for-result operation. Concurrent callers may
|
||||
block inside `Run` while waiting for their selected backend, then receive the
|
||||
ordinary result or error from that invocation.
|
||||
|
||||
For a configured pool, the active-generation limit plus the waiting capacity
|
||||
defines the maximum number of concurrent `Run` invocations that Promptkit will
|
||||
accept for that backend. A waiting capacity of zero therefore accepts no more
|
||||
runs than the active limit. Admission is immediate: a call either reserves one
|
||||
of those bounded slots or receives the capacity error. An accepted run may
|
||||
then wait internally for active-generation capacity.
|
||||
|
||||
For a limited backend, Promptkit will bound the number of accepted runs before
|
||||
expensive artifact loading, prompt rendering, and large defensive copies where
|
||||
practical. Lightweight prompt, profile, and backend resolution may occur first
|
||||
when it is required to identify the correct capacity pool. This pre-admission
|
||||
resolution must not become a second execution-precedence path with behavior
|
||||
that can drift from `Prepare`.
|
||||
|
||||
An accepted run retains its admission until it completes or fails. Every
|
||||
actual model-generation call for that run must separately observe the
|
||||
backend's active-generation limit. This includes:
|
||||
|
||||
- the initial generation;
|
||||
- every output-repair generation; and
|
||||
- calls made through either the built-in or an injected model client.
|
||||
|
||||
Preparation and output validation should not hold an active-generation permit.
|
||||
A repair remains part of its already-admitted run, but reacquires active
|
||||
generation capacity so repairs cannot exceed the backend limit. It must not be
|
||||
rejected merely because new runs filled the waiting queue after its initial
|
||||
generation.
|
||||
|
||||
Within one backend pool, waiting generation calls should be served in FIFO
|
||||
order, subject to canceled calls being removed. Different backend pools make
|
||||
progress independently; a saturated local backend must not consume
|
||||
OpenRouter's active or waiting capacity.
|
||||
|
||||
The feature will not promise ordering across backend pools or completion order
|
||||
among admitted runs.
|
||||
|
||||
## Capacity Failure And Cancellation
|
||||
|
||||
When a backend's bounded waiting capacity is full, a new `Run` call will fail
|
||||
promptly rather than waiting outside the bounded admission system. The public
|
||||
API will expose a recognizable capacity-exhaustion error identity distinct
|
||||
from invalid configuration, invalid requests, and model-client failures.
|
||||
Rejected calls return no partial result and do not invoke the model client.
|
||||
|
||||
Waiting within the admitted backlog or for active-generation capacity must
|
||||
honor the caller's context. Cancellation or deadline expiry will:
|
||||
|
||||
- stop waiting promptly;
|
||||
- release any admission or generation capacity held by that invocation;
|
||||
- preserve the applicable context error identity; and
|
||||
- avoid invoking the model client if cancellation wins before generation
|
||||
starts.
|
||||
|
||||
Capacity must also be released after preparation, generation, validation,
|
||||
repair, or collaborator failure. One failed or canceled run must not reduce
|
||||
the backend's future usable capacity.
|
||||
|
||||
Elapsed `Run` timing will include time spent waiting after the call is
|
||||
accepted. `PreparedRun` timing will continue to describe preparation rather
|
||||
than queue waiting.
|
||||
|
||||
## Engine And Client Boundaries
|
||||
|
||||
All pools and queued state belong to one `Engine`. Separate engines do not
|
||||
share capacity, even when they register the same backend ID or endpoint. The
|
||||
feature introduces no process-global scheduler.
|
||||
|
||||
The engine will apply policy consistently to the built-in model client and an
|
||||
injected `LLMClient`. Consumers calling their own client outside Promptkit are
|
||||
outside this boundary. Injected clients remain responsible for their internal
|
||||
thread safety and cancellation behavior.
|
||||
|
||||
Backend policy is keyed by the resolved backend ID rather than endpoint text.
|
||||
This preserves stable routing when a selected backend's endpoint is overridden
|
||||
and avoids accidentally combining unrelated registrations that happen to use
|
||||
the same URL.
|
||||
|
||||
## Queue Lifetime And Observability
|
||||
|
||||
Admission state is buffered, ephemeral, and in-process. It is not persisted
|
||||
and has no survival guarantee across engine disposal or process termination.
|
||||
Promptkit will not introduce background job ownership or require consumers to
|
||||
start or stop workers.
|
||||
|
||||
The initial feature does not require public queue-depth metrics, callbacks, or
|
||||
inspection APIs. Capacity errors and ordinary call timing provide the
|
||||
consumer-visible behavior. Operational observability can be added later
|
||||
without coupling the scheduling mechanism to an application logging or
|
||||
metrics system.
|
||||
|
||||
## Compatibility
|
||||
|
||||
Consumer-registered backends and endpoint-only profiles remain unlimited
|
||||
unless concurrency is explicitly configured, preserving their existing
|
||||
behavior.
|
||||
|
||||
The built-in OpenRouter backend will change from unlimited concurrency to a
|
||||
limit of 16 with a bounded waiting capacity of 1024. Ordinary synchronous
|
||||
calls remain unchanged, while unusually high concurrent use may now wait or
|
||||
return the capacity error. This behavioral change must be identified in the
|
||||
release notes for the version that publishes it.
|
||||
|
||||
Adding backend policy fields and a public capacity error is otherwise
|
||||
additive. The change will use a pre-`v1` minor release under Promptkit's
|
||||
[release policy](../release.md#release-model).
|
||||
|
||||
## Non-Goals
|
||||
|
||||
This scope does not include:
|
||||
|
||||
- asynchronous job handles, polling, or detached result delivery;
|
||||
- durable or cross-process queues;
|
||||
- persistence or recovery across engine or process shutdown;
|
||||
- priorities, scheduling weights, or consumer-defined fairness classes;
|
||||
- automatic retries, backoff, rate-limit interpretation, or provider quota
|
||||
discovery;
|
||||
- token-per-minute or request-per-minute rate limiting;
|
||||
- dynamic reconfiguration after engine construction;
|
||||
- per-profile or per-run concurrency overrides;
|
||||
- endpoint-keyed pooling for profiles without a backend ID;
|
||||
- process-global coordination across engines;
|
||||
- application worker lifecycle, logging, tracing, or metrics policy; or
|
||||
- changes to prompt, profile, schema, or model-provider wire formats.
|
||||
|
||||
## Target End State
|
||||
|
||||
This roadmap reaches its target end state when:
|
||||
|
||||
- each engine independently coordinates configured backend capacity;
|
||||
- the built-in OpenRouter backend allows 16 active generations and up to 1024
|
||||
waiting runs;
|
||||
- consumer backends can opt into their own active and waiting limits while
|
||||
remaining unlimited by default;
|
||||
- endpoint overrides retain the selected backend's capacity pool and
|
||||
endpoint-only profiles remain unlimited;
|
||||
- synchronous `Run` callers wait for and receive their ordinary result;
|
||||
- admission is bounded before expensive preparation work where practical;
|
||||
- every initial and repair generation observes the backend's active limit
|
||||
without serializing preparation or validation;
|
||||
- a full waiting queue returns a recognizable capacity error without invoking
|
||||
the model client;
|
||||
- cancellation and all failure paths promptly release capacity and preserve
|
||||
context error identity;
|
||||
- built-in and injected model clients receive the same scheduling behavior;
|
||||
- pools remain ephemeral, engine-scoped, and independent across backend IDs;
|
||||
and
|
||||
- current-state GoDoc, consumer, internal, and release documentation describe
|
||||
the implemented behavior once it lands.
|
||||
61
docs/roadmap/deferred.md
Normal file
61
docs/roadmap/deferred.md
Normal file
@@ -0,0 +1,61 @@
|
||||
# Deferred Feature Ideas
|
||||
|
||||
## Purpose
|
||||
|
||||
This document catalogs feature ideas that remain potentially useful but have
|
||||
been deliberately postponed. These ideas are not awaiting ordinary selection
|
||||
from the [future feature catalog](future.md); each has a stated reason to wait
|
||||
and should be reconsidered only when its trigger becomes relevant.
|
||||
|
||||
Deferred entries are not commitments, schedules, active implementation plans,
|
||||
or descriptions of current behavior. When an entry is reactivated, move it to
|
||||
`future.md` for evaluation or directly into a focused roadmap after its open
|
||||
design dependencies have been resolved.
|
||||
|
||||
## Deferred Ideas
|
||||
|
||||
### Semantic Execution-Target Fingerprints
|
||||
|
||||
**Reason for deferral:** A stable digest requires a deliberate semantic-
|
||||
equality and versioning design. Notarius can safely use conservative source
|
||||
hashes and a Promptkit release marker today, while Weatherreporter does not
|
||||
currently reuse LLM-dependent checkpoints.
|
||||
|
||||
Promptkit could expose an opaque equality value for a resolved profile and its
|
||||
effective generation target. This would let checkpointing consumers detect
|
||||
generation-affecting configuration changes without hashing YAML presentation
|
||||
or depending on Promptkit's built-in catalog layout.
|
||||
|
||||
The digest should change with semantically relevant state such as the resolved
|
||||
model, endpoint, backend routing identity, request defaults, extra parameters,
|
||||
profile generation settings, and selected built-in profile semantics. It
|
||||
should exclude credential values, concurrency and queue policy, source paths,
|
||||
comments, formatting, and other representation-only changes. Whether a
|
||||
credential environment-variable name affects equality must be decided
|
||||
explicitly. The encoding should remain opaque and internally versioned so
|
||||
Promptkit can deliberately invalidate earlier digests when its resolution
|
||||
semantics change.
|
||||
|
||||
Reconsider this idea when a downstream consumer needs Promptkit-owned
|
||||
checkpoint equality or when a broader semantic identity design is selected.
|
||||
|
||||
### Eager Source Validation
|
||||
|
||||
**Reason for deferral:** Exact prompt and profile inspection may already
|
||||
provide a sufficiently small validation surface. Experience from downstream
|
||||
adoption should establish whether an engine-wide operation would add enough
|
||||
value to justify its broader contract.
|
||||
|
||||
Promptkit could provide an explicit offline operation that discovers and
|
||||
structurally validates configured prompt, profile, and schema sources without
|
||||
model generation. The normal `NewEngine` path would remain lazy.
|
||||
|
||||
An eager operation would need coherent handling for duplicate prompt IDs and
|
||||
versions, strict YAML decoding, referenced content files, profile/backend
|
||||
membership, schema syntax and transitive references, context cancellation,
|
||||
and source-specific public errors. Credential declarations must remain
|
||||
separate from credential values; checking current environment availability,
|
||||
if supported at all, should be an explicit option and must not expose secrets.
|
||||
|
||||
Reconsider this idea after downstream use of `InspectPrompt`,
|
||||
`InspectProfile`, and fixture-based preparation demonstrates a concrete gap.
|
||||
@@ -12,6 +12,9 @@ consumer value, and important scope boundaries. Defer API design,
|
||||
implementation details, sequencing, and acceptance criteria until an idea is
|
||||
selected.
|
||||
|
||||
Ideas that have been deliberately postponed rather than left available for
|
||||
ordinary selection belong in the [deferred catalog](deferred.md).
|
||||
|
||||
## Using This Catalog
|
||||
|
||||
- Add an idea when its purpose and likely value can be stated clearly.
|
||||
@@ -23,6 +26,8 @@ selected.
|
||||
- When an idea is selected, move its active planning to a focused roadmap or,
|
||||
when it requires a durable architectural decision, an ADR. Update
|
||||
current-state documentation only when implementation lands.
|
||||
- Move an idea to `deferred.md` when maintainers decide to retain it but wait
|
||||
for a stated design dependency, demand signal, or reconsideration trigger.
|
||||
- Remove ideas that are no longer relevant. Retain a rejected idea only when
|
||||
its rationale is likely to prevent repeated reconsideration.
|
||||
|
||||
@@ -33,9 +38,33 @@ consumers.
|
||||
|
||||
## Ideas
|
||||
|
||||
No ideas are currently cataloged. Backend-specific concurrency management has
|
||||
been selected for active planning in the
|
||||
[focused concurrency roadmap](concurrency.md).
|
||||
### 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.
|
||||
|
||||
## Entry Format
|
||||
|
||||
|
||||
@@ -1,841 +0,0 @@
|
||||
# Backend-Specific Concurrency Management Implementation Plan
|
||||
|
||||
**Status:** Complete.
|
||||
|
||||
## Purpose
|
||||
|
||||
This document is the decision-complete implementation plan for
|
||||
[backend-specific concurrency management](concurrency.md). It is written for a
|
||||
coding agent that will implement each stage in order.
|
||||
|
||||
The feature roadmap owns the intended capability, consumer value, policy
|
||||
choices, compatibility decision, and target end state. This document owns the
|
||||
concrete API, internal representation, scheduling architecture, implementation
|
||||
sequence, test ownership, documentation updates, and completion gates.
|
||||
|
||||
## Implementation Rules
|
||||
|
||||
- Complete the stages in order. Keep the repository compiling and the focused
|
||||
tests passing at every stage boundary.
|
||||
- Preserve unrelated working-tree changes. In particular, `concurrency.md` and
|
||||
the removal of its source idea from `future.md` may already be uncommitted
|
||||
when implementation begins; retain both.
|
||||
- Follow every policy under `docs/policy/`, the task-specific reading guide in
|
||||
`docs/development.md`, and the target behavior in `concurrency.md`.
|
||||
- Keep the public API in the root `promptkit` package and implementation
|
||||
details under `internal/`. Do not expose scheduler types or create another
|
||||
public package.
|
||||
- Use only the Go standard library for scheduling. Do not add a queue,
|
||||
semaphore, worker-pool, or metrics dependency.
|
||||
- Preserve synchronous, wait-for-result `Run`, unrestricted `Prepare`,
|
||||
engine-local state, endpoint-only profiles, backend-selected profiles,
|
||||
backend identity through endpoint overrides, and injected `LLMClient`
|
||||
behavior.
|
||||
- Do not broaden the work into asynchronous jobs, durable queues, retries,
|
||||
rate limiting, dynamic configuration, priorities, worker lifecycle,
|
||||
endpoint-keyed pools, or public queue observability.
|
||||
- Keep all tests deterministic, bounded, offline, and race-safe. Coordinate
|
||||
concurrent tests with channels and barriers rather than timing assumptions
|
||||
or live providers.
|
||||
- Update exact GoDoc with each exported declaration change. Update durable
|
||||
current-state documents only after the corresponding behavior is
|
||||
implemented.
|
||||
- Test configurable mechanisms with small test-owned limits. Assert the exact
|
||||
OpenRouter `16` and default queue `1024` values only at the registry contract
|
||||
that owns those operational defaults.
|
||||
- Do not create a release, change a module version, or tag a commit. The final
|
||||
implementation handoff must identify the built-in OpenRouter behavior change
|
||||
for the next pre-`v1` minor release.
|
||||
|
||||
## Fixed Design
|
||||
|
||||
### Public Backend Configuration
|
||||
|
||||
Append these fields to the existing root `Backend` type in `backends.go`:
|
||||
|
||||
```go
|
||||
type Backend struct {
|
||||
// Existing fields remain unchanged and in their current order.
|
||||
|
||||
ConcurrencyLimit int
|
||||
QueueCapacity *int
|
||||
}
|
||||
```
|
||||
|
||||
Use these exact semantics:
|
||||
|
||||
| Public values | Meaning |
|
||||
| --- | --- |
|
||||
| `ConcurrencyLimit == 0`, `QueueCapacity == nil` | Unlimited backend; preserve current behavior. |
|
||||
| `ConcurrencyLimit > 0`, `QueueCapacity == nil` | Limit active generations and use the default waiting capacity of 1024. |
|
||||
| `ConcurrencyLimit > 0`, `QueueCapacity != nil` | Limit active generations and use the pointed-to capacity exactly, including zero. |
|
||||
| `ConcurrencyLimit < 0` | Invalid engine configuration. |
|
||||
| `QueueCapacity != nil` and `*QueueCapacity < 0` | Invalid engine configuration. |
|
||||
| `ConcurrencyLimit == 0` and `QueueCapacity != nil` | Invalid engine configuration because a queue without an active limit has no defined consumer value. |
|
||||
|
||||
`ConcurrencyLimit` counts simultaneous calls to the engine-owned internal
|
||||
model-client boundary for this backend. `QueueCapacity` controls additional
|
||||
accepted `Run` invocations beyond that limit. The maximum admitted runs for a
|
||||
limited backend is therefore:
|
||||
|
||||
```text
|
||||
ConcurrencyLimit + effective QueueCapacity
|
||||
```
|
||||
|
||||
Guard that addition against integer overflow during backend validation.
|
||||
Do not impose an arbitrary upper bound beyond non-negativity and overflow
|
||||
safety.
|
||||
|
||||
The `QueueCapacity` pointer exists only to distinguish omission from explicit
|
||||
zero. `WithBackend` and `NewEngine` must not retain the caller's pointer.
|
||||
`Backend` continues to have no stable JSON representation, and consumers
|
||||
remain directed to keyed literals.
|
||||
|
||||
Do not add concurrency fields to `Profile`, `ExecutionTarget`,
|
||||
`ExecutionTargetOverride`, `RunRequest`, prompt or profile files, or stable
|
||||
prepared/result JSON.
|
||||
|
||||
### Built-In And Custom Defaults
|
||||
|
||||
The backend registry owns these exact operational defaults:
|
||||
|
||||
```go
|
||||
const (
|
||||
openRouterConcurrencyLimit = 16
|
||||
defaultQueueCapacity = 1024
|
||||
)
|
||||
```
|
||||
|
||||
The built-in `openrouter` definition has a normalized concurrency limit of 16
|
||||
and queue capacity of 1024.
|
||||
|
||||
Consumer registrations remain unlimited when concurrency is omitted. For a
|
||||
consumer backend with a positive limit and omitted queue capacity, normalize
|
||||
the queue capacity to 1024. Preserve an explicitly configured zero.
|
||||
|
||||
Consumers still cannot replace the reserved `openrouter` registration.
|
||||
Endpoint-only profiles have no backend policy and remain unlimited. A selected
|
||||
backend retains its pool when a profile or request overrides only its endpoint.
|
||||
|
||||
### Internal Backend Representation
|
||||
|
||||
Extend `internal/domain.Backend` with scalar policy values and explicit
|
||||
presence rather than retaining a pointer:
|
||||
|
||||
```go
|
||||
type Backend struct {
|
||||
// Existing fields...
|
||||
ConcurrencyLimit int
|
||||
QueueCapacity int
|
||||
QueueCapacitySet bool
|
||||
}
|
||||
|
||||
type BackendCapacityPolicy struct {
|
||||
ConcurrencyLimit int
|
||||
QueueCapacity int
|
||||
}
|
||||
```
|
||||
|
||||
`WithBackend` converts the public pointer into `QueueCapacity` plus
|
||||
`QueueCapacitySet`. Registry normalization validates the combinations above,
|
||||
fills the default, and leaves every limited stored backend with
|
||||
`QueueCapacitySet == true`. Unlimited stored backends retain zero values and
|
||||
`QueueCapacitySet == false`.
|
||||
|
||||
Add this internal registry method:
|
||||
|
||||
```go
|
||||
func (r *Registry) CapacityPolicies() map[string]domain.BackendCapacityPolicy
|
||||
```
|
||||
|
||||
It returns a newly allocated map containing only limited backends. Values are
|
||||
scalars, so callers cannot mutate registry state. The built-in OpenRouter
|
||||
policy is included. `GetBackend` continues returning a defensive backend copy,
|
||||
now including normalized scalar capacity metadata.
|
||||
|
||||
Capacity policy is operational registry metadata. Do not merge it into an
|
||||
execution target or expose it to injected model clients.
|
||||
|
||||
### Public Capacity Error
|
||||
|
||||
Add this root sentinel beside the other run errors in `engine.go`:
|
||||
|
||||
```go
|
||||
var ErrCapacityExceeded = errors.New("backend capacity exceeded")
|
||||
```
|
||||
|
||||
Its GoDoc must state that it identifies a `Run` rejected because the selected
|
||||
backend has already admitted `ConcurrencyLimit + QueueCapacity` runs. It is
|
||||
not an invalid request, an LLM/provider rate-limit response, or an
|
||||
`ErrLLMGenerate` failure.
|
||||
|
||||
The internal capacity component owns a corresponding internal
|
||||
`ErrCapacityExceeded`. Add its mapping in `publicErrorFor` before the broader
|
||||
generation and invalid-request cases. The public error must preserve the
|
||||
internal error through wrapping while matching `ErrCapacityExceeded` with
|
||||
`errors.Is`.
|
||||
|
||||
A capacity rejection returns no partial result and must not invoke the
|
||||
artifact reader, renderer, schema loader, validator, or model client. Prompt,
|
||||
profile, and backend loading needed to select the pool may already have
|
||||
occurred.
|
||||
|
||||
### Internal Capacity Component
|
||||
|
||||
Add `internal/capacity` as the single owner of engine-local run admission and
|
||||
active-generation permits.
|
||||
|
||||
Use these package-level boundaries:
|
||||
|
||||
```go
|
||||
var ErrCapacityExceeded error
|
||||
|
||||
type Manager struct {
|
||||
// Private immutable pool map.
|
||||
}
|
||||
|
||||
func NewManager(
|
||||
policies map[string]domain.BackendCapacityPolicy,
|
||||
) (*Manager, error)
|
||||
|
||||
func (m *Manager) Admit(
|
||||
ctx context.Context,
|
||||
backendID string,
|
||||
) (release func(), err error)
|
||||
|
||||
func NewClient(m *Manager, next llm.Client) llm.Client
|
||||
```
|
||||
|
||||
`NewManager` copies the supplied map and creates one independent pool per
|
||||
limited backend. Defensively reject blank IDs, non-positive concurrency
|
||||
limits, negative queue capacities, or total-capacity overflow even though the
|
||||
registry normally supplies normalized values. Construction creates no worker
|
||||
goroutines.
|
||||
|
||||
An absent manager, blank backend ID, or ID absent from the policy map is
|
||||
unlimited:
|
||||
|
||||
- `Admit` succeeds with a non-nil no-op release function; and
|
||||
- the client wrapper calls the next client directly.
|
||||
|
||||
For a limited pool, `Admit` is immediate and context-aware:
|
||||
|
||||
1. return `ctx.Err()` if the context is already done;
|
||||
2. under the pool lock, compare admitted runs with
|
||||
`ConcurrencyLimit + QueueCapacity`;
|
||||
3. return an error matching internal `ErrCapacityExceeded` when full; or
|
||||
4. increment admitted runs and return an idempotent release function.
|
||||
|
||||
The release function decrements admission exactly once, even if accidentally
|
||||
called more than once. It does not release an active-generation permit; those
|
||||
permits have their own lifetime.
|
||||
|
||||
### FIFO Generation Permits
|
||||
|
||||
`NewClient` returns an internal `llm.Client` wrapper around either the built-in
|
||||
client or the public-client adapter. It must preserve requests, successful
|
||||
responses, nil responses, and collaborator error identities exactly.
|
||||
`next` must be non-nil; `NewEngine` and internal runner construction maintain
|
||||
that invariant. A nil manager returns `next` unchanged.
|
||||
|
||||
For a configured backend ID, the wrapper:
|
||||
|
||||
1. acquires one active-generation permit from the matching pool;
|
||||
2. waits in FIFO order when the active count equals `ConcurrencyLimit`;
|
||||
3. removes a canceled waiter and returns `ctx.Err()` when cancellation wins
|
||||
before the permit is granted;
|
||||
4. invokes the next client only after a permit is granted; and
|
||||
5. releases the permit with `defer` after every success, nil response,
|
||||
collaborator error, panic unwinding, or context outcome.
|
||||
|
||||
Implement FIFO and cancellation explicitly with a mutex and an ordered waiter
|
||||
list. A channel used only as a counting semaphore is insufficient because it
|
||||
does not define FIFO ordering or safe removal of canceled waiters.
|
||||
|
||||
Permit grant and cancellation must have one lock-protected linearization
|
||||
point. If cancellation removes the waiter first, do not invoke the next
|
||||
client. If grant wins first, invoke the next client with the caller's context;
|
||||
the next client may then observe cancellation normally. Never lose or
|
||||
double-release a permit in this race.
|
||||
|
||||
Releasing a permit transfers it to the oldest non-canceled waiter before
|
||||
making it generally available. Different backend pools never share admission
|
||||
or active counts.
|
||||
|
||||
The active wrapper enforces its limit even if an internal caller invokes it
|
||||
without a run admission lease. Bounded backlog is guaranteed for ordinary
|
||||
engine `Run` calls by the runner admission path; no public API exposes the
|
||||
wrapped internal client directly.
|
||||
|
||||
### Engine Assembly
|
||||
|
||||
In `NewEngine`, after constructing the validated backend registry:
|
||||
|
||||
1. obtain `backendRegistry.CapacityPolicies()`;
|
||||
2. construct one `capacity.Manager`;
|
||||
3. construct the selected base internal LLM client exactly as today;
|
||||
4. wrap that base client with `capacity.NewClient`; and
|
||||
5. pass both the wrapped client and manager-as-admitter to the runner.
|
||||
|
||||
Every `NewEngine` call constructs a distinct manager. Do not cache managers,
|
||||
pools, or policies in package globals. The wrapper must be applied after a
|
||||
public injected client is adapted to `internal/llm.Client`, so built-in and
|
||||
injected clients receive identical scheduling behavior.
|
||||
|
||||
If `NewManager` reports a defensive configuration error, make `NewEngine`
|
||||
return an error matching `ErrInvalidConfig`.
|
||||
|
||||
`Prepare` does not use the manager. An injected client remains required to be
|
||||
safe for concurrent calls because different backend pools and unlimited
|
||||
backends may still invoke it concurrently.
|
||||
|
||||
### Shared Two-Phase Preparation
|
||||
|
||||
Refactor `internal/usecase.Runner` so `Prepare` and `Run` share one preparation
|
||||
pipeline with two private phases. Do not duplicate prompt/profile/backend
|
||||
selection or execution precedence.
|
||||
|
||||
The first phase resolves only the state required before admission:
|
||||
|
||||
1. validate `PromptID`;
|
||||
2. normalize the direct session ID;
|
||||
3. load the prompt definition;
|
||||
4. hash the original prompt definition at its existing error-order position;
|
||||
5. select and load the execution profile;
|
||||
6. resolve the selected backend;
|
||||
7. resolve and validate the effective execution target and credentials; and
|
||||
8. resolve the effective output contract without loading its schema.
|
||||
|
||||
Return a private state value containing the loaded definition, normalized
|
||||
direct session, prompt-definition hash, selected profile ID, effective target,
|
||||
numeric-presence metadata, effective output contract, and preparation start
|
||||
time. Keep this value private to `internal/usecase`.
|
||||
|
||||
The second phase consumes that state and performs:
|
||||
|
||||
1. structured-output schema loading;
|
||||
2. artifact loading and input hashing;
|
||||
3. message and prompt-session rendering;
|
||||
4. direct-session application;
|
||||
5. rendered-prompt hashing; and
|
||||
6. `PreparedRun` construction and timing.
|
||||
|
||||
Preserve every existing precedence rule, error identity, direct-session
|
||||
template bypass, hash input, selected identity, copy guarantee, and timing
|
||||
field. Do not reload the prompt, profile, or backend between phases.
|
||||
|
||||
`Runner.Prepare` records its start time, runs both phases consecutively, and
|
||||
never calls admission. Its behavior and error ordering remain unchanged.
|
||||
|
||||
`Runner.Run` records its existing run start time, runs the first preparation
|
||||
phase, and then calls:
|
||||
|
||||
```go
|
||||
release, err := r.admitter.Admit(ctx, effectiveBackendID)
|
||||
```
|
||||
|
||||
Use a narrow use-case-owned interface with the same signature:
|
||||
|
||||
```go
|
||||
type RunAdmitter interface {
|
||||
Admit(context.Context, string) (func(), error)
|
||||
}
|
||||
```
|
||||
|
||||
A nil admitter means unlimited behavior for internal constructors and tests.
|
||||
On successful admission, immediately `defer release()` around the remainder of
|
||||
the run. Then run the second preparation phase, initial generation,
|
||||
validation, and all repair attempts.
|
||||
|
||||
If admission returns internal `capacity.ErrCapacityExceeded`, add useful
|
||||
backend context without changing its identity. If it returns `ctx.Err()`,
|
||||
preserve that identity directly rather than recategorizing it as invalid
|
||||
request or generation failure.
|
||||
|
||||
This refactor intentionally replaces the current literal `Run`-calls-`Prepare`
|
||||
implementation with shared private phases. Update current-state documentation
|
||||
to describe one shared pipeline rather than retaining an inaccurate call-graph
|
||||
claim.
|
||||
|
||||
### Generation And Repair Lifetime
|
||||
|
||||
The admission lease covers the entire accepted run:
|
||||
|
||||
- second-phase preparation;
|
||||
- initial generation;
|
||||
- validation;
|
||||
- every repair; and
|
||||
- all failure and cancellation exits.
|
||||
|
||||
Preparation and validation do not hold an active-generation permit. The
|
||||
wrapped client acquires a permit only around each actual `Generate` call.
|
||||
|
||||
The runner's initial generation already carries the effective backend ID in
|
||||
`GenerateRequest.Target`. Preserve that value. `RepairRequest.Target` and the
|
||||
default repairer's generated request must continue carrying the same backend
|
||||
ID, allowing each repair to reacquire the same pool's active permit.
|
||||
|
||||
When testing or constructing `NewRunnerWithRepairer`, pass the same wrapped
|
||||
client to both the runner and `NewDefaultOutputRepairer`. Do not add capacity
|
||||
state to `RepairRequest`, `ExecutionTarget`, or public generation values.
|
||||
|
||||
A repair remains within its existing admission lease. It waits for a FIFO
|
||||
active permit but never performs a second bounded admission and therefore
|
||||
cannot fail merely because later runs filled the admission capacity.
|
||||
|
||||
### Error And Cancellation Semantics
|
||||
|
||||
The required public outcomes are:
|
||||
|
||||
| Situation | Required error identity |
|
||||
| --- | --- |
|
||||
| Admission capacity is full | `ErrCapacityExceeded` only; not `ErrInvalidRequest` or `ErrLLMGenerate`. |
|
||||
| Context is done before admission succeeds | Preserve `ctx.Err()`; do not return capacity exhaustion. |
|
||||
| Context cancels while waiting for an active permit | Preserve `ctx.Err()` through the existing `ErrLLMGenerate` generation category. |
|
||||
| Wrapped client fails after permit acquisition | Preserve existing `ErrLLMGenerate` and collaborator identities. |
|
||||
| Preparation or validation fails after admission | Preserve its existing category and release admission. |
|
||||
|
||||
Maintain the existing rule that `Run` returns no partial result on any
|
||||
operational error. Do not add queue status to errors or results.
|
||||
|
||||
`RunResult.Duration` continues to start at runner entry and therefore includes
|
||||
pre-admission resolution, accepted preparation, and active-permit waiting.
|
||||
`PreparedRun.DurationMS` continues to cover only its shared preparation phases;
|
||||
it does not include later generation waiting. Capacity-rejected calls have no
|
||||
result or timing value.
|
||||
|
||||
### Ownership And Concurrency Safety
|
||||
|
||||
The registry, capacity policy map, pool map, and per-pool limits are immutable
|
||||
after engine construction. Only admission counts, active counts, and waiter
|
||||
lists are mutable and must be protected by the owning pool mutex.
|
||||
|
||||
Do not retain public queue pointers, caller request values, contexts, or
|
||||
generation requests after their call completes. A canceled waiter must be
|
||||
unlinked so its context and request cannot remain reachable from the pool.
|
||||
|
||||
Do not hold a pool mutex while:
|
||||
|
||||
- loading or rendering prompts;
|
||||
- reading artifacts or schemas;
|
||||
- invoking a model client;
|
||||
- validating output;
|
||||
- closing a waiter notification channel if the implementation could re-enter
|
||||
pool code; or
|
||||
- calling consumer code.
|
||||
|
||||
No scheduler operation may spawn a goroutine whose lifetime outlasts the
|
||||
calling `Run`. The zero steady-state goroutine count is part of the
|
||||
in-process/no-worker-lifecycle design.
|
||||
|
||||
## Test Ownership
|
||||
|
||||
Use this ownership split and avoid repeating the full policy matrix at every
|
||||
layer:
|
||||
|
||||
- `internal/backend/registry_test.go` owns normalization, validation, the exact
|
||||
OpenRouter policy, the custom default queue, explicit zero, unlimited
|
||||
omission, and policy-map copying.
|
||||
- `internal/capacity/manager_test.go` owns admission bounds, idempotent release,
|
||||
FIFO active permits, cancellation races, capacity recovery, independent
|
||||
pools, unlimited IDs, and observed peak concurrency.
|
||||
- `internal/capacity/client_test.go` owns wrapper request/response/error
|
||||
transparency and the rule that cancellation before grant does not invoke the
|
||||
next client. Combine these with manager tests if one coherent package test
|
||||
expresses the behavior more clearly.
|
||||
- `internal/usecase/runner_test.go` owns two-phase preparation parity, pool
|
||||
selection, admission before expensive work, admission release across run
|
||||
exits, `Prepare` bypass, and repair reuse of the admitted backend.
|
||||
- Root external-package tests own public configuration conversion, assembled
|
||||
engine-local behavior, endpoint-override routing, injected-client limiting,
|
||||
and public capacity/context error identities.
|
||||
- Existing model-client HTTP tests remain unchanged because scheduling does
|
||||
not alter the OpenAI-compatible wire contract.
|
||||
|
||||
Concurrency tests must use test-owned limits such as one or two and
|
||||
channel-controlled blocking clients. Record observed active and peak counts
|
||||
under a mutex or atomics. Do not use `time.Sleep` to infer queue state.
|
||||
Package-internal tests may inspect a waiter list under its mutex through a
|
||||
small test helper when necessary to establish deterministic FIFO ordering; do
|
||||
not add production metrics or hooks solely for tests.
|
||||
|
||||
Do not add separate tests for trivial scalar copies when registry or assembled
|
||||
behavior already protects them.
|
||||
|
||||
## Stage 1 — Backend Policy And Public Configuration
|
||||
|
||||
**Status:** Complete.
|
||||
|
||||
### Goal
|
||||
|
||||
Add the public and internal backend policy representation, normalize all
|
||||
configured states, and expose immutable normalized policies without changing
|
||||
runtime scheduling yet.
|
||||
|
||||
### Work
|
||||
|
||||
1. Add `ConcurrencyLimit` and `QueueCapacity` to `Backend` in `backends.go`
|
||||
with exact GoDoc for unlimited, defaulted, explicit-zero, invalid, and
|
||||
engine-scoped behavior.
|
||||
2. Convert the public queue pointer into scalar value plus presence in
|
||||
`WithBackend`; do not retain the pointer.
|
||||
3. Add the internal backend policy fields and
|
||||
`BackendCapacityPolicy` to `internal/domain/domain.go`.
|
||||
4. Add the two registry-owned constants and configure the built-in OpenRouter
|
||||
definition with 16 and 1024.
|
||||
5. Extend `normalizeBackend` with the fixed validation, defaulting, explicit
|
||||
zero, and overflow rules.
|
||||
6. Add `Registry.CapacityPolicies`, returning only limited policies in a fresh
|
||||
map.
|
||||
7. Update existing backend composite literals and assertions only where the
|
||||
new fields are relevant. Continue using keyed literals.
|
||||
|
||||
### Tests
|
||||
|
||||
1. Extend the exact built-in registry test with the OpenRouter limit and queue.
|
||||
2. Add one coherent table covering unlimited omission, default queue,
|
||||
explicit-zero queue, negative values, queue-without-limit, and total
|
||||
overflow.
|
||||
3. Extend the registry copy/normalization test to prove returned policy maps
|
||||
cannot mutate registry state.
|
||||
4. Add root coverage only if needed to prove the public pointer/presence
|
||||
conversion; do not reproduce registry validation cases at the facade.
|
||||
|
||||
### Focused Validation
|
||||
|
||||
Run:
|
||||
|
||||
```sh
|
||||
gofmt -w backends.go internal/domain/domain.go \
|
||||
internal/backend/registry.go internal/backend/registry_test.go
|
||||
go test . ./internal/backend
|
||||
go vet . ./internal/backend
|
||||
git diff --check
|
||||
```
|
||||
|
||||
Include another touched Go test file in `gofmt` only if it actually changed.
|
||||
|
||||
### Completion Gate
|
||||
|
||||
This stage is complete when every public configuration state has one normalized
|
||||
internal meaning, OpenRouter exposes exactly 16/1024, custom backends remain
|
||||
unlimited by omission, and no runtime call is scheduled yet.
|
||||
|
||||
## Stage 2 — Engine-Local Capacity Manager
|
||||
|
||||
**Status:** Complete.
|
||||
|
||||
### Goal
|
||||
|
||||
Implement and prove the bounded admission mechanism and FIFO active-generation
|
||||
client wrapper independently of runner orchestration.
|
||||
|
||||
### Work
|
||||
|
||||
1. Add `internal/capacity/manager.go` with the manager, immutable policy copy,
|
||||
per-backend pools, internal error, immediate admission, idempotent release,
|
||||
and FIFO context-aware active permits.
|
||||
2. Add `internal/capacity/client.go` with the transparent `llm.Client` wrapper.
|
||||
3. Use mutex-protected waiter state and an ordered list; explicitly resolve
|
||||
grant-versus-cancel races.
|
||||
4. Ensure unlimited and independent-pool fast paths avoid queue allocation.
|
||||
5. Do not start workers, timers, cleanup goroutines, or process-global state.
|
||||
|
||||
### Tests
|
||||
|
||||
1. Add a compact constructor-validation table for blank IDs, non-positive
|
||||
limits, negative queues, and total-capacity overflow.
|
||||
2. With a small configured policy, prove that exactly
|
||||
`limit + queueCapacity` admissions succeed, the next matches
|
||||
`ErrCapacityExceeded`, and a release permits another admission.
|
||||
3. Prove release is idempotent.
|
||||
4. Drive more blocked client calls than the active limit and assert observed
|
||||
peak concurrency never exceeds that limit.
|
||||
5. Prove FIFO order with deterministic queue-entry synchronization.
|
||||
6. Cancel the first and a middle waiter and prove they are removed, never call
|
||||
the wrapped client, and do not block later waiters.
|
||||
7. Exercise the grant/cancel race repeatedly under `go test -race`, asserting
|
||||
no permit leak or double invocation.
|
||||
8. Prove different backend IDs proceed independently and blank, unknown, or
|
||||
nil-manager paths remain unlimited.
|
||||
9. Prove request values, successful and nil responses, and collaborator errors
|
||||
pass through unchanged after permit acquisition.
|
||||
|
||||
### Focused Validation
|
||||
|
||||
Run:
|
||||
|
||||
```sh
|
||||
gofmt -w internal/capacity/manager.go \
|
||||
internal/capacity/manager_test.go \
|
||||
internal/capacity/client.go \
|
||||
internal/capacity/client_test.go
|
||||
go test ./internal/capacity
|
||||
go test -race ./internal/capacity
|
||||
go vet ./internal/capacity
|
||||
git diff --check
|
||||
```
|
||||
|
||||
If tests are combined into one file, omit the nonexistent file from `gofmt`.
|
||||
|
||||
### Completion Gate
|
||||
|
||||
This stage is complete when the standalone component enforces relational
|
||||
admission and active limits, FIFO cancellation is race-safe, separate pools
|
||||
are independent, and the wrapper is transparent apart from waiting.
|
||||
|
||||
## Stage 3 — Shared Preparation And Early Run Admission
|
||||
|
||||
**Status:** Complete.
|
||||
|
||||
### Goal
|
||||
|
||||
Refactor runner preparation into one shared two-phase pipeline and place
|
||||
bounded admission after backend resolution but before schema, artifact, and
|
||||
rendering work.
|
||||
|
||||
### Work
|
||||
|
||||
1. Add the private pre-admission preparation state and split the existing
|
||||
`Prepare` logic according to the fixed design.
|
||||
2. Make `Runner.Prepare` call both phases without an admitter.
|
||||
3. Add the `RunAdmitter` interface and runner field.
|
||||
4. Update `NewRunner` and `NewRunnerWithRepairer` to accept the optional
|
||||
admitter; update internal call sites with nil until root assembly is wired.
|
||||
5. Change `Runner.Run` to use the first phase, admit by effective backend ID,
|
||||
defer the returned release, and then use the second phase.
|
||||
6. Preserve all existing error precedence, target resolution, hashes,
|
||||
metadata, session behavior, and timing.
|
||||
7. Return capacity and context errors with the fixed identities. Do not invoke
|
||||
later collaborators after rejection.
|
||||
|
||||
### Tests
|
||||
|
||||
1. Keep the existing `Run`/`Prepare` parity coverage passing to prove the
|
||||
shared phases do not drift.
|
||||
2. Add a fake admitter that records backend IDs and release calls.
|
||||
3. Prove a backend-selected run admits with the selected ID even when the
|
||||
endpoint is overridden.
|
||||
4. Prove an endpoint-only run uses the unlimited/blank identity and that
|
||||
`Prepare` never calls admission.
|
||||
5. Reject admission and assert schema, artifact, renderer, validator, repairer,
|
||||
and LLM collaborators are not invoked.
|
||||
6. Prove admission is released after one successful run and representative
|
||||
second-phase, generation, and validation errors. Prefer a small table around
|
||||
the single `defer` invariant rather than duplicating every error test.
|
||||
7. Retain direct-session, backend precedence, credential, hashing, and repair
|
||||
tests unchanged except for constructor arguments.
|
||||
|
||||
### Focused Validation
|
||||
|
||||
Run:
|
||||
|
||||
```sh
|
||||
gofmt -w internal/usecase/runner.go \
|
||||
internal/usecase/runner_test.go
|
||||
go test ./internal/usecase
|
||||
go test -race ./internal/usecase
|
||||
go vet ./internal/usecase
|
||||
git diff --check
|
||||
```
|
||||
|
||||
### Completion Gate
|
||||
|
||||
This stage is complete when `Prepare` remains unrestricted, `Run` admits after
|
||||
one canonical routing phase and before expensive completion work, every exit
|
||||
releases admission, and existing preparation semantics remain unchanged.
|
||||
|
||||
## Stage 4 — Engine Assembly And Public Runtime Contract
|
||||
|
||||
**Status:** Complete.
|
||||
|
||||
### Goal
|
||||
|
||||
Wire one manager into each engine, schedule built-in and injected clients,
|
||||
expose the capacity error, and prove assembled runtime behavior.
|
||||
|
||||
### Work
|
||||
|
||||
1. Add public `ErrCapacityExceeded` and its exact GoDoc in `engine.go`.
|
||||
2. Map internal capacity exhaustion in `errors.go`.
|
||||
3. Construct the manager from the registry policy snapshot in `NewEngine`.
|
||||
4. Wrap the selected internal client after built-in or injected-client
|
||||
selection and pass the manager and wrapped client to the runner.
|
||||
5. Update `Engine`, `NewEngine`, `Run`, `WithLLMClient`, and `LLMClient` GoDoc
|
||||
only where concurrency, capacity, or cancellation statements change.
|
||||
6. Ensure manager-construction errors match `ErrInvalidConfig`.
|
||||
7. For internal repair coverage, construct the default repairer with the same
|
||||
wrapped client used by its runner and confirm repair target backend identity
|
||||
remains intact.
|
||||
|
||||
### Tests
|
||||
|
||||
1. Add an external-package assembled test with a small custom limit and a
|
||||
blocking injected client; assert peak generation equals or remains below
|
||||
the configured limit.
|
||||
2. With queue capacity zero, block one accepted run before generation and
|
||||
assert the next matching-backend run returns `ErrCapacityExceeded`, does not
|
||||
match `ErrInvalidRequest` or `ErrLLMGenerate`, returns no result, and never
|
||||
reaches expensive collaborators or the client.
|
||||
3. In the same or another focused workflow, prove an endpoint override remains
|
||||
in the selected backend's pool.
|
||||
4. Prove two engines with the same backend ID have independent capacity.
|
||||
5. Prove an unlimited custom backend and an endpoint-only profile preserve
|
||||
concurrent behavior.
|
||||
6. Cancel a call waiting for an active permit; assert it matches both
|
||||
`context.Canceled` and `ErrLLMGenerate`, never invokes the injected client,
|
||||
and leaves capacity reusable.
|
||||
7. Add one internal repair workflow with concurrent runs or controlled permits
|
||||
showing initial and repair generations never exceed the same backend limit
|
||||
and repairs do not perform a second admission.
|
||||
8. Extend the public error sentinel contract test with
|
||||
`ErrCapacityExceeded`.
|
||||
|
||||
Avoid a second HTTP-level concurrency suite: the capacity client tests and one
|
||||
assembled injected-client workflow already protect the shared wrapper used by
|
||||
the built-in client.
|
||||
|
||||
### Focused Validation
|
||||
|
||||
Run:
|
||||
|
||||
```sh
|
||||
gofmt -w engine.go errors.go backends.go \
|
||||
internal/usecase/runner.go internal/usecase/runner_test.go \
|
||||
engine_test.go public_contract_test.go
|
||||
go test . ./internal/backend ./internal/capacity ./internal/usecase
|
||||
go test -race . ./internal/capacity ./internal/usecase
|
||||
go vet . ./internal/backend ./internal/capacity ./internal/usecase
|
||||
git diff --check
|
||||
```
|
||||
|
||||
Add any newly created capacity files to `gofmt` when they changed in this
|
||||
stage.
|
||||
|
||||
### Completion Gate
|
||||
|
||||
This stage is complete when every engine has independent pools, limited runs
|
||||
are bounded and FIFO at generation, endpoint routing is correct, capacity and
|
||||
context errors are stable, repairs reuse admission, and both client kinds pass
|
||||
through the same wrapper.
|
||||
|
||||
## Stage 5 — Durable Documentation And Final Validation
|
||||
|
||||
**Status:** Complete.
|
||||
|
||||
### Goal
|
||||
|
||||
Move implemented contracts into their durable owners, record compatibility
|
||||
impact, and validate the complete repository.
|
||||
|
||||
### Work
|
||||
|
||||
1. Review every changed exported declaration. Ensure GoDoc is the canonical
|
||||
owner of exact field types, nil/zero semantics, defaulting, error identity,
|
||||
engine scope, concurrency safety, cancellation, and source compatibility.
|
||||
2. Update `doc.go` so its concurrency summary acknowledges backend scheduling
|
||||
while continuing to require injected collaborators to be concurrency-safe.
|
||||
3. Update `docs/consumers/pkg-promptkit.md` with task-oriented examples for:
|
||||
- a limited local backend;
|
||||
- omitted queue capacity selecting 1024;
|
||||
- explicit zero queue capacity; and
|
||||
- handling `ErrCapacityExceeded`.
|
||||
Keep exact field semantics in GoDoc rather than duplicating a full table.
|
||||
4. Add `docs/internal/capacity.md` as the durable owner of pool lifecycle,
|
||||
admission, FIFO active permits, cancellation, client wrapping, and test
|
||||
ownership.
|
||||
5. Add `internal/capacity` to `docs/internal/overview.md`.
|
||||
6. Update `docs/policy/architecture.md` to include the implemented component
|
||||
and root assembly dependency without turning policy into an API reference.
|
||||
7. Update `docs/internal/runner.md` to describe the shared two-phase
|
||||
preparation pipeline, early bounded admission, lease lifetime, generation
|
||||
permits, repairs, capacity failures, and cancellation.
|
||||
8. Review `docs/formats.md`; add only a concise link or clarification if needed
|
||||
to explain that endpoint overrides preserve backend capacity identity.
|
||||
Do not add concurrency fields to YAML.
|
||||
9. Do not change the OpenAI-compatible integration contract or
|
||||
`docs/internal/llm.md` unless implementation changes their current
|
||||
statements; scheduling is outside the provider wire contract and concrete
|
||||
model-client implementation.
|
||||
10. Record in the implementation handoff that built-in OpenRouter now limits
|
||||
active generations to 16 with queue capacity 1024 and that the release
|
||||
must be a pre-`v1` minor release. Do not edit the release procedure or
|
||||
create a tag.
|
||||
11. After every check passes, set `concurrency.md`, this implementation plan,
|
||||
and each stage status to `Complete`. Do not remove the roadmaps in the
|
||||
implementation change; lifecycle retirement follows review.
|
||||
|
||||
### Full Validation
|
||||
|
||||
Run the complete sequence from `docs/development.md`:
|
||||
|
||||
```sh
|
||||
go test ./...
|
||||
go test -race ./...
|
||||
go vet ./...
|
||||
go build ./...
|
||||
go run ./examples/go-library/prepare
|
||||
gofmt -l $(git ls-files '*.go')
|
||||
git diff --check
|
||||
```
|
||||
|
||||
The formatting command must produce no paths. Follow every added or changed
|
||||
Markdown link and confirm its target and heading exist.
|
||||
|
||||
Also inspect:
|
||||
|
||||
```sh
|
||||
git status --short
|
||||
git diff --stat
|
||||
git diff
|
||||
```
|
||||
|
||||
Confirm that:
|
||||
|
||||
- only intended backend, capacity, runner, facade, test, documentation, and
|
||||
roadmap files changed;
|
||||
- no `go.work`, `go.work.sum`, local module replacement, credential,
|
||||
generated binary, coverage output, or unrelated change was introduced;
|
||||
- the built-in OpenRouter policy is exactly 16/1024;
|
||||
- custom and endpoint-only backends remain unlimited by omission;
|
||||
- explicit queue zero is distinguishable from omission;
|
||||
- no capacity value enters execution targets, generated requests, stable JSON,
|
||||
prompt/profile YAML, or provider payloads;
|
||||
- every engine owns distinct pools with no package-global mutable state;
|
||||
- every initial and repair generation uses the active permit wrapper;
|
||||
- capacity and waiter state is released on success, error, panic unwinding,
|
||||
and cancellation;
|
||||
- concurrency tests use deterministic coordination rather than sleeps;
|
||||
- current-state documentation describes implemented behavior rather than
|
||||
referring readers to the roadmaps; and
|
||||
- the feature and implementation roadmaps contain no unresolved work marked
|
||||
complete.
|
||||
|
||||
### Completion Gate
|
||||
|
||||
The implementation is complete only when every target-end-state item in
|
||||
`concurrency.md` is implemented, race-enabled tests demonstrate the configured
|
||||
limits and cancellation safety, durable contracts no longer depend on roadmap
|
||||
prose, and the OpenRouter compatibility change is clearly reported for the
|
||||
next minor release.
|
||||
|
||||
## Implementation Handoff
|
||||
|
||||
Backend-specific capacity management is implemented and has passed the complete
|
||||
repository validation sequence. The built-in OpenRouter backend now permits 16
|
||||
active generations and a waiting capacity of 1024. Custom backends remain
|
||||
unlimited when their limit is omitted, and endpoint-only profiles remain
|
||||
unlimited.
|
||||
|
||||
Publishing this behavior requires a pre-`v1` minor release. Its release notes
|
||||
must identify that unusually high concurrent OpenRouter use can now wait or
|
||||
return `ErrCapacityExceeded`. This implementation does not change a module
|
||||
version or create a tag.
|
||||
|
||||
## Open Questions
|
||||
|
||||
None. The feature roadmap and this plan fix the public representation,
|
||||
registry defaults, admission bound, FIFO generation behavior, early-routing
|
||||
refactor, cancellation races, error identities, engine and repair lifetimes,
|
||||
test ownership, compatibility treatment, and non-goals required for
|
||||
implementation.
|
||||
71
docs/roadmap/structured-generation-errors.md
Normal file
71
docs/roadmap/structured-generation-errors.md
Normal file
@@ -0,0 +1,71 @@
|
||||
# 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.
|
||||
339
engine.go
339
engine.go
@@ -63,9 +63,10 @@ var (
|
||||
// ErrPromptRender identifies a failure to render prompt messages or the
|
||||
// session ID from the resolved inputs and variables.
|
||||
ErrPromptRender = errors.New("failed to render prompt")
|
||||
// ErrCapacityExceeded identifies a Run rejected because the selected backend
|
||||
// already admitted ConcurrencyLimit + QueueCapacity calls. It is not an
|
||||
// invalid request, an LLM or provider rate-limit response, or ErrLLMGenerate.
|
||||
// ErrCapacityExceeded identifies a Run or RunPrepared rejected because the
|
||||
// selected backend already admitted ConcurrencyLimit + QueueCapacity calls.
|
||||
// A [CapacityError] reports the selected backend ID. It is not an invalid
|
||||
// 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
|
||||
@@ -77,12 +78,15 @@ var (
|
||||
ErrValidation = errors.New("failed to validate output")
|
||||
)
|
||||
|
||||
// Engine prepares and runs Promptkit prompt requests.
|
||||
// Engine inspects prompts and profiles and prepares and runs Promptkit prompt
|
||||
// requests.
|
||||
//
|
||||
// An Engine is safe for concurrent calls to [Engine.Prepare] and [Engine.Run].
|
||||
// Each Engine owns independent backend-capacity pools that coordinate Run
|
||||
// admission and model generation. Injected collaborators may still be invoked
|
||||
// concurrently across different backend pools or for unlimited backends.
|
||||
// An Engine is safe for concurrent calls to [Engine.InspectPrompt],
|
||||
// [Engine.InspectProfile], [Engine.Prepare], [Engine.PrepareExecution],
|
||||
// [Engine.Run], and [Engine.RunPrepared]. Each Engine owns independent
|
||||
// backend-capacity pools that coordinate Run and RunPrepared admission and
|
||||
// model generation. Injected collaborators may still be invoked concurrently
|
||||
// across different backend pools or for unlimited backends.
|
||||
type Engine struct {
|
||||
runner *usecase.Runner
|
||||
}
|
||||
@@ -94,9 +98,10 @@ type Config struct {
|
||||
// It is required unless a WithPromptFS or WithPromptFile option supplies the
|
||||
// prompt source.
|
||||
PromptDir string
|
||||
// ProfileDir is an optional directory whose profiles take precedence over
|
||||
// embedded built-in profiles. An empty value selects only built-ins unless
|
||||
// profile options are also supplied.
|
||||
// ProfileDir is an optional ordinary configured source whose profiles take
|
||||
// precedence over application fallback and embedded built-in profiles. An
|
||||
// empty value selects the lower-precedence sources unless a profile-source
|
||||
// option supplies the ordinary source.
|
||||
ProfileDir string
|
||||
// SchemaDir is the root for JSON Schema files. An empty value uses the
|
||||
// current directory. WithSchemaFS or WithSchemaFile replaces this source.
|
||||
@@ -115,12 +120,12 @@ type Config struct {
|
||||
// Option customizes engine construction.
|
||||
//
|
||||
// NewEngine applies options in argument order and ignores nil options. Within
|
||||
// each prompt-source, profile-source, in-memory-profile, schema-source,
|
||||
// model-client, and artifact-reader category, the last non-nil valid option
|
||||
// replaces earlier options in that category. WithBackend is the additive
|
||||
// exception: unique registrations accumulate, and a repeated backend ID is an
|
||||
// error rather than a replacement. An invalid option fails construction even
|
||||
// if a later option would replace it.
|
||||
// each prompt-source, ordinary-profile-source, fallback-profile-source,
|
||||
// in-memory-profile, schema-source, model-client, and artifact-reader
|
||||
// category, the last non-nil valid option replaces earlier options in that
|
||||
// category. WithBackend is the additive exception: unique registrations
|
||||
// accumulate, and a repeated backend ID is an error rather than a replacement.
|
||||
// An invalid option fails construction even if a later option would replace it.
|
||||
type Option interface {
|
||||
apply(*engineOptions) error
|
||||
}
|
||||
@@ -132,26 +137,30 @@ func (f optionFunc) apply(options *engineOptions) error {
|
||||
}
|
||||
|
||||
type engineOptions struct {
|
||||
llmClient llm.Client
|
||||
artifactReader artifactadapter.Reader
|
||||
promptDefs promptdef.Repository
|
||||
profiles profile.Repository
|
||||
memoryProfiles profile.Repository
|
||||
backends []domain.Backend
|
||||
validator validate.Validator
|
||||
promptSource bool
|
||||
profileSource bool
|
||||
memorySource bool
|
||||
validatorSource bool
|
||||
artifactSource bool
|
||||
llmClient llm.Client
|
||||
artifactReader artifactadapter.Reader
|
||||
promptDefs promptdef.Repository
|
||||
profiles profile.Repository
|
||||
fallbackProfiles profile.Repository
|
||||
memoryProfiles profile.Repository
|
||||
backends []domain.Backend
|
||||
validator validate.Validator
|
||||
promptSource bool
|
||||
profileSource bool
|
||||
fallbackProfileSource bool
|
||||
memorySource bool
|
||||
validatorSource bool
|
||||
artifactSource bool
|
||||
}
|
||||
|
||||
// WithLLMClient replaces the built-in model client used by [Engine.Run].
|
||||
// WithLLMClient replaces the built-in model client used by [Engine.Run] and
|
||||
// [Engine.RunPrepared].
|
||||
//
|
||||
// A nil client makes NewEngine fail with ErrInvalidConfig. The Engine schedules
|
||||
// Generate calls according to the selected backend's capacity policy, but the
|
||||
// client may still be called concurrently across different backend pools or for
|
||||
// unlimited backends. The client is not used by [Engine.Prepare].
|
||||
// unlimited backends. The client is not used by [Engine.Prepare] or
|
||||
// [Engine.PrepareExecution].
|
||||
func WithLLMClient(client LLMClient) Option {
|
||||
return optionFunc(func(options *engineOptions) error {
|
||||
if client == nil {
|
||||
@@ -210,7 +219,7 @@ func WithPromptFile(path string) Option {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
options.promptDefs = promptdef.NewFSRepository(fsys, root)
|
||||
options.promptDefs = promptdef.NewFileRepository(fsys, root, filepath.Dir(path))
|
||||
options.promptSource = true
|
||||
return nil
|
||||
})
|
||||
@@ -218,12 +227,12 @@ func WithPromptFile(path string) Option {
|
||||
|
||||
// WithProfileFS loads execution profiles from fsys under root.
|
||||
//
|
||||
// Profiles from this source overlay built-in profiles. Profile YAML must use
|
||||
// api_key_env for environment-based credentials; raw API keys are rejected.
|
||||
// fsys must be non-nil and root must be non-empty; otherwise NewEngine fails
|
||||
// with ErrInvalidConfig. This option replaces Config.ProfileDir and earlier
|
||||
// file or FS profile-source options, but remains below WithProfiles in
|
||||
// precedence.
|
||||
// Profiles from this ordinary configured source take precedence over
|
||||
// application fallback and built-in profiles. Profile YAML must use api_key_env
|
||||
// for environment-based credentials; raw API keys are rejected. fsys must be
|
||||
// non-nil and root must be non-empty; otherwise NewEngine fails with
|
||||
// ErrInvalidConfig. This option replaces Config.ProfileDir and earlier file or
|
||||
// FS profile-source options, but remains below WithProfiles in precedence.
|
||||
func WithProfileFS(fsys fs.FS, root string) Option {
|
||||
return optionFunc(func(options *engineOptions) error {
|
||||
if fsys == nil {
|
||||
@@ -240,11 +249,12 @@ func WithProfileFS(fsys fs.FS, root string) Option {
|
||||
|
||||
// WithProfileFile loads execution profiles from the single profile file at path.
|
||||
//
|
||||
// The profile overlays built-in profiles. Profile YAML must use api_key_env for
|
||||
// environment-based credentials; raw API keys are rejected. path must name an
|
||||
// existing non-directory file when NewEngine applies the option. This option
|
||||
// replaces Config.ProfileDir and earlier file or FS profile-source options,
|
||||
// but remains below WithProfiles in precedence.
|
||||
// The profile takes precedence over application fallback and built-in profiles.
|
||||
// Profile YAML must use api_key_env for environment-based credentials; raw API
|
||||
// keys are rejected. path must name an existing non-directory file when
|
||||
// NewEngine applies the option. This option replaces Config.ProfileDir and
|
||||
// earlier file or FS profile-source options, but remains below WithProfiles in
|
||||
// precedence.
|
||||
func WithProfileFile(path string) Option {
|
||||
return optionFunc(func(options *engineOptions) error {
|
||||
fsys, root, err := fileSource(path)
|
||||
@@ -257,8 +267,41 @@ func WithProfileFile(path string) Option {
|
||||
})
|
||||
}
|
||||
|
||||
// WithFallbackProfileFS supplies application-owned fallback profile
|
||||
// definitions from fsys under root.
|
||||
//
|
||||
// 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
|
||||
// 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
|
||||
// format failure stops resolution.
|
||||
//
|
||||
// Files use the ordinary strict profile YAML and api_key_env credential rules.
|
||||
// Loading and validation are lazy: NewEngine validates this option's arguments
|
||||
// but does not read profile files. fsys must be non-nil and root must be
|
||||
// nonblank; otherwise NewEngine returns an error matching ErrInvalidConfig.
|
||||
// Repeating this option replaces the earlier valid fallback source.
|
||||
//
|
||||
// This option controls profile-definition lookup, not provider or generation
|
||||
// failover.
|
||||
func WithFallbackProfileFS(fsys fs.FS, root string) Option {
|
||||
return optionFunc(func(options *engineOptions) error {
|
||||
if fsys == nil {
|
||||
return ErrInvalidConfig
|
||||
}
|
||||
if strings.TrimSpace(root) == "" {
|
||||
return ErrInvalidConfig
|
||||
}
|
||||
options.fallbackProfiles = profile.NewFSRepository(fsys, root)
|
||||
options.fallbackProfileSource = true
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// WithProfiles configures in-memory profiles that take precedence over
|
||||
// configured profile files and built-in profiles.
|
||||
// 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
|
||||
@@ -344,13 +387,7 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
|
||||
promptDefs = promptdef.NewFilesystemRepository(cfg.PromptDir)
|
||||
}
|
||||
|
||||
profiles := builtin.NewRepositoryWithDirectory(cfg.ProfileDir)
|
||||
if options.profileSource {
|
||||
profiles = builtin.NewRepositoryWithPrimary(options.profiles)
|
||||
}
|
||||
if options.memorySource {
|
||||
profiles = profile.NewOverlayRepository(options.memoryProfiles, profiles)
|
||||
}
|
||||
profiles := newProfileRepository(cfg.ProfileDir, options)
|
||||
|
||||
backendRegistry, err := backend.NewRegistry(options.backends)
|
||||
if err != nil {
|
||||
@@ -403,26 +440,132 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newProfileRepository(profileDir string, options engineOptions) profile.Repository {
|
||||
repository := builtin.NewRepository()
|
||||
|
||||
if options.fallbackProfileSource {
|
||||
repository = profile.NewOverlayRepository(options.fallbackProfiles, repository)
|
||||
}
|
||||
|
||||
if options.profileSource {
|
||||
repository = profile.NewOverlayRepository(options.profiles, repository)
|
||||
} else if strings.TrimSpace(profileDir) != "" {
|
||||
repository = profile.NewOverlayRepository(profile.NewFilesystemRepository(profileDir), repository)
|
||||
}
|
||||
|
||||
if options.memorySource {
|
||||
repository = profile.NewOverlayRepository(options.memoryProfiles, repository)
|
||||
}
|
||||
|
||||
return repository
|
||||
}
|
||||
|
||||
func fileSource(name string) (fs.FS, string, error) {
|
||||
cleanName := strings.TrimSpace(name)
|
||||
if cleanName == "" {
|
||||
if strings.TrimSpace(name) == "" {
|
||||
return nil, "", ErrInvalidConfig
|
||||
}
|
||||
dir := filepath.Dir(cleanName)
|
||||
base := filepath.Base(cleanName)
|
||||
if base == "." || base == string(filepath.Separator) || strings.TrimSpace(base) == "" {
|
||||
dir := filepath.Dir(name)
|
||||
base := filepath.Base(name)
|
||||
if base == "." || base == string(filepath.Separator) {
|
||||
return nil, "", ErrInvalidConfig
|
||||
}
|
||||
info, err := os.Stat(cleanName)
|
||||
info, err := os.Stat(name)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("%w: failed to access source file %q: %v", ErrInvalidConfig, cleanName, err)
|
||||
return nil, "", fmt.Errorf("%w: failed to access source file %q: %v", ErrInvalidConfig, name, err)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return nil, "", fmt.Errorf("%w: source path %q must be a file", ErrInvalidConfig, cleanName)
|
||||
return nil, "", fmt.Errorf("%w: source path %q must be a file", ErrInvalidConfig, name)
|
||||
}
|
||||
return os.DirFS(dir), filepath.ToSlash(base), nil
|
||||
}
|
||||
|
||||
// InspectPrompt resolves one explicit prompt definition without selecting a
|
||||
// profile or starting execution work.
|
||||
//
|
||||
// InspectPrompt requires a nonblank promptID. It passes nonblank promptID and
|
||||
// promptVersion values unchanged to the engine's ordinary, case-sensitive
|
||||
// prompt selection. An empty version succeeds only when that source has one
|
||||
// selected ID; a nonempty version selects one exact ID/version pair. The
|
||||
// configured prompt source is used without merging, fallback, or enumeration.
|
||||
//
|
||||
// A successful result proves that the selected definition and any referenced
|
||||
// message content files were structurally loaded. Inputs are returned in
|
||||
// definition order. DefaultProfileID is declared metadata only and is not
|
||||
// resolved. OutputContract is the normalized declared contract, with a JSON
|
||||
// Schema path when declared but without loading or compiling that schema.
|
||||
// PromptHash is the same opaque equality value as PreparedRun.PromptHash for
|
||||
// the selected definition and observed source state; its spelling, length,
|
||||
// encoding, algorithm, and security properties are not contracts.
|
||||
//
|
||||
// This method does not return prompt bodies, templates, source paths, schemas,
|
||||
// rendered messages, or execution settings. It does not resolve a profile or
|
||||
// credential, read artifacts or schemas, render, validate, admit capacity,
|
||||
// contact a provider, or generate model output. The returned PromptInspection
|
||||
// and its input slice are caller-owned. Filesystem-backed inspection is a
|
||||
// point-in-time lookup and does not freeze a definition for later execution.
|
||||
//
|
||||
// A nil Engine returns an error matching ErrInvalidConfig. A blank prompt ID
|
||||
// matches ErrInvalidRequest. An absent exact ID or version matches
|
||||
// ErrPromptNotFound and not ErrPromptLoad. Malformed, unreadable, duplicate,
|
||||
// ambiguous, referenced-content, or hashing failures match ErrPromptLoad.
|
||||
// Cancellation during lookup matches ErrPromptLoad while preserving the
|
||||
// context error. InspectPrompt returns no partial result on error.
|
||||
func (e *Engine) InspectPrompt(
|
||||
ctx context.Context,
|
||||
promptID string,
|
||||
promptVersion string,
|
||||
) (*PromptInspection, error) {
|
||||
if e == nil || e.runner == nil {
|
||||
return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig)
|
||||
}
|
||||
|
||||
inspection, err := e.runner.InspectPrompt(ctx, promptID, promptVersion)
|
||||
if err != nil {
|
||||
return nil, mapPublicError(err)
|
||||
}
|
||||
return fromDomainPromptInspection(inspection), nil
|
||||
}
|
||||
|
||||
// InspectProfile resolves one explicit profile without selecting a prompt or
|
||||
// starting execution work.
|
||||
//
|
||||
// InspectProfile trims surrounding whitespace from profileID and looks up the
|
||||
// resulting nonblank ID exactly and case-sensitively through the engine's
|
||||
// in-memory, ordinary configured-source, application fallback, and built-in
|
||||
// profile precedence. It applies the framework timeout baseline, selected
|
||||
// backend, and then selected profile to EffectiveModelParams without a request
|
||||
// override. BackendID is empty for an endpoint-only profile.
|
||||
//
|
||||
// APIKeyEnv in the returned target is an environment-variable name, never its
|
||||
// value. APIKeyRequired instead reports a direct credential requirement and is
|
||||
// mutually exclusive with a nonblank APIKeyEnv. InspectProfile neither derives
|
||||
// an ID from a prompt default_profile nor checks credential availability, so an
|
||||
// absent or blank named environment variable is not an error.
|
||||
//
|
||||
// The returned ProfileInspection and all nested mutable values are
|
||||
// caller-owned. Filesystem-backed inspection is a point-in-time lookup and
|
||||
// does not freeze the profile for a later execution. This method does not load
|
||||
// a prompt, render, read artifacts or schemas, admit backend capacity, contact
|
||||
// a provider, or generate model output.
|
||||
//
|
||||
// A nil Engine returns an error matching ErrInvalidConfig. A blank profile ID
|
||||
// matches ErrInvalidRequest. An absent exact ID matches ErrProfileNotFound and
|
||||
// not ErrProfileLoad. Malformed or unreadable profile data, an unknown backend,
|
||||
// or an invalid resolved target matches ErrProfileLoad. Cancellation during
|
||||
// profile loading matches ErrProfileLoad while preserving the context error.
|
||||
// InspectProfile returns no partial result on error.
|
||||
func (e *Engine) InspectProfile(ctx context.Context, profileID string) (*ProfileInspection, error) {
|
||||
if e == nil || e.runner == nil {
|
||||
return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig)
|
||||
}
|
||||
|
||||
inspection, err := e.runner.InspectProfile(ctx, profileID)
|
||||
if err != nil {
|
||||
return nil, mapPublicError(err)
|
||||
}
|
||||
return fromDomainProfileInspection(inspection), nil
|
||||
}
|
||||
|
||||
// Prepare resolves and renders a prompt request without calling an LLM.
|
||||
//
|
||||
// Prepare selects the prompt and profile, resolves any selected backend and
|
||||
@@ -457,6 +600,38 @@ func (e *Engine) Prepare(ctx context.Context, req RunRequest) (*PreparedRun, err
|
||||
return fromDomainPreparedRun(prepared), nil
|
||||
}
|
||||
|
||||
// PrepareExecution completely prepares a prompt request without calling the
|
||||
// configured LLMClient or reserving backend admission capacity.
|
||||
//
|
||||
// The returned opaque handle is bound to this Engine and permits one
|
||||
// [Engine.RunPrepared] invocation. Preparation freezes the selected sources,
|
||||
// rendered messages, effective settings, inputs, provider structured-output
|
||||
// metadata, and validation resources needed by that invocation. The handle
|
||||
// retains a direct RunRequest.APIKey only in private execution state;
|
||||
// [PreparedExecution.Details] is credential-redacted.
|
||||
//
|
||||
// The context governs preparation only. Cancellation after this method
|
||||
// returns does not invalidate the handle or propagate to RunPrepared.
|
||||
// PrepareExecution returns the same error categories as [Engine.Prepare] and
|
||||
// returns no handle on error. A nil Engine returns an error matching
|
||||
// ErrInvalidConfig.
|
||||
func (e *Engine) PrepareExecution(ctx context.Context, req RunRequest) (*PreparedExecution, error) {
|
||||
if e == nil || e.runner == nil {
|
||||
return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig)
|
||||
}
|
||||
|
||||
domainReq, err := toDomainRunRequest(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
prepared, err := e.runner.PrepareExecution(ctx, domainReq)
|
||||
if err != nil {
|
||||
return nil, mapPublicError(err)
|
||||
}
|
||||
return &PreparedExecution{internal: prepared}, nil
|
||||
}
|
||||
|
||||
// Run prepares a request, invokes the configured LLMClient, and validates the
|
||||
// generated output.
|
||||
//
|
||||
@@ -467,9 +642,10 @@ func (e *Engine) Prepare(ctx context.Context, req RunRequest) (*PreparedRun, err
|
||||
// single-pass even when OutputContract.RepairAttempts is positive.
|
||||
//
|
||||
// Run can return every error category documented by [Engine.Prepare], plus
|
||||
// ErrCapacityExceeded and ErrLLMGenerate. ErrCapacityExceeded identifies
|
||||
// rejection before artifacts, schemas, rendering, or model generation because
|
||||
// the selected backend's admission capacity is full; it does not match
|
||||
// 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
|
||||
@@ -491,3 +667,42 @@ func (e *Engine) Run(ctx context.Context, req RunRequest) (*RunResult, error) {
|
||||
}
|
||||
return fromDomainRunResult(result), nil
|
||||
}
|
||||
|
||||
// RunPrepared atomically claims and executes a handle created by
|
||||
// [Engine.PrepareExecution].
|
||||
//
|
||||
// A valid owning-Engine invocation consumes the handle's one attempt before
|
||||
// credential revalidation, backend admission, generation, or validation.
|
||||
// Cancellation, capacity rejection, generation failure, operational
|
||||
// validation failure, and success all leave the handle unusable. A nil,
|
||||
// zero-value, foreign-Engine, discarded, claimed, or used handle returns an
|
||||
// error matching ErrInvalidRequest; a nil Engine returns ErrInvalidConfig and
|
||||
// does not claim the handle.
|
||||
//
|
||||
// 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.
|
||||
//
|
||||
// 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.
|
||||
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)
|
||||
}
|
||||
|
||||
var internal *usecase.PreparedExecution
|
||||
if prepared != nil {
|
||||
internal = prepared.internal
|
||||
}
|
||||
result, err := e.runner.RunPrepared(ctx, internal)
|
||||
if err != nil {
|
||||
return nil, mapPublicError(err)
|
||||
}
|
||||
return fromDomainRunResult(result), nil
|
||||
}
|
||||
|
||||
1095
engine_test.go
1095
engine_test.go
File diff suppressed because it is too large
Load Diff
@@ -3,6 +3,7 @@ package promptkit
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/capacity"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
@@ -14,6 +15,11 @@ func mapPublicError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
var internalCapacityError *usecase.CapacityError
|
||||
if errors.As(err, &internalCapacityError) && internalCapacityError != nil &&
|
||||
strings.TrimSpace(internalCapacityError.BackendID) != "" {
|
||||
return &CapacityError{BackendID: internalCapacityError.BackendID}
|
||||
}
|
||||
publicErr := publicErrorFor(err)
|
||||
if publicErr == nil {
|
||||
return err
|
||||
|
||||
@@ -20,3 +20,31 @@ func TestMapPublicErrorPreservesGenerationCancellation(t *testing.T) {
|
||||
t.Fatalf("mapped error=%v, want context.Canceled", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapPublicErrorTranslatesCapacityError(t *testing.T) {
|
||||
internalErr := &usecase.CapacityError{BackendID: "limited"}
|
||||
|
||||
err := mapPublicError(internalErr)
|
||||
var publicErr *CapacityError
|
||||
if !errors.As(err, &publicErr) || publicErr == nil {
|
||||
t.Fatalf("mapped error=%v, want public CapacityError", err)
|
||||
}
|
||||
if publicErr.BackendID != "limited" {
|
||||
t.Fatalf("mapped backend ID=%q, want limited", publicErr.BackendID)
|
||||
}
|
||||
if !errors.Is(err, ErrCapacityExceeded) {
|
||||
t.Fatalf("mapped error=%v, want ErrCapacityExceeded", err)
|
||||
}
|
||||
if errors.Is(err, ErrInvalidRequest) || errors.Is(err, ErrLLMGenerate) {
|
||||
t.Fatalf("mapped capacity error has an unrelated category: %v", err)
|
||||
}
|
||||
var leakedInternalErr *usecase.CapacityError
|
||||
if errors.As(err, &leakedInternalErr) {
|
||||
t.Fatalf("mapped error exposes internal CapacityError: %v", err)
|
||||
}
|
||||
|
||||
internalErr.BackendID = "changed"
|
||||
if publicErr.BackendID != "limited" {
|
||||
t.Fatalf("mapped backend ID changed with source error: %q", publicErr.BackendID)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,10 +16,12 @@ import (
|
||||
|
||||
var (
|
||||
ErrUnsupportedRefType = errors.New("unsupported artifact reference type")
|
||||
ErrMissingInlineBody = errors.New("missing body for inline artifact")
|
||||
ErrMissingFilePath = errors.New("missing file path for file artifact")
|
||||
ErrUnsupportedFile = errors.New("file artifact path is not a regular file")
|
||||
)
|
||||
|
||||
const fileReadChunkSize = 64 * 1024
|
||||
|
||||
// Reader resolves artifact references into actual artifacts.
|
||||
type Reader interface {
|
||||
Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error)
|
||||
@@ -34,7 +36,7 @@ type CompositeReader struct {
|
||||
func NewCompositeReader() Reader {
|
||||
return &CompositeReader{
|
||||
inlineReader: &inlineReader{},
|
||||
fileReader: &fileReader{},
|
||||
fileReader: &fileReader{open: openArtifactFile},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -64,10 +66,6 @@ func (r *inlineReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domai
|
||||
default:
|
||||
}
|
||||
|
||||
if ref.Body == "" {
|
||||
return nil, ErrMissingInlineBody
|
||||
}
|
||||
|
||||
body := []byte(ref.Body)
|
||||
return &domain.Artifact{
|
||||
ContentType: defaults.ContentTypeTextPlain,
|
||||
@@ -78,7 +76,15 @@ func (r *inlineReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domai
|
||||
}, nil
|
||||
}
|
||||
|
||||
type fileReader struct{}
|
||||
type artifactFile interface {
|
||||
Read([]byte) (int, error)
|
||||
Stat() (os.FileInfo, error)
|
||||
Close() error
|
||||
}
|
||||
|
||||
type fileReader struct {
|
||||
open func(string) (artifactFile, error)
|
||||
}
|
||||
|
||||
func (r *fileReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error) {
|
||||
select {
|
||||
@@ -91,25 +97,71 @@ func (r *fileReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.
|
||||
return nil, ErrMissingFilePath
|
||||
}
|
||||
|
||||
return readFileArtifact(ref.URI)
|
||||
return readFileArtifact(ctx, ref.URI, r.open)
|
||||
}
|
||||
|
||||
func readFileArtifact(path string) (*domain.Artifact, error) {
|
||||
file, err := os.Open(path)
|
||||
func openArtifactFile(path string) (artifactFile, error) {
|
||||
return os.Open(path)
|
||||
}
|
||||
|
||||
func readFileArtifact(ctx context.Context, path string, open func(string) (artifactFile, error)) (*domain.Artifact, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read file %s: %w", path, err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return nil, fmt.Errorf("%w: %s", ErrUnsupportedFile, path)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
file, err := open(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read file %s: %w", path, err)
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
data, err := io.ReadAll(file)
|
||||
openedInfo, err := file.Stat()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read file %s: %w", path, err)
|
||||
return nil, fmt.Errorf("failed to inspect opened file %s: %w", path, err)
|
||||
}
|
||||
if !openedInfo.Mode().IsRegular() {
|
||||
return nil, fmt.Errorf("%w: %s", ErrUnsupportedFile, path)
|
||||
}
|
||||
|
||||
data := make([]byte, 0)
|
||||
chunk := make([]byte, fileReadChunkSize)
|
||||
for {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
n, readErr := file.Read(chunk)
|
||||
if n > 0 {
|
||||
data = append(data, chunk[:n]...)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if errors.Is(readErr, io.EOF) {
|
||||
break
|
||||
}
|
||||
if readErr != nil {
|
||||
return nil, fmt.Errorf("failed to read file %s: %w", path, readErr)
|
||||
}
|
||||
}
|
||||
|
||||
contentType := mime.TypeByExtension(filepath.Ext(path))
|
||||
if contentType == "" {
|
||||
contentType = defaults.ContentTypeTextPlain
|
||||
}
|
||||
hash := fmt.Sprintf("%x", sha256.Sum256(data))
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &domain.Artifact{
|
||||
Name: filepath.Base(path),
|
||||
@@ -117,6 +169,6 @@ func readFileArtifact(path string) (*domain.Artifact, error) {
|
||||
Body: data,
|
||||
URI: path,
|
||||
Size: int64(len(data)),
|
||||
Hash: fmt.Sprintf("%x", sha256.Sum256(data)),
|
||||
Hash: hash,
|
||||
}, nil
|
||||
}
|
||||
|
||||
43
internal/artifact/reader_fifo_linux_test.go
Normal file
43
internal/artifact/reader_fifo_linux_test.go
Normal file
@@ -0,0 +1,43 @@
|
||||
//go:build linux
|
||||
|
||||
package artifact
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
func TestFileReaderRejectsFIFOBeforeOpen(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "artifact.fifo")
|
||||
if err := syscall.Mkfifo(path, 0o600); err != nil {
|
||||
t.Fatalf("create fifo: %v", err)
|
||||
}
|
||||
|
||||
type result struct {
|
||||
artifact *domain.Artifact
|
||||
err error
|
||||
}
|
||||
done := make(chan result, 1)
|
||||
go func() {
|
||||
artifact, err := NewCompositeReader().Read(context.Background(), domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefFile,
|
||||
URI: path,
|
||||
})
|
||||
done <- result{artifact: artifact, err: err}
|
||||
}()
|
||||
|
||||
select {
|
||||
case got := <-done:
|
||||
if got.artifact != nil || !errors.Is(got.err, ErrUnsupportedFile) {
|
||||
t.Fatalf("artifact=%#v err=%v, want nil/ErrUnsupportedFile", got.artifact, got.err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("FIFO read blocked instead of rejecting the non-regular file")
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package artifact
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
@@ -11,54 +12,97 @@ import (
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
func TestCompositeReader_Read(t *testing.T) {
|
||||
func TestCompositeReaderRejectsUnsupportedReferences(t *testing.T) {
|
||||
_, err := NewCompositeReader().Read(context.Background(), domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefType("unsupported"),
|
||||
URI: "unsupported://bucket/key",
|
||||
})
|
||||
if !errors.Is(err, ErrUnsupportedRefType) {
|
||||
t.Fatalf("expected ErrUnsupportedRefType, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompositeReaderSourceParityAndOpaqueHashes(t *testing.T) {
|
||||
reader := NewCompositeReader()
|
||||
ctx := context.Background()
|
||||
hashes := make(map[string]string)
|
||||
tests := []struct {
|
||||
name string
|
||||
content string
|
||||
}{
|
||||
{name: "empty", content: ""},
|
||||
{name: "ordinary", content: "same content"},
|
||||
{name: "changed", content: "changed content"},
|
||||
}
|
||||
|
||||
t.Run("inline artifact", func(t *testing.T) {
|
||||
ref := domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefInline,
|
||||
Body: "hello world",
|
||||
}
|
||||
art, err := reader.Read(ctx, ref)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if string(art.Body) != "hello world" {
|
||||
t.Errorf("expected 'hello world', got %s", string(art.Body))
|
||||
}
|
||||
if art.ContentType != "text/plain" {
|
||||
t.Errorf("expected text/plain content type, got %q", art.ContentType)
|
||||
}
|
||||
if art.Hash != "b94d27b9934d3e08a52e52d7da7dabfac484efe37a5380ee9088f7ace2efcde9" {
|
||||
t.Errorf("unexpected hash: %s", art.Hash)
|
||||
}
|
||||
if art.Size != int64(len(ref.Body)) {
|
||||
t.Errorf("expected size %d, got %d", len(ref.Body), art.Size)
|
||||
}
|
||||
})
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
filePath := filepath.Join(t.TempDir(), "artifact.txt")
|
||||
if err := os.WriteFile(filePath, []byte(tc.content), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
t.Run("inline artifact missing body", func(t *testing.T) {
|
||||
ref := domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefInline,
|
||||
Body: "",
|
||||
}
|
||||
_, err := reader.Read(ctx, ref)
|
||||
if !errors.Is(err, ErrMissingInlineBody) {
|
||||
t.Errorf("expected ErrMissingInlineBody, got %v", err)
|
||||
}
|
||||
})
|
||||
sources := []struct {
|
||||
name string
|
||||
ref domain.ArtifactRef
|
||||
wantURI string
|
||||
}{
|
||||
{
|
||||
name: "inline",
|
||||
ref: domain.ArtifactRef{Type: domain.ArtifactRefInline, Body: tc.content},
|
||||
},
|
||||
{
|
||||
name: "inline with uri",
|
||||
ref: domain.ArtifactRef{Type: domain.ArtifactRefInline, URI: "memory://input", Body: tc.content},
|
||||
wantURI: "memory://input",
|
||||
},
|
||||
{
|
||||
name: "file",
|
||||
ref: domain.ArtifactRef{Type: domain.ArtifactRefFile, URI: filePath},
|
||||
wantURI: filePath,
|
||||
},
|
||||
}
|
||||
|
||||
t.Run("unsupported ref type", func(t *testing.T) {
|
||||
ref := domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefType("unsupported"),
|
||||
URI: "unsupported://bucket/key",
|
||||
}
|
||||
_, err := reader.Read(ctx, ref)
|
||||
if !errors.Is(err, ErrUnsupportedRefType) {
|
||||
t.Error("expected error for unsupported type")
|
||||
}
|
||||
})
|
||||
var sourceHash string
|
||||
for _, source := range sources {
|
||||
t.Run(source.name, func(t *testing.T) {
|
||||
first, err := reader.Read(context.Background(), source.ref)
|
||||
if err != nil {
|
||||
t.Fatalf("first read: %v", err)
|
||||
}
|
||||
second, err := reader.Read(context.Background(), source.ref)
|
||||
if err != nil {
|
||||
t.Fatalf("second read: %v", err)
|
||||
}
|
||||
if string(first.Body) != tc.content || first.Size != int64(len(tc.content)) {
|
||||
t.Fatalf("body=%q size=%d, want %q/%d", first.Body, first.Size, tc.content, len(tc.content))
|
||||
}
|
||||
if first.URI != source.wantURI {
|
||||
t.Fatalf("URI = %q, want %q", first.URI, source.wantURI)
|
||||
}
|
||||
if first.Hash == "" || first.Hash != second.Hash {
|
||||
t.Fatalf("hashes are not non-empty and stable: %q/%q", first.Hash, second.Hash)
|
||||
}
|
||||
if sourceHash == "" {
|
||||
sourceHash = first.Hash
|
||||
} else if first.Hash != sourceHash {
|
||||
t.Fatalf("equal content hashes differ: %q/%q", sourceHash, first.Hash)
|
||||
}
|
||||
if source.ref.Type == domain.ArtifactRefFile {
|
||||
if first.Name != filepath.Base(filePath) || !strings.HasPrefix(first.ContentType, "text/plain") {
|
||||
t.Fatalf("unexpected file metadata: %+v", first)
|
||||
}
|
||||
} else if first.ContentType != "text/plain" {
|
||||
t.Fatalf("inline content type = %q", first.ContentType)
|
||||
}
|
||||
})
|
||||
}
|
||||
hashes[tc.name] = sourceHash
|
||||
})
|
||||
}
|
||||
|
||||
if hashes["empty"] == hashes["ordinary"] || hashes["ordinary"] == hashes["changed"] {
|
||||
t.Fatalf("changed content did not change opaque hash: %#v", hashes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompositeReaderCopiesInlineData(t *testing.T) {
|
||||
@@ -87,94 +131,154 @@ func TestCompositeReaderCopiesInlineData(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompositeReaderHonorsCancellation(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
func TestCompositeReaderHonorsPreCancellation(t *testing.T) {
|
||||
filePath := filepath.Join(t.TempDir(), "artifact.txt")
|
||||
if err := os.WriteFile(filePath, []byte("ignored"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
ref domain.ArtifactRef
|
||||
}{
|
||||
{name: "inline", ref: domain.ArtifactRef{Type: domain.ArtifactRefInline, Body: "ignored"}},
|
||||
{name: "file", ref: domain.ArtifactRef{Type: domain.ArtifactRefFile, URI: filePath}},
|
||||
}
|
||||
|
||||
_, err := NewCompositeReader().Read(ctx, domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefInline,
|
||||
Body: "ignored",
|
||||
})
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
artifact, err := NewCompositeReader().Read(ctx, tc.ref)
|
||||
if artifact != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("artifact=%#v err=%v, want nil/context.Canceled", artifact, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileReader_Read(t *testing.T) {
|
||||
content := []byte("test file content")
|
||||
filePath := filepath.Join(t.TempDir(), "artifact.txt")
|
||||
if err := os.WriteFile(filePath, content, 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
func TestFileReaderFailuresAndMetadata(t *testing.T) {
|
||||
reader := NewCompositeReader()
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("file artifact loading", func(t *testing.T) {
|
||||
ref := domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefFile,
|
||||
URI: filePath,
|
||||
}
|
||||
art, err := reader.Read(ctx, ref)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if string(art.Body) != string(content) {
|
||||
t.Errorf("expected %s, got %s", string(content), string(art.Body))
|
||||
}
|
||||
if art.Name != filepath.Base(filePath) {
|
||||
t.Errorf("expected name %q, got %q", filepath.Base(filePath), art.Name)
|
||||
}
|
||||
if !strings.HasPrefix(art.ContentType, "text/plain") {
|
||||
t.Errorf("expected text content type, got %q", art.ContentType)
|
||||
}
|
||||
if art.URI != filePath {
|
||||
t.Errorf("expected URI %q, got %q", filePath, art.URI)
|
||||
}
|
||||
if art.Size != int64(len(content)) {
|
||||
t.Errorf("expected size %d, got %d", len(content), art.Size)
|
||||
}
|
||||
if art.Hash != "60f5237ed4049f0382661ef009d2bc42e48c3ceb3edb6600f7024e7ab3b838f3" {
|
||||
t.Errorf("unexpected hash: %s", art.Hash)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing file path", func(t *testing.T) {
|
||||
ref := domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefFile,
|
||||
URI: "",
|
||||
}
|
||||
_, err := reader.Read(ctx, ref)
|
||||
_, err := reader.Read(context.Background(), domain.ArtifactRef{Type: domain.ArtifactRefFile})
|
||||
if !errors.Is(err, ErrMissingFilePath) {
|
||||
t.Errorf("expected ErrMissingFilePath, got %v", err)
|
||||
t.Fatalf("expected ErrMissingFilePath, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing file", func(t *testing.T) {
|
||||
ref := domain.ArtifactRef{
|
||||
_, err := reader.Read(context.Background(), domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefFile,
|
||||
URI: filepath.Join(t.TempDir(), "missing.txt"),
|
||||
}
|
||||
if _, err := reader.Read(ctx, ref); err == nil {
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected missing file error")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("directory rejected before open", func(t *testing.T) {
|
||||
artifact, err := reader.Read(context.Background(), domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefFile,
|
||||
URI: t.TempDir(),
|
||||
})
|
||||
if artifact != nil || !errors.Is(err, ErrUnsupportedFile) {
|
||||
t.Fatalf("artifact=%#v err=%v, want nil/ErrUnsupportedFile", artifact, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-regular opened target rejected", func(t *testing.T) {
|
||||
filePath := filepath.Join(t.TempDir(), "artifact.txt")
|
||||
if err := os.WriteFile(filePath, []byte("content"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
directoryInfo, err := os.Stat(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
fileReader := &fileReader{open: func(path string) (artifactFile, error) {
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &reportedInfoFile{artifactFile: file, info: directoryInfo}, nil
|
||||
}}
|
||||
|
||||
artifact, err := fileReader.Read(context.Background(), domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefFile,
|
||||
URI: filePath,
|
||||
})
|
||||
if artifact != nil || !errors.Is(err, ErrUnsupportedFile) {
|
||||
t.Fatalf("artifact=%#v err=%v, want nil/ErrUnsupportedFile", artifact, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unknown extension uses text fallback", func(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "artifact.unknownextension")
|
||||
if err := os.WriteFile(path, content, 0o600); err != nil {
|
||||
filePath := filepath.Join(t.TempDir(), "artifact.unknownextension")
|
||||
if err := os.WriteFile(filePath, []byte("content"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
art, err := reader.Read(ctx, domain.ArtifactRef{
|
||||
artifact, err := reader.Read(context.Background(), domain.ArtifactRef{
|
||||
Type: domain.ArtifactRefFile,
|
||||
URI: path,
|
||||
URI: filePath,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
t.Fatalf("read artifact: %v", err)
|
||||
}
|
||||
if art.ContentType != "text/plain" {
|
||||
t.Errorf("expected text/plain fallback, got %q", art.ContentType)
|
||||
if artifact.ContentType != "text/plain" {
|
||||
t.Fatalf("content type = %q", artifact.ContentType)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestFileReaderCancelsAfterReadProgress(t *testing.T) {
|
||||
filePath := filepath.Join(t.TempDir(), "artifact.bin")
|
||||
content := bytes.Repeat([]byte("x"), fileReadChunkSize*2)
|
||||
if err := os.WriteFile(filePath, content, 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
var opened *cancelAfterProgressFile
|
||||
reader := &fileReader{open: func(path string) (artifactFile, error) {
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
opened = &cancelAfterProgressFile{artifactFile: file, cancel: cancel}
|
||||
return opened, nil
|
||||
}}
|
||||
|
||||
artifact, err := reader.Read(ctx, domain.ArtifactRef{Type: domain.ArtifactRefFile, URI: filePath})
|
||||
if artifact != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("artifact=%#v err=%v, want nil/context.Canceled", artifact, err)
|
||||
}
|
||||
if opened == nil || opened.reads != 1 {
|
||||
t.Fatalf("read count = %v, want one progressing read", opened)
|
||||
}
|
||||
}
|
||||
|
||||
type reportedInfoFile struct {
|
||||
artifactFile
|
||||
info os.FileInfo
|
||||
}
|
||||
|
||||
func (f *reportedInfoFile) Stat() (os.FileInfo, error) {
|
||||
return f.info, nil
|
||||
}
|
||||
|
||||
type cancelAfterProgressFile struct {
|
||||
artifactFile
|
||||
cancel context.CancelFunc
|
||||
reads int
|
||||
}
|
||||
|
||||
func (f *cancelAfterProgressFile) Read(buffer []byte) (int, error) {
|
||||
n, err := f.artifactFile.Read(buffer)
|
||||
if n > 0 {
|
||||
f.reads++
|
||||
f.cancel()
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ package backend
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
@@ -109,10 +108,11 @@ func (r *Registry) CapacityPolicies() map[string]domain.BackendCapacityPolicy {
|
||||
}
|
||||
|
||||
func normalizeBackend(definition domain.Backend) (domain.Backend, error) {
|
||||
definition.Endpoint = strings.TrimSpace(definition.Endpoint)
|
||||
if err := validateEndpoint(definition.Endpoint); err != nil {
|
||||
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(definition.Endpoint)
|
||||
if err != nil {
|
||||
return domain.Backend{}, fmt.Errorf("backend %q endpoint: %w", definition.ID, err)
|
||||
}
|
||||
definition.Endpoint = endpoint
|
||||
|
||||
definition.APIKeyEnv = strings.TrimSpace(definition.APIKeyEnv)
|
||||
if definition.APIKeyEnv != "" && !environmentVariableName.MatchString(definition.APIKeyEnv) {
|
||||
@@ -182,31 +182,3 @@ func normalizeBackend(definition domain.Backend) (domain.Backend, error) {
|
||||
definition.ExtraParams = extraParams
|
||||
return definition, nil
|
||||
}
|
||||
|
||||
func validateEndpoint(endpoint string) error {
|
||||
if endpoint == "" {
|
||||
return errors.New("must not be blank")
|
||||
}
|
||||
if strings.Contains(endpoint, "#") {
|
||||
return errors.New("must not contain a fragment")
|
||||
}
|
||||
|
||||
parsed, err := url.Parse(endpoint)
|
||||
if err != nil {
|
||||
return fmt.Errorf("must be a valid URL: %w", err)
|
||||
}
|
||||
scheme := strings.ToLower(parsed.Scheme)
|
||||
if scheme != "http" && scheme != "https" {
|
||||
return errors.New("must use http or https")
|
||||
}
|
||||
if !parsed.IsAbs() || parsed.Hostname() == "" {
|
||||
return errors.New("must be absolute and include a host")
|
||||
}
|
||||
if parsed.User != nil {
|
||||
return errors.New("must not contain user information")
|
||||
}
|
||||
if parsed.RawQuery != "" || parsed.ForceQuery {
|
||||
return errors.New("must not contain a query string")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -124,6 +124,11 @@ func TestManagerAdmissionHonorsContextAndUnlimitedBackends(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("construct manager: %v", err)
|
||||
}
|
||||
release, err := manager.Admit(context.Background(), "limited")
|
||||
if err != nil {
|
||||
t.Fatalf("fill limited pool: %v", err)
|
||||
}
|
||||
defer release()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
@@ -14,21 +14,12 @@ const (
|
||||
ContentTypeApplicationJSON = "application/json"
|
||||
OpenAIChatCompletionsPath = "/chat/completions"
|
||||
|
||||
ExecutionDefaultTemperature = 0.0
|
||||
ExecutionDefaultMaxTokens = 0
|
||||
ExecutionDefaultTopP = 1.0
|
||||
ExecutionDefaultTimeoutSeconds = 600
|
||||
)
|
||||
|
||||
var (
|
||||
LLMRequestTimeoutDefault = 10 * time.Minute
|
||||
LLMRequestTimeoutDefault = 10 * time.Minute
|
||||
)
|
||||
|
||||
func ExecutionTargetDefault() domain.ExecutionTarget {
|
||||
return domain.ExecutionTarget{
|
||||
Temperature: ExecutionDefaultTemperature,
|
||||
MaxTokens: ExecutionDefaultMaxTokens,
|
||||
TopP: ExecutionDefaultTopP,
|
||||
TimeoutSeconds: ExecutionDefaultTimeoutSeconds,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -97,22 +97,22 @@ type RunResult struct {
|
||||
// PreparedRun contains pre-LLM execution state from the prepare/render phase.
|
||||
// It must never include resolved API key values, model output, or validation data.
|
||||
type PreparedRun struct {
|
||||
PromptID string `json:"prompt_id"`
|
||||
PromptVersion string `json:"prompt_version,omitempty"`
|
||||
PromptHash string `json:"prompt_hash,omitempty"`
|
||||
SelectedProfileID string `json:"selected_profile_id"`
|
||||
SelectedBackendID string `json:"selected_backend_id,omitempty"`
|
||||
EffectiveModelParams ExecutionTarget `json:"effective_model_params"`
|
||||
TargetPresence ExecutionTargetPresence `json:"-"`
|
||||
OutputContract OutputContract `json:"output_contract"`
|
||||
StructuredOutput *StructuredOutputSpec `json:"structured_output,omitempty"`
|
||||
InputHashes map[string]string `json:"input_hashes,omitempty"`
|
||||
SessionID string `json:"session_id,omitempty"`
|
||||
RenderedPromptHash string `json:"rendered_prompt_hash"`
|
||||
Messages []RenderedMessage `json:"messages"`
|
||||
StartTime time.Time `json:"start_time,omitempty"`
|
||||
EndTime time.Time `json:"end_time,omitempty"`
|
||||
DurationMS int64 `json:"duration_ms,omitempty"`
|
||||
PromptID string
|
||||
PromptVersion string
|
||||
PromptHash string
|
||||
SelectedProfileID string
|
||||
SelectedBackendID string
|
||||
EffectiveModelParams ExecutionTarget
|
||||
TargetPresence ExecutionTargetPresence
|
||||
OutputContract OutputContract
|
||||
StructuredOutput *StructuredOutputSpec
|
||||
InputHashes map[string]string
|
||||
SessionID string
|
||||
RenderedPromptHash string
|
||||
Messages []RenderedMessage
|
||||
StartTime time.Time
|
||||
EndTime time.Time
|
||||
DurationMS int64
|
||||
}
|
||||
|
||||
// ArtifactRef represents a reference to an input artifact.
|
||||
@@ -145,6 +145,16 @@ type PromptDefinition struct {
|
||||
Validation OutputContract `yaml:"validation"`
|
||||
}
|
||||
|
||||
// PromptInspection is the resolved result of exact prompt inspection.
|
||||
type PromptInspection struct {
|
||||
PromptID string
|
||||
PromptVersion string
|
||||
PromptHash string
|
||||
DefaultProfileID string
|
||||
Inputs []PromptInput
|
||||
OutputContract OutputContract
|
||||
}
|
||||
|
||||
// PromptInput describes one named input expected by a prompt definition.
|
||||
type PromptInput struct {
|
||||
Name string `yaml:"name"`
|
||||
@@ -236,6 +246,13 @@ type ExecutionTarget struct {
|
||||
ExtraParams map[string]any `yaml:"extra_params" json:"extra_params"`
|
||||
}
|
||||
|
||||
// ProfileInspection is the resolved result of exact profile inspection.
|
||||
type ProfileInspection struct {
|
||||
ProfileID string
|
||||
EffectiveModelParams ExecutionTarget
|
||||
APIKeyRequired bool
|
||||
}
|
||||
|
||||
// OutputContract defines the requirements for the output artifact.
|
||||
type OutputContract struct {
|
||||
Format OutputFormat `yaml:"format"`
|
||||
|
||||
39
internal/domain/endpoint.go
Normal file
39
internal/domain/endpoint.go
Normal file
@@ -0,0 +1,39 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// NormalizeOpenAICompatibleBaseEndpoint trims and validates a source-neutral
|
||||
// OpenAI-compatible provider base endpoint.
|
||||
func NormalizeOpenAICompatibleBaseEndpoint(endpoint string) (string, error) {
|
||||
endpoint = strings.TrimSpace(endpoint)
|
||||
if endpoint == "" {
|
||||
return "", errors.New("endpoint must not be blank")
|
||||
}
|
||||
if strings.Contains(endpoint, "#") {
|
||||
return "", errors.New("endpoint must not contain a fragment")
|
||||
}
|
||||
|
||||
parsed, err := url.Parse(endpoint)
|
||||
if err != nil {
|
||||
return "", errors.New("endpoint must be a valid URL")
|
||||
}
|
||||
parsed.Scheme = strings.ToLower(parsed.Scheme)
|
||||
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
||||
return "", errors.New("endpoint must use http or https")
|
||||
}
|
||||
if !parsed.IsAbs() || parsed.Hostname() == "" {
|
||||
return "", errors.New("endpoint must be absolute and include a host")
|
||||
}
|
||||
if parsed.User != nil {
|
||||
return "", errors.New("endpoint must not contain user information")
|
||||
}
|
||||
if parsed.RawQuery != "" || parsed.ForceQuery {
|
||||
return "", errors.New("endpoint must not contain a query string")
|
||||
}
|
||||
|
||||
return parsed.String(), nil
|
||||
}
|
||||
47
internal/domain/endpoint_test.go
Normal file
47
internal/domain/endpoint_test.go
Normal file
@@ -0,0 +1,47 @@
|
||||
package domain
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestNormalizeOpenAICompatibleBaseEndpoint(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
endpoint string
|
||||
want string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "http host", endpoint: "http://provider.example", want: "http://provider.example"},
|
||||
{name: "https nested path and whitespace", endpoint: " HTTPS://provider.example/api/openai/v1 ", want: "https://provider.example/api/openai/v1"},
|
||||
{name: "IPv4 host and port", endpoint: "http://127.0.0.1:8080/v1", want: "http://127.0.0.1:8080/v1"},
|
||||
{name: "IPv6 host and port", endpoint: "https://[::1]:8443/v1", want: "https://[::1]:8443/v1"},
|
||||
{name: "repeated trailing slashes", endpoint: "https://provider.example/v1///", want: "https://provider.example/v1///"},
|
||||
{name: "blank", endpoint: " \t\n ", wantErr: true},
|
||||
{name: "relative path", endpoint: "/api/v1", wantErr: true},
|
||||
{name: "scheme relative", endpoint: "//provider.example/v1", wantErr: true},
|
||||
{name: "missing host", endpoint: "https:///v1", wantErr: true},
|
||||
{name: "unsupported scheme", endpoint: "ftp://provider.example/v1", wantErr: true},
|
||||
{name: "user information", endpoint: "https://user:secret@provider.example/v1", wantErr: true},
|
||||
{name: "query", endpoint: "https://provider.example/v1?mode=chat", wantErr: true},
|
||||
{name: "empty query", endpoint: "https://provider.example/v1?", wantErr: true},
|
||||
{name: "fragment", endpoint: "https://provider.example/v1#chat", wantErr: true},
|
||||
{name: "empty fragment", endpoint: "https://provider.example/v1#", wantErr: true},
|
||||
{name: "malformed URL", endpoint: "https://provider.example/%zz", wantErr: true},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := NormalizeOpenAICompatibleBaseEndpoint(tc.endpoint)
|
||||
if tc.wantErr {
|
||||
if err == nil {
|
||||
t.Fatalf("expected endpoint error, got %q", got)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("normalize endpoint: %v", err)
|
||||
}
|
||||
if got != tc.want {
|
||||
t.Fatalf("normalized endpoint = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
34
internal/domain/execution_settings.go
Normal file
34
internal/domain/execution_settings.go
Normal file
@@ -0,0 +1,34 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"time"
|
||||
)
|
||||
|
||||
const maxExecutionTimeoutSeconds int64 = math.MaxInt64 / int64(time.Second)
|
||||
|
||||
// ValidateExecutionTargetSettings validates source-neutral execution-setting
|
||||
// invariants on a resolved target.
|
||||
func ValidateExecutionTargetSettings(target ExecutionTarget) error {
|
||||
if !isFinite(target.Temperature) || target.Temperature < 0 || target.Temperature > 2 {
|
||||
return errors.New("temperature must be finite and between 0 and 2")
|
||||
}
|
||||
if target.MaxTokens < 0 {
|
||||
return errors.New("max_tokens must be greater than or equal to 0")
|
||||
}
|
||||
if !isFinite(target.TopP) || target.TopP < 0 || target.TopP > 1 {
|
||||
return errors.New("top_p must be finite and between 0 and 1")
|
||||
}
|
||||
if target.TimeoutSeconds < 0 {
|
||||
return errors.New("timeout_seconds must be greater than or equal to 0")
|
||||
}
|
||||
if int64(target.TimeoutSeconds) > maxExecutionTimeoutSeconds {
|
||||
return errors.New("timeout_seconds exceeds the maximum supported duration")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isFinite(value float64) bool {
|
||||
return !math.IsNaN(value) && !math.IsInf(value, 0)
|
||||
}
|
||||
73
internal/domain/execution_settings_test.go
Normal file
73
internal/domain/execution_settings_test.go
Normal file
@@ -0,0 +1,73 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateExecutionTargetSettings(t *testing.T) {
|
||||
valid := ExecutionTarget{
|
||||
Temperature: 1,
|
||||
MaxTokens: 1,
|
||||
TopP: 0.5,
|
||||
TimeoutSeconds: 1,
|
||||
}
|
||||
|
||||
type testCase struct {
|
||||
name string
|
||||
change func(*ExecutionTarget)
|
||||
wantErr string
|
||||
}
|
||||
tests := []testCase{
|
||||
{name: "temperature lower boundary", change: func(v *ExecutionTarget) { v.Temperature = 0 }},
|
||||
{name: "temperature finite lower neighbor", change: func(v *ExecutionTarget) { v.Temperature = math.Nextafter(0, 1) }},
|
||||
{name: "temperature finite upper neighbor", change: func(v *ExecutionTarget) { v.Temperature = math.Nextafter(2, 0) }},
|
||||
{name: "temperature upper boundary", change: func(v *ExecutionTarget) { v.Temperature = 2 }},
|
||||
{name: "temperature below lower boundary", change: func(v *ExecutionTarget) { v.Temperature = math.Nextafter(0, math.Inf(-1)) }, wantErr: "temperature"},
|
||||
{name: "temperature above upper boundary", change: func(v *ExecutionTarget) { v.Temperature = math.Nextafter(2, math.Inf(1)) }, wantErr: "temperature"},
|
||||
{name: "temperature NaN", change: func(v *ExecutionTarget) { v.Temperature = math.NaN() }, wantErr: "temperature"},
|
||||
{name: "temperature positive infinity", change: func(v *ExecutionTarget) { v.Temperature = math.Inf(1) }, wantErr: "temperature"},
|
||||
{name: "temperature negative infinity", change: func(v *ExecutionTarget) { v.Temperature = math.Inf(-1) }, wantErr: "temperature"},
|
||||
{name: "max tokens lower boundary", change: func(v *ExecutionTarget) { v.MaxTokens = 0 }},
|
||||
{name: "max tokens finite neighbor", change: func(v *ExecutionTarget) { v.MaxTokens = 1 }},
|
||||
{name: "max tokens below lower boundary", change: func(v *ExecutionTarget) { v.MaxTokens = -1 }, wantErr: "max_tokens"},
|
||||
{name: "top p lower boundary", change: func(v *ExecutionTarget) { v.TopP = 0 }},
|
||||
{name: "top p finite lower neighbor", change: func(v *ExecutionTarget) { v.TopP = math.Nextafter(0, 1) }},
|
||||
{name: "top p finite upper neighbor", change: func(v *ExecutionTarget) { v.TopP = math.Nextafter(1, 0) }},
|
||||
{name: "top p upper boundary", change: func(v *ExecutionTarget) { v.TopP = 1 }},
|
||||
{name: "top p below lower boundary", change: func(v *ExecutionTarget) { v.TopP = math.Nextafter(0, math.Inf(-1)) }, wantErr: "top_p"},
|
||||
{name: "top p above upper boundary", change: func(v *ExecutionTarget) { v.TopP = math.Nextafter(1, math.Inf(1)) }, wantErr: "top_p"},
|
||||
{name: "top p NaN", change: func(v *ExecutionTarget) { v.TopP = math.NaN() }, wantErr: "top_p"},
|
||||
{name: "top p positive infinity", change: func(v *ExecutionTarget) { v.TopP = math.Inf(1) }, wantErr: "top_p"},
|
||||
{name: "top p negative infinity", change: func(v *ExecutionTarget) { v.TopP = math.Inf(-1) }, wantErr: "top_p"},
|
||||
{name: "timeout lower boundary", change: func(v *ExecutionTarget) { v.TimeoutSeconds = 0 }},
|
||||
{name: "timeout finite neighbor", change: func(v *ExecutionTarget) { v.TimeoutSeconds = 1 }},
|
||||
{name: "timeout below lower boundary", change: func(v *ExecutionTarget) { v.TimeoutSeconds = -1 }, wantErr: "timeout_seconds"},
|
||||
}
|
||||
if strconv.IntSize == 64 {
|
||||
durationLimit := maxExecutionTimeoutSeconds
|
||||
tests = append(tests,
|
||||
testCase{name: "timeout duration boundary", change: func(v *ExecutionTarget) { v.TimeoutSeconds = int(durationLimit) }},
|
||||
testCase{name: "timeout above duration boundary", change: func(v *ExecutionTarget) { v.TimeoutSeconds = int(durationLimit) + 1 }, wantErr: "timeout_seconds"},
|
||||
)
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
target := valid
|
||||
tt.change(&target)
|
||||
err := ValidateExecutionTargetSettings(target)
|
||||
if tt.wantErr == "" {
|
||||
if err != nil {
|
||||
t.Fatalf("validate execution settings: %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Fatalf("error = %v, want diagnostic containing %q", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
30
internal/domain/output_contract.go
Normal file
30
internal/domain/output_contract.go
Normal file
@@ -0,0 +1,30 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ValidateOutputContract validates source-neutral output-contract invariants.
|
||||
func ValidateOutputContract(contract OutputContract) error {
|
||||
switch contract.Format {
|
||||
case FormatText, FormatMarkdown, FormatJSON:
|
||||
default:
|
||||
return fmt.Errorf("invalid output format: %q", contract.Format)
|
||||
}
|
||||
|
||||
switch contract.ValidationMode {
|
||||
case ValidationNone, ValidationBasic, ValidationJSON, ValidationJSONSchema:
|
||||
default:
|
||||
return fmt.Errorf("invalid validation mode: %q", contract.ValidationMode)
|
||||
}
|
||||
|
||||
if contract.ValidationMode == ValidationJSONSchema && strings.TrimSpace(contract.SchemaPath) == "" {
|
||||
return errors.New("schema_path is required when validation_mode is json_schema")
|
||||
}
|
||||
if contract.RepairAttempts < 0 {
|
||||
return errors.New("repair_attempts must be greater than or equal to 0")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
68
internal/domain/output_contract_test.go
Normal file
68
internal/domain/output_contract_test.go
Normal file
@@ -0,0 +1,68 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateOutputContract(t *testing.T) {
|
||||
valid := OutputContract{
|
||||
Format: FormatText,
|
||||
ValidationMode: ValidationNone,
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
change func(*OutputContract)
|
||||
wantErr string
|
||||
}{
|
||||
{name: "text format", change: func(c *OutputContract) { c.Format = FormatText }},
|
||||
{name: "markdown format", change: func(c *OutputContract) { c.Format = FormatMarkdown }},
|
||||
{name: "json format", change: func(c *OutputContract) { c.Format = FormatJSON }},
|
||||
{name: "empty format", change: func(c *OutputContract) { c.Format = "" }, wantErr: "format"},
|
||||
{name: "unsupported format", change: func(c *OutputContract) { c.Format = OutputFormat("binary") }, wantErr: "format"},
|
||||
{name: "none validation", change: func(c *OutputContract) { c.ValidationMode = ValidationNone }},
|
||||
{name: "basic validation", change: func(c *OutputContract) { c.ValidationMode = ValidationBasic }},
|
||||
{name: "json validation", change: func(c *OutputContract) { c.ValidationMode = ValidationJSON }},
|
||||
{name: "json schema validation", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationJSONSchema
|
||||
c.SchemaPath = "schema.json"
|
||||
}},
|
||||
{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: "zero repair attempts", change: func(c *OutputContract) { c.RepairAttempts = 0 }},
|
||||
{name: "positive repair attempts", change: func(c *OutputContract) { c.RepairAttempts = 1 }},
|
||||
{name: "json schema empty path", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationJSONSchema
|
||||
c.SchemaPath = ""
|
||||
}, wantErr: "schema_path"},
|
||||
{name: "json schema whitespace path", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationJSONSchema
|
||||
c.SchemaPath = " \t "
|
||||
}, wantErr: "schema_path"},
|
||||
{name: "json schema nonblank path", change: func(c *OutputContract) {
|
||||
c.ValidationMode = ValidationJSONSchema
|
||||
c.SchemaPath = " schema.json "
|
||||
}},
|
||||
{name: "non-schema empty path", change: func(c *OutputContract) { c.SchemaPath = "" }},
|
||||
{name: "non-schema populated path", change: func(c *OutputContract) { c.SchemaPath = "ignored.json" }},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
contract := valid
|
||||
tt.change(&contract)
|
||||
err := ValidateOutputContract(contract)
|
||||
if tt.wantErr == "" {
|
||||
if err != nil {
|
||||
t.Fatalf("validate output contract: %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Fatalf("error = %v, want diagnostic containing %q", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,141 +0,0 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestPreparedRunJSONDoesNotIncludeSecretValues(t *testing.T) {
|
||||
const envName = "PROMPTKIT_TEST_API_KEY"
|
||||
const secret = "super-secret-value"
|
||||
t.Setenv(envName, secret)
|
||||
|
||||
prepared := PreparedRun{
|
||||
PromptID: "prompt.id",
|
||||
PromptVersion: "v1",
|
||||
PromptHash: "prompt-hash",
|
||||
SelectedProfileID: "local-fast",
|
||||
EffectiveModelParams: ExecutionTarget{
|
||||
Endpoint: "http://llm/v1",
|
||||
Model: "gpt-test",
|
||||
APIKeyEnv: envName,
|
||||
APIKey: secret,
|
||||
},
|
||||
InputHashes: map[string]string{"transcript": "hash-1"},
|
||||
RenderedPromptHash: "rendered-hash",
|
||||
Messages: []RenderedMessage{
|
||||
{Role: "system", Content: "You are helpful."},
|
||||
{Role: "user", Content: "Summarize this."},
|
||||
},
|
||||
}
|
||||
|
||||
b, err := json.Marshal(prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
|
||||
out := string(b)
|
||||
if strings.Contains(out, secret) {
|
||||
t.Fatalf("prepared run JSON unexpectedly contains secret value: %s", out)
|
||||
}
|
||||
if !strings.Contains(out, `"api_key_env":"`+envName+`"`) {
|
||||
t.Fatalf("prepared run JSON should include api_key_env name: %s", out)
|
||||
}
|
||||
|
||||
var top map[string]any
|
||||
if err := json.Unmarshal(b, &top); err != nil {
|
||||
t.Fatalf("unmarshal failed: %v", err)
|
||||
}
|
||||
|
||||
for _, forbidden := range []string{"raw_output", "validation", "artifact"} {
|
||||
if _, ok := top[forbidden]; ok {
|
||||
t.Fatalf("prepared run JSON should not include %q", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedRunJSONIncludesMessageCacheControlOnlyWhenPresent(t *testing.T) {
|
||||
prepared := PreparedRun{
|
||||
PromptID: "prompt.id",
|
||||
SelectedProfileID: "local-fast",
|
||||
EffectiveModelParams: ExecutionTarget{
|
||||
Endpoint: "http://llm/v1",
|
||||
Model: "gpt-test",
|
||||
},
|
||||
RenderedPromptHash: "rendered-hash",
|
||||
Messages: []RenderedMessage{
|
||||
{
|
||||
Role: "system",
|
||||
Content: "You are helpful.",
|
||||
CacheControl: &CacheControl{
|
||||
Type: CacheControlEphemeral,
|
||||
TTL: "1h",
|
||||
},
|
||||
},
|
||||
{Role: "user", Content: "Summarize this."},
|
||||
},
|
||||
}
|
||||
|
||||
b, err := json.Marshal(prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var decoded struct {
|
||||
Messages []map[string]any `json:"messages"`
|
||||
}
|
||||
if err := json.Unmarshal(b, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal failed: %v", err)
|
||||
}
|
||||
if len(decoded.Messages) != 2 {
|
||||
t.Fatalf("expected 2 messages, got %d", len(decoded.Messages))
|
||||
}
|
||||
|
||||
cacheControl, ok := decoded.Messages[0]["cache_control"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("expected cache_control on first message, got %#v", decoded.Messages[0])
|
||||
}
|
||||
if cacheControl["type"] != string(CacheControlEphemeral) || cacheControl["ttl"] != "1h" {
|
||||
t.Fatalf("unexpected cache_control payload: %#v", cacheControl)
|
||||
}
|
||||
if _, ok := decoded.Messages[1]["cache_control"]; ok {
|
||||
t.Fatalf("expected second message to omit cache_control, got %#v", decoded.Messages[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedRunJSONIncludesSessionIDOnlyWhenPresent(t *testing.T) {
|
||||
prepared := PreparedRun{
|
||||
PromptID: "prompt.id",
|
||||
SelectedProfileID: "local-fast",
|
||||
EffectiveModelParams: ExecutionTarget{
|
||||
Endpoint: "http://llm/v1",
|
||||
Model: "gpt-test",
|
||||
},
|
||||
SessionID: "session-123",
|
||||
RenderedPromptHash: "rendered-hash",
|
||||
Messages: []RenderedMessage{{Role: "user", Content: "Summarize this."}},
|
||||
}
|
||||
|
||||
b, err := json.Marshal(prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var decoded map[string]any
|
||||
if err := json.Unmarshal(b, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal failed: %v", err)
|
||||
}
|
||||
if decoded["session_id"] != "session-123" {
|
||||
t.Fatalf("expected session_id in prepared run JSON, got %#v", decoded["session_id"])
|
||||
}
|
||||
|
||||
prepared.SessionID = ""
|
||||
b, err = json.Marshal(prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
if strings.Contains(string(b), "session_id") {
|
||||
t.Fatalf("expected empty session_id to be omitted, got %s", b)
|
||||
}
|
||||
}
|
||||
@@ -8,6 +8,9 @@ import (
|
||||
|
||||
// NormalizeSessionID applies the shared session identifier rule.
|
||||
func NormalizeSessionID(raw string) (string, error) {
|
||||
if !utf8.ValidString(raw) {
|
||||
return "", fmt.Errorf("session_id must contain valid UTF-8")
|
||||
}
|
||||
normalized := strings.TrimSpace(raw)
|
||||
if normalized == "" {
|
||||
return "", nil
|
||||
|
||||
@@ -7,10 +7,10 @@ import (
|
||||
|
||||
func TestNormalizeSessionID(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
raw string
|
||||
want string
|
||||
wantErr bool
|
||||
name string
|
||||
raw string
|
||||
want string
|
||||
wantErrContains string
|
||||
}{
|
||||
{
|
||||
name: "trims surrounding Unicode whitespace",
|
||||
@@ -28,21 +28,24 @@ func TestNormalizeSessionID(t *testing.T) {
|
||||
want: strings.Repeat("界", SessionIDMaxLength),
|
||||
},
|
||||
{
|
||||
name: "one Unicode code point over maximum is rejected",
|
||||
raw: strings.Repeat("界", SessionIDMaxLength+1),
|
||||
wantErr: true,
|
||||
name: "one Unicode code point over maximum is rejected",
|
||||
raw: strings.Repeat("界", SessionIDMaxLength+1),
|
||||
wantErrContains: "exceeds maximum",
|
||||
},
|
||||
{name: "invalid UTF-8 before valid content", raw: string([]byte{0xff}) + "session", wantErrContains: "valid UTF-8"},
|
||||
{name: "invalid UTF-8 within valid content", raw: "ses" + string([]byte{0xff}) + "sion", wantErrContains: "valid UTF-8"},
|
||||
{name: "invalid UTF-8 after valid content", raw: "session" + string([]byte{0xff}), wantErrContains: "valid UTF-8"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := NormalizeSessionID(tt.raw)
|
||||
if tt.wantErr {
|
||||
if tt.wantErrContains != "" {
|
||||
if err == nil {
|
||||
t.Fatal("expected normalization error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "exceeds maximum") {
|
||||
t.Fatalf("expected useful length diagnostic, got %v", err)
|
||||
if !strings.Contains(err.Error(), tt.wantErrContains) {
|
||||
t.Fatalf("expected diagnostic containing %q, got %v", tt.wantErrContains, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -36,7 +36,8 @@ func FindYAMLFiles(ctx context.Context, root string) ([]string, error) {
|
||||
return files, err
|
||||
}
|
||||
|
||||
// FindFSYAMLFiles returns sorted paths for .yaml and .yml files under root in fsys.
|
||||
// FindFSYAMLFiles returns root itself when it names a file. For a directory
|
||||
// root, it returns sorted paths for .yaml and .yml files beneath that root.
|
||||
func FindFSYAMLFiles(ctx context.Context, fsys fs.FS, root string) ([]string, error) {
|
||||
cleanRoot := CleanFSRoot(root)
|
||||
var files []string
|
||||
@@ -52,6 +53,10 @@ func FindFSYAMLFiles(ctx context.Context, fsys fs.FS, root string) ([]string, er
|
||||
if d.IsDir() {
|
||||
return nil
|
||||
}
|
||||
if name == cleanRoot {
|
||||
files = append(files, name)
|
||||
return nil
|
||||
}
|
||||
if !IsYAMLFile(d.Name()) {
|
||||
return nil
|
||||
}
|
||||
@@ -71,10 +76,10 @@ func RelativePath(root string, filePath string) string {
|
||||
return filepath.Clean(rel)
|
||||
}
|
||||
|
||||
// CleanFSRoot normalizes a root path for use with fs.FS.
|
||||
// CleanFSRoot normalizes a root path for use with fs.FS while preserving
|
||||
// nonblank leading and trailing whitespace.
|
||||
func CleanFSRoot(root string) string {
|
||||
root = strings.TrimSpace(root)
|
||||
if root == "" || root == "." {
|
||||
if strings.TrimSpace(root) == "" || root == "." {
|
||||
return "."
|
||||
}
|
||||
return path.Clean(root)
|
||||
@@ -97,22 +102,21 @@ func DisplayPath(root string, name string) string {
|
||||
// ResolveFSPath resolves userPath from baseDir and keeps it inside root.
|
||||
func ResolveFSPath(root string, baseDir string, userPath string) (string, string, error) {
|
||||
cleanRoot := CleanFSRoot(root)
|
||||
cleanBase := path.Clean(strings.TrimSpace(baseDir))
|
||||
if cleanBase == "" {
|
||||
cleanBase := path.Clean(baseDir)
|
||||
if strings.TrimSpace(baseDir) == "" {
|
||||
cleanBase = cleanRoot
|
||||
}
|
||||
if !containsFSPath(cleanRoot, cleanBase) {
|
||||
return "", "", fmt.Errorf("base path %q is outside source root %q", cleanBase, cleanRoot)
|
||||
}
|
||||
|
||||
cleanUserPath := strings.TrimSpace(userPath)
|
||||
if cleanUserPath == "" {
|
||||
if strings.TrimSpace(userPath) == "" {
|
||||
return "", "", fmt.Errorf("path is required")
|
||||
}
|
||||
cleanUserPath = path.Clean(cleanUserPath)
|
||||
if path.IsAbs(cleanUserPath) {
|
||||
if path.IsAbs(userPath) {
|
||||
return "", "", fmt.Errorf("path %q must be relative", userPath)
|
||||
}
|
||||
cleanUserPath := path.Clean(userPath)
|
||||
|
||||
resolved := path.Clean(path.Join(cleanBase, cleanUserPath))
|
||||
if !containsFSPath(cleanRoot, resolved) {
|
||||
@@ -130,13 +134,6 @@ func containsFSPath(root string, name string) bool {
|
||||
return name == root || strings.HasPrefix(name, strings.TrimSuffix(root, "/")+"/")
|
||||
}
|
||||
|
||||
// Stem strips .yaml or .yml from a file name.
|
||||
func Stem(name string) string {
|
||||
name = strings.TrimSuffix(name, ".yaml")
|
||||
name = strings.TrimSuffix(name, ".yml")
|
||||
return name
|
||||
}
|
||||
|
||||
func IsYAMLFile(name string) bool {
|
||||
return strings.HasSuffix(name, ".yaml") || strings.HasSuffix(name, ".yml")
|
||||
}
|
||||
|
||||
@@ -54,7 +54,7 @@ func TestFindFSYAMLFilesNestedSortedAndFiltered(t *testing.T) {
|
||||
"other/ignored.yaml": &fstest.MapFile{Data: []byte("id: ignored")},
|
||||
}
|
||||
|
||||
got, err := FindFSYAMLFiles(context.Background(), fsys, " prompts ")
|
||||
got, err := FindFSYAMLFiles(context.Background(), fsys, "prompts")
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
@@ -98,8 +98,11 @@ func TestCleanFSRoot(t *testing.T) {
|
||||
want string
|
||||
}{
|
||||
{name: "empty", root: "", want: "."},
|
||||
{name: "whitespace only", root: " \t ", want: "."},
|
||||
{name: "dot", root: ".", want: "."},
|
||||
{name: "trimmed", root: " prompts/../profiles ", want: "profiles"},
|
||||
{name: "cleaned", root: "prompts/../profiles", want: "profiles"},
|
||||
{name: "leading whitespace preserved", root: " profiles", want: " profiles"},
|
||||
{name: "trailing whitespace preserved", root: "profiles ", want: "profiles "},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
@@ -158,6 +161,22 @@ func TestResolveFSPath(t *testing.T) {
|
||||
wantPath: "prompts/shared/user.tmpl",
|
||||
wantDisplay: "shared/user.tmpl",
|
||||
},
|
||||
{
|
||||
name: "leading whitespace preserved",
|
||||
root: "prompts",
|
||||
baseDir: "prompts/nested",
|
||||
userPath: " user.tmpl",
|
||||
wantPath: "prompts/nested/ user.tmpl",
|
||||
wantDisplay: "nested/ user.tmpl",
|
||||
},
|
||||
{
|
||||
name: "trailing whitespace preserved",
|
||||
root: "prompts",
|
||||
baseDir: "prompts/nested",
|
||||
userPath: "user.tmpl ",
|
||||
wantPath: "prompts/nested/user.tmpl ",
|
||||
wantDisplay: "nested/user.tmpl ",
|
||||
},
|
||||
{
|
||||
name: "escape rejected",
|
||||
root: "prompts",
|
||||
@@ -218,26 +237,6 @@ func TestResolveFSPath(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestStemStripsYAMLExtensions(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
want string
|
||||
}{
|
||||
{name: "yaml", in: "prompt.yaml", want: "prompt"},
|
||||
{name: "yml", in: "profile.yml", want: "profile"},
|
||||
{name: "other", in: "file.txt", want: "file.txt"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := Stem(tc.in); got != tc.want {
|
||||
t.Fatalf("expected %q, got %q", tc.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsYAMLFile(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
// Package jsonvalue validates and defensively copies JSON-compatible value
|
||||
// trees used by public configuration and request boundaries.
|
||||
// Package jsonvalue validates and defensively copies bounded JSON-compatible
|
||||
// value trees used by configuration, request, and prepared-state boundaries.
|
||||
package jsonvalue
|
||||
|
||||
import (
|
||||
@@ -8,23 +8,39 @@ import (
|
||||
"math"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
const maxSafeJSONInteger = 1<<53 - 1
|
||||
const (
|
||||
maxContainerDepth = 100
|
||||
maxProducedNodes = 100_000
|
||||
)
|
||||
|
||||
type visit struct {
|
||||
typ reflect.Type
|
||||
ptr uintptr
|
||||
}
|
||||
|
||||
type traversalState struct {
|
||||
active map[visit]struct{}
|
||||
producedNodes int
|
||||
}
|
||||
|
||||
// Copy validates and deeply copies a JSON-compatible value while preserving
|
||||
// compatible concrete map, slice, array, scalar, and number types. It rejects
|
||||
// cycles and values that exceed the package's traversal limits.
|
||||
func Copy(src any) (any, error) {
|
||||
return copyValue(reflect.ValueOf(src), "value", newTraversalState(), true, 0)
|
||||
}
|
||||
|
||||
// CopyMap validates and deeply copies an extra-parameter map while preserving
|
||||
// compatible concrete map, slice, array, scalar, and number types.
|
||||
// compatible concrete map, slice, array, scalar, and number types. It rejects
|
||||
// empty object keys, cycles, and values that exceed the package's traversal
|
||||
// limits.
|
||||
func CopyMap(src map[string]any) (map[string]any, error) {
|
||||
if src == nil {
|
||||
return nil, nil
|
||||
}
|
||||
copied, err := copyValue(reflect.ValueOf(src), "extra_params", make(map[visit]struct{}))
|
||||
copied, err := copyValue(reflect.ValueOf(src), "extra_params", newTraversalState(), false, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -35,88 +51,121 @@ func CopyMap(src map[string]any) (map[string]any, error) {
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func copyValue(value reflect.Value, path string, seen map[visit]struct{}) (any, error) {
|
||||
if !value.IsValid() {
|
||||
func copyValue(
|
||||
value reflect.Value,
|
||||
path string,
|
||||
state *traversalState,
|
||||
allowEmptyMapKeys bool,
|
||||
containerDepth int,
|
||||
) (any, error) {
|
||||
resolved, cleanup, isNull, err := state.resolveIndirection(value, path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer cleanup()
|
||||
if isNull {
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
if value.Kind() == reflect.Interface {
|
||||
if value.IsNil() {
|
||||
return nil, nil
|
||||
}
|
||||
return copyValue(value.Elem(), path, seen)
|
||||
}
|
||||
value = resolved
|
||||
if !value.CanInterface() {
|
||||
return nil, fmt.Errorf("%s: value cannot be copied", path)
|
||||
}
|
||||
if number, ok := value.Interface().(json.Number); ok {
|
||||
if _, err := json.Marshal(number); err != nil {
|
||||
if !validJSONNumber(number) {
|
||||
return nil, fmt.Errorf("%s: invalid JSON number", path)
|
||||
}
|
||||
f, err := strconv.ParseFloat(number.String(), 64)
|
||||
if err != nil || math.IsNaN(f) || math.IsInf(f, 0) {
|
||||
return nil, fmt.Errorf("%s: invalid JSON number", path)
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return number, nil
|
||||
}
|
||||
|
||||
switch value.Kind() {
|
||||
case reflect.Bool, reflect.String:
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return value.Interface(), nil
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
if value.Int() < -maxSafeJSONInteger || value.Int() > maxSafeJSONInteger {
|
||||
return nil, fmt.Errorf("%s: integer is outside the JSON-safe range", path)
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return value.Interface(), nil
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
|
||||
if value.Uint() > maxSafeJSONInteger {
|
||||
return nil, fmt.Errorf("%s: integer is outside the JSON-safe range", path)
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return value.Interface(), nil
|
||||
case reflect.Float32, reflect.Float64:
|
||||
number := value.Convert(reflect.TypeOf(float64(0))).Float()
|
||||
number := value.Float()
|
||||
if math.IsNaN(number) || math.IsInf(number, 0) {
|
||||
return nil, fmt.Errorf("%s: floating-point value must be finite", path)
|
||||
}
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return value.Interface(), nil
|
||||
case reflect.Pointer:
|
||||
case reflect.Map:
|
||||
if value.IsNil() {
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
current := visit{typ: value.Type(), ptr: value.Pointer()}
|
||||
if _, ok := seen[current]; ok {
|
||||
return nil, fmt.Errorf("%s: cyclic value is not supported", path)
|
||||
nextDepth, err := state.enterContainer(path, containerDepth)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
seen[current] = struct{}{}
|
||||
defer delete(seen, current)
|
||||
return copyValue(value.Elem(), path, seen)
|
||||
case reflect.Map:
|
||||
return copyMapValue(value, path, seen)
|
||||
return copyMapValue(value, path, state, allowEmptyMapKeys, nextDepth)
|
||||
case reflect.Slice:
|
||||
if value.IsNil() {
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
return copySequenceValue(value, path, seen)
|
||||
nextDepth, err := state.enterContainer(path, containerDepth)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return copySequenceValue(value, path, state, allowEmptyMapKeys, nextDepth)
|
||||
case reflect.Array:
|
||||
return copySequenceValue(value, path, seen)
|
||||
nextDepth, err := state.enterContainer(path, containerDepth)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return copySequenceValue(value, path, state, allowEmptyMapKeys, nextDepth)
|
||||
default:
|
||||
return nil, fmt.Errorf("%s: unsupported JSON value type %s", path, value.Type())
|
||||
}
|
||||
}
|
||||
|
||||
func copyMapValue(value reflect.Value, path string, seen map[visit]struct{}) (any, error) {
|
||||
if value.IsNil() {
|
||||
return nil, nil
|
||||
}
|
||||
func copyMapValue(
|
||||
value reflect.Value,
|
||||
path string,
|
||||
state *traversalState,
|
||||
allowEmptyMapKeys bool,
|
||||
containerDepth int,
|
||||
) (any, error) {
|
||||
if value.Type().Key().Kind() != reflect.String {
|
||||
return nil, fmt.Errorf("%s: map key type %s is not supported", path, value.Type().Key())
|
||||
}
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := state.ensureChildCapacity(path, value.Len()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
current := visit{typ: value.Type(), ptr: value.Pointer()}
|
||||
if _, ok := seen[current]; ok {
|
||||
if _, ok := state.active[current]; ok {
|
||||
return nil, fmt.Errorf("%s: cyclic value is not supported", path)
|
||||
}
|
||||
seen[current] = struct{}{}
|
||||
defer delete(seen, current)
|
||||
state.active[current] = struct{}{}
|
||||
defer delete(state.active, current)
|
||||
|
||||
keys := value.MapKeys()
|
||||
sort.Slice(keys, func(i, j int) bool {
|
||||
@@ -133,10 +182,16 @@ func copyMapValue(value reflect.Value, path string, seen map[visit]struct{}) (an
|
||||
elementType := value.Type().Elem()
|
||||
for _, key := range keys {
|
||||
name := key.String()
|
||||
if name == "" {
|
||||
if name == "" && !allowEmptyMapKeys {
|
||||
return nil, fmt.Errorf("%s: map key must not be empty", path)
|
||||
}
|
||||
copied, err := copyValue(value.MapIndex(key), path+"."+name, seen)
|
||||
copied, err := copyValue(
|
||||
value.MapIndex(key),
|
||||
path+"."+name,
|
||||
state,
|
||||
allowEmptyMapKeys,
|
||||
containerDepth,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -171,22 +226,41 @@ func copyMapValue(value reflect.Value, path string, seen map[visit]struct{}) (an
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func copySequenceValue(value reflect.Value, path string, seen map[visit]struct{}) (any, error) {
|
||||
func copySequenceValue(
|
||||
value reflect.Value,
|
||||
path string,
|
||||
state *traversalState,
|
||||
allowEmptyMapKeys bool,
|
||||
containerDepth int,
|
||||
) (any, error) {
|
||||
if err := state.produceNode(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := state.ensureChildCapacity(path, value.Len()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var current visit
|
||||
if value.Kind() == reflect.Slice {
|
||||
current = visit{typ: value.Type(), ptr: value.Pointer()}
|
||||
if _, ok := seen[current]; ok {
|
||||
if _, ok := state.active[current]; ok {
|
||||
return nil, fmt.Errorf("%s: cyclic value is not supported", path)
|
||||
}
|
||||
seen[current] = struct{}{}
|
||||
defer delete(seen, current)
|
||||
state.active[current] = struct{}{}
|
||||
defer delete(state.active, current)
|
||||
}
|
||||
|
||||
values := make([]any, value.Len())
|
||||
preserveType := true
|
||||
elementType := value.Type().Elem()
|
||||
for i := 0; i < value.Len(); i++ {
|
||||
copied, err := copyValue(value.Index(i), fmt.Sprintf("%s[%d]", path, i), seen)
|
||||
copied, err := copyValue(
|
||||
value.Index(i),
|
||||
fmt.Sprintf("%s[%d]", path, i),
|
||||
state,
|
||||
allowEmptyMapKeys,
|
||||
containerDepth,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -222,6 +296,73 @@ func copySequenceValue(value reflect.Value, path string, seen map[visit]struct{}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func validJSONNumber(number json.Number) bool {
|
||||
var parsed json.Number
|
||||
if err := json.Unmarshal([]byte(number.String()), &parsed); err != nil {
|
||||
return false
|
||||
}
|
||||
return parsed.String() == number.String()
|
||||
}
|
||||
|
||||
func newTraversalState() *traversalState {
|
||||
return &traversalState{active: make(map[visit]struct{})}
|
||||
}
|
||||
|
||||
func (state *traversalState) produceNode(path string) error {
|
||||
if state.producedNodes >= maxProducedNodes {
|
||||
return fmt.Errorf("%s: JSON value work limit exceeded", path)
|
||||
}
|
||||
state.producedNodes++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (state *traversalState) enterContainer(path string, depth int) (int, error) {
|
||||
depth++
|
||||
if depth > maxContainerDepth {
|
||||
return 0, fmt.Errorf("%s: JSON container depth limit exceeded", path)
|
||||
}
|
||||
return depth, nil
|
||||
}
|
||||
|
||||
func (state *traversalState) ensureChildCapacity(path string, count int) error {
|
||||
if count > maxProducedNodes-state.producedNodes {
|
||||
return fmt.Errorf("%s: JSON value work limit exceeded", path)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (state *traversalState) resolveIndirection(
|
||||
value reflect.Value,
|
||||
path string,
|
||||
) (reflect.Value, func(), bool, error) {
|
||||
var visits []visit
|
||||
cleanup := func() {
|
||||
for _, current := range visits {
|
||||
delete(state.active, current)
|
||||
}
|
||||
}
|
||||
|
||||
for value.IsValid() && (value.Kind() == reflect.Interface || value.Kind() == reflect.Pointer) {
|
||||
if value.IsNil() {
|
||||
return reflect.Value{}, cleanup, true, nil
|
||||
}
|
||||
if value.Kind() == reflect.Pointer {
|
||||
current := visit{typ: value.Type(), ptr: value.Pointer()}
|
||||
if _, ok := state.active[current]; ok {
|
||||
cleanup()
|
||||
return reflect.Value{}, nil, false, fmt.Errorf("%s: cyclic value is not supported", path)
|
||||
}
|
||||
state.active[current] = struct{}{}
|
||||
visits = append(visits, current)
|
||||
}
|
||||
value = value.Elem()
|
||||
}
|
||||
if !value.IsValid() {
|
||||
return reflect.Value{}, cleanup, true, nil
|
||||
}
|
||||
return value, cleanup, false, nil
|
||||
}
|
||||
|
||||
func canAssignNil(typ reflect.Type) bool {
|
||||
switch typ.Kind() {
|
||||
case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice:
|
||||
|
||||
@@ -1,97 +1,332 @@
|
||||
package jsonvalue_test
|
||||
package jsonvalue
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"math"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
)
|
||||
|
||||
func TestCopyMapPreservesTypesAndIsolatesMutations(t *testing.T) {
|
||||
nested := map[string]int{"limit": 2}
|
||||
sequence := []string{"one", "two"}
|
||||
input := map[string]any{
|
||||
"count": int64(7),
|
||||
"number": json.Number("-1.25e+2"),
|
||||
"nested": nested,
|
||||
"sequence": sequence,
|
||||
}
|
||||
|
||||
copied, err := jsonvalue.CopyMap(input)
|
||||
if err != nil {
|
||||
t.Fatalf("copy map: %v", err)
|
||||
}
|
||||
nested["limit"] = 99
|
||||
sequence[0] = "changed"
|
||||
input["added"] = true
|
||||
|
||||
if got, ok := copied["count"].(int64); !ok || got != 7 {
|
||||
t.Fatalf("integer type or value changed: %#v", copied["count"])
|
||||
}
|
||||
if got, ok := copied["number"].(json.Number); !ok || got != "-1.25e+2" {
|
||||
t.Fatalf("JSON number type or value changed: %#v", copied["number"])
|
||||
}
|
||||
if got := copied["nested"].(map[string]int)["limit"]; got != 2 {
|
||||
t.Fatalf("nested map was not isolated: %d", got)
|
||||
}
|
||||
if got := copied["sequence"].([]string)[0]; got != "one" {
|
||||
t.Fatalf("sequence was not isolated: %q", got)
|
||||
}
|
||||
if _, ok := copied["added"]; ok {
|
||||
t.Fatalf("top-level map was not isolated: %#v", copied)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyMapRejectsInvalidValues(t *testing.T) {
|
||||
cyclicMap := map[string]any{}
|
||||
cyclicMap["self"] = cyclicMap
|
||||
cyclicSlice := []any{nil}
|
||||
cyclicSlice[0] = cyclicSlice
|
||||
type (
|
||||
namedBool bool
|
||||
namedString string
|
||||
namedInt64 int64
|
||||
namedUint64 uint64
|
||||
namedFloat32 float32
|
||||
namedFloat64 float64
|
||||
namedKey string
|
||||
namedMap map[namedKey]namedInt64
|
||||
namedSlice []namedString
|
||||
namedArray [1]map[string]int
|
||||
)
|
||||
|
||||
func TestCopyPreservesSupportedScalarAndNumberTypes(t *testing.T) {
|
||||
maxInt := int(^uint(0) >> 1)
|
||||
minInt := -maxInt - 1
|
||||
tests := []struct {
|
||||
name string
|
||||
value any
|
||||
}{
|
||||
{name: "empty nested key", value: map[string]int{"": 1}},
|
||||
{name: "non-string map key", value: map[int]string{1: "one"}},
|
||||
{name: "unsupported value", value: make(chan int)},
|
||||
{name: "cyclic map", value: cyclicMap},
|
||||
{name: "cyclic slice", value: cyclicSlice},
|
||||
{name: "NaN", value: math.NaN()},
|
||||
{name: "positive infinity", value: math.Inf(1)},
|
||||
{name: "unsafe signed integer", value: int64(1 << 53)},
|
||||
{name: "unsafe unsigned integer", value: uint64(1 << 53)},
|
||||
{name: "bool", value: true},
|
||||
{name: "named bool", value: namedBool(true)},
|
||||
{name: "string", value: "value"},
|
||||
{name: "named string", value: namedString("value")},
|
||||
{name: "int", value: minInt},
|
||||
{name: "int8", value: int8(-1 << 7)},
|
||||
{name: "int16", value: int16(-1 << 15)},
|
||||
{name: "int32", value: int32(-1 << 31)},
|
||||
{name: "int64", value: int64(-1 << 63)},
|
||||
{name: "named int64", value: namedInt64(1<<63 - 1)},
|
||||
{name: "uint", value: ^uint(0)},
|
||||
{name: "uint8", value: ^uint8(0)},
|
||||
{name: "uint16", value: ^uint16(0)},
|
||||
{name: "uint32", value: ^uint32(0)},
|
||||
{name: "uint64", value: ^uint64(0)},
|
||||
{name: "uintptr", value: ^uintptr(0)},
|
||||
{name: "named uint64", value: namedUint64(^uint64(0))},
|
||||
{name: "float32", value: float32(1.25)},
|
||||
{name: "float64", value: float64(-2.5e100)},
|
||||
{name: "named float32", value: namedFloat32(3.5)},
|
||||
{name: "named float64", value: namedFloat64(-4.5e200)},
|
||||
{name: "JSON number integer", value: json.Number("18446744073709551615")},
|
||||
{name: "JSON number fraction", value: json.Number("-1.25e+2")},
|
||||
{name: "JSON number beyond float64", value: json.Number("1e9999")},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if _, err := jsonvalue.CopyMap(map[string]any{"value": tc.value}); err == nil {
|
||||
got, err := Copy(tc.value)
|
||||
if err != nil {
|
||||
t.Fatalf("copy value: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(got, tc.value) {
|
||||
t.Fatalf("value or concrete type changed: got %#v (%T), want %#v (%T)", got, got, tc.value, tc.value)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyRejectsInvalidNumbers(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
value any
|
||||
}{
|
||||
{name: "float32 NaN", value: float32(math.NaN())},
|
||||
{name: "float64 NaN", value: math.NaN()},
|
||||
{name: "named float NaN", value: namedFloat64(math.NaN())},
|
||||
{name: "positive infinity", value: math.Inf(1)},
|
||||
{name: "negative infinity", value: math.Inf(-1)},
|
||||
{name: "empty JSON number", value: json.Number("")},
|
||||
{name: "leading zero JSON number", value: json.Number("01")},
|
||||
{name: "leading plus JSON number", value: json.Number("+1")},
|
||||
{name: "trailing decimal JSON number", value: json.Number("1.")},
|
||||
{name: "leading decimal JSON number", value: json.Number(".1")},
|
||||
{name: "non-number JSON number", value: json.Number("NaN")},
|
||||
{name: "spaced JSON number", value: json.Number(" 1")},
|
||||
{name: "quoted JSON number", value: json.Number(`"1"`)},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if _, err := Copy(tc.value); err == nil {
|
||||
t.Fatal("expected validation error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyMapValidatesJSONNumberSyntaxAndRange(t *testing.T) {
|
||||
for _, number := range []json.Number{"0", "-1", "1.25", "-1.25e+2"} {
|
||||
t.Run("valid "+number.String(), func(t *testing.T) {
|
||||
got, err := jsonvalue.CopyMap(map[string]any{"value": number})
|
||||
func TestCopyPreservesCompatibleCollectionsAndNilEmptyDistinctions(t *testing.T) {
|
||||
collections := []struct {
|
||||
name string
|
||||
value any
|
||||
}{
|
||||
{name: "unnamed map", value: map[string]int{"limit": 2}},
|
||||
{name: "named map", value: namedMap{"limit": 2}},
|
||||
{name: "unnamed slice", value: []string{"one", "two"}},
|
||||
{name: "named slice", value: namedSlice{"one", "two"}},
|
||||
{name: "unnamed array", value: [2]int{1, 2}},
|
||||
{name: "named array", value: namedArray{{"limit": 2}}},
|
||||
{name: "empty map", value: map[string]int{}},
|
||||
{name: "empty named map", value: namedMap{}},
|
||||
{name: "empty slice", value: []string{}},
|
||||
{name: "empty named slice", value: namedSlice{}},
|
||||
{name: "empty array", value: [0]string{}},
|
||||
}
|
||||
|
||||
for _, tc := range collections {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := Copy(tc.value)
|
||||
if err != nil {
|
||||
t.Fatalf("copy valid JSON number: %v", err)
|
||||
t.Fatalf("copy collection: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(got["value"], number) {
|
||||
t.Fatalf("JSON number changed: got %#v want %#v", got["value"], number)
|
||||
if !reflect.DeepEqual(got, tc.value) || reflect.TypeOf(got) != reflect.TypeOf(tc.value) {
|
||||
t.Fatalf("collection changed: got %#v (%T), want %#v (%T)", got, got, tc.value, tc.value)
|
||||
}
|
||||
kind := reflect.ValueOf(got).Kind()
|
||||
if (kind == reflect.Map || kind == reflect.Slice) && reflect.ValueOf(got).IsNil() {
|
||||
t.Fatal("non-nil collection became nil")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
for _, number := range []json.Number{"", "01", "+1", "1.", ".1", "1e9999", "not-a-number"} {
|
||||
t.Run("invalid "+number.String(), func(t *testing.T) {
|
||||
if _, err := jsonvalue.CopyMap(map[string]any{"value": number}); err == nil {
|
||||
t.Fatal("expected invalid JSON number error")
|
||||
var nilMap map[string]int
|
||||
var nilSlice []string
|
||||
var nilPointer *namedInt64
|
||||
for _, value := range []any{nil, nilMap, nilSlice, nilPointer} {
|
||||
got, err := Copy(value)
|
||||
if err != nil {
|
||||
t.Fatalf("copy null value: %v", err)
|
||||
}
|
||||
if got != nil {
|
||||
t.Fatalf("null value became %#v (%T)", got, got)
|
||||
}
|
||||
}
|
||||
|
||||
gotNil, err := CopyMap(nil)
|
||||
if err != nil || gotNil != nil {
|
||||
t.Fatalf("nil CopyMap result = %#v, %v", gotNil, err)
|
||||
}
|
||||
gotEmpty, err := CopyMap(map[string]any{})
|
||||
if err != nil || gotEmpty == nil || len(gotEmpty) != 0 {
|
||||
t.Fatalf("empty CopyMap result = %#v, %v", gotEmpty, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyHandlesIndirectionAndIsolatesNestedMutations(t *testing.T) {
|
||||
integer := namedInt64(7)
|
||||
nestedMap := namedMap{"limit": 2}
|
||||
nestedSlice := namedSlice{"original"}
|
||||
nestedArray := namedArray{{"limit": 3}}
|
||||
shared := []any{map[string]int{"value": 4}}
|
||||
input := map[string]any{
|
||||
"integer": &integer,
|
||||
"map": nestedMap,
|
||||
"slice": nestedSlice,
|
||||
"array": nestedArray,
|
||||
"first": shared,
|
||||
"second": shared,
|
||||
}
|
||||
|
||||
copiedValue, err := Copy(input)
|
||||
if err != nil {
|
||||
t.Fatalf("copy mixed tree: %v", err)
|
||||
}
|
||||
copied := copiedValue.(map[string]any)
|
||||
nestedMap["limit"] = 20
|
||||
nestedSlice[0] = "changed"
|
||||
nestedArray[0]["limit"] = 30
|
||||
shared[0].(map[string]int)["value"] = 40
|
||||
|
||||
if got, ok := copied["integer"].(namedInt64); !ok || got != 7 {
|
||||
t.Fatalf("pointer target changed: %#v", copied["integer"])
|
||||
}
|
||||
if got := copied["map"].(namedMap)["limit"]; got != 2 {
|
||||
t.Fatalf("nested map aliased input: %d", got)
|
||||
}
|
||||
if got := copied["slice"].(namedSlice)[0]; got != "original" {
|
||||
t.Fatalf("nested slice aliased input: %q", got)
|
||||
}
|
||||
if got := copied["array"].(namedArray)[0]["limit"]; got != 3 {
|
||||
t.Fatalf("nested array aliased input: %d", got)
|
||||
}
|
||||
first := copied["first"].([]any)
|
||||
second := copied["second"].([]any)
|
||||
if got := first[0].(map[string]int)["value"]; got != 4 {
|
||||
t.Fatalf("shared child aliased input: %d", got)
|
||||
}
|
||||
first[0].(map[string]int)["value"] = 99
|
||||
if got := second[0].(map[string]int)["value"]; got != 4 {
|
||||
t.Fatalf("repeated acyclic value shared copied output: %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyAndCopyMapApplyDistinctEmptyKeyRules(t *testing.T) {
|
||||
nested := map[string]any{"": []any{"original"}}
|
||||
copiedValue, err := Copy(nested)
|
||||
if err != nil {
|
||||
t.Fatalf("Copy rejected empty schema key: %v", err)
|
||||
}
|
||||
nested[""].([]any)[0] = "changed"
|
||||
if got := copiedValue.(map[string]any)[""].([]any)[0]; got != "original" {
|
||||
t.Fatalf("copied schema value was not isolated: %v", got)
|
||||
}
|
||||
|
||||
_, err = CopyMap(map[string]any{"nested": map[string]any{"": true}})
|
||||
if err == nil || !strings.Contains(err.Error(), "extra_params.nested") {
|
||||
t.Fatalf("CopyMap empty-key error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyRejectsUnsupportedValuesAndActiveCycles(t *testing.T) {
|
||||
cyclicMap := map[string]any{}
|
||||
cyclicMap["self"] = cyclicMap
|
||||
cyclicSlice := []any{nil}
|
||||
cyclicSlice[0] = cyclicSlice
|
||||
var cyclicPointer any
|
||||
cyclicPointer = &cyclicPointer
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value any
|
||||
wantPath string
|
||||
}{
|
||||
{name: "non-string map key", value: map[int]string{1: "one"}, wantPath: "value"},
|
||||
{name: "unsupported channel", value: make(chan int), wantPath: "value"},
|
||||
{name: "deterministic map path", value: map[string]any{"z": make(chan int), "a": make(chan int)}, wantPath: "value.a"},
|
||||
{name: "cyclic map", value: cyclicMap, wantPath: "value.self"},
|
||||
{name: "cyclic slice", value: cyclicSlice, wantPath: "value[0]"},
|
||||
{name: "cyclic pointer", value: cyclicPointer, wantPath: "value"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := Copy(tc.value)
|
||||
if err == nil || !strings.Contains(err.Error(), tc.wantPath) {
|
||||
t.Fatalf("error = %v, want structural path %q", err, tc.wantPath)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyEnforcesContainerDepth(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
depth int
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "just below", depth: maxContainerDepth - 1},
|
||||
{name: "at limit", depth: maxContainerDepth},
|
||||
{name: "over limit", depth: maxContainerDepth + 1, wantErr: true},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := Copy(alternatingContainers(tc.depth))
|
||||
if tc.wantErr {
|
||||
if err == nil || !strings.HasPrefix(err.Error(), "value") || !strings.Contains(err.Error(), "container depth limit") {
|
||||
t.Fatalf("depth error = %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("copy depth %d: %v", tc.depth, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyEnforcesProducedNodeBudgetForRepeatedAcyclicValues(t *testing.T) {
|
||||
shared := []any{true}
|
||||
sharedOccurrences := (maxProducedNodes - 2) / 2
|
||||
justBelow := repeatedValues(shared, sharedOccurrences, 0)
|
||||
atLimit := repeatedValues(shared, sharedOccurrences, 1)
|
||||
overLimit := repeatedValues(shared, sharedOccurrences, 2)
|
||||
|
||||
for name, value := range map[string]any{
|
||||
"just below": justBelow,
|
||||
"at limit": atLimit,
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if _, err := Copy(value); err != nil {
|
||||
t.Fatalf("copy value within work budget: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
_, err := Copy(overLimit)
|
||||
if err == nil || !strings.HasPrefix(err.Error(), "value[") || !strings.Contains(err.Error(), "value work limit") {
|
||||
t.Fatalf("work-budget error = %v", err)
|
||||
}
|
||||
|
||||
_, err = Copy(make([]any, maxProducedNodes))
|
||||
if err == nil || !strings.HasPrefix(err.Error(), "value:") || !strings.Contains(err.Error(), "value work limit") {
|
||||
t.Fatalf("flat work-budget error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func alternatingContainers(depth int) any {
|
||||
var value any = true
|
||||
for level := 0; level < depth; level++ {
|
||||
switch level % 3 {
|
||||
case 0:
|
||||
value = map[string]any{"child": value}
|
||||
case 1:
|
||||
value = []any{value}
|
||||
default:
|
||||
value = [1]any{value}
|
||||
}
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func repeatedValues(shared []any, occurrences, leadingScalars int) []any {
|
||||
values := make([]any, 0, leadingScalars+occurrences)
|
||||
for i := 0; i < leadingScalars; i++ {
|
||||
values = append(values, false)
|
||||
}
|
||||
for i := 0; i < occurrences; i++ {
|
||||
values = append(values, shared)
|
||||
}
|
||||
return values
|
||||
}
|
||||
|
||||
@@ -25,6 +25,20 @@ var (
|
||||
ErrMalformedResponse = errors.New("malformed llm response")
|
||||
)
|
||||
|
||||
const maxOpenAIChatResponseBytes int64 = 16 << 20
|
||||
|
||||
type requestFailedError struct {
|
||||
cause error
|
||||
}
|
||||
|
||||
func (e *requestFailedError) Error() string {
|
||||
return ErrRequestFailed.Error()
|
||||
}
|
||||
|
||||
func (e *requestFailedError) Unwrap() []error {
|
||||
return []error{ErrRequestFailed, e.cause}
|
||||
}
|
||||
|
||||
type OpenAICompatibleConfig struct {
|
||||
BaseURL string
|
||||
Model string
|
||||
@@ -39,9 +53,11 @@ type OpenAICompatibleClient struct {
|
||||
}
|
||||
|
||||
func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) {
|
||||
baseURL := strings.TrimSpace(cfg.BaseURL)
|
||||
if baseURL != "" {
|
||||
if _, err := url.ParseRequestURI(baseURL); err != nil {
|
||||
baseURL := ""
|
||||
if strings.TrimSpace(cfg.BaseURL) != "" {
|
||||
var err error
|
||||
baseURL, err = domain.NormalizeOpenAICompatibleBaseEndpoint(cfg.BaseURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err)
|
||||
}
|
||||
}
|
||||
@@ -63,25 +79,29 @@ func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleCli
|
||||
}
|
||||
|
||||
return &OpenAICompatibleClient{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
baseURL: baseURL,
|
||||
defaultModel: cfg.Model,
|
||||
httpClient: client,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) {
|
||||
if req.Target.TimeoutSeconds < 0 {
|
||||
return nil, fmt.Errorf("%w: timeout_seconds must be greater than or equal to 0", ErrInvalidRequest)
|
||||
if err := domain.ValidateExecutionTargetSettings(req.Target); err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
endpoint := strings.TrimSpace(req.Target.Endpoint)
|
||||
if endpoint == "" {
|
||||
endpoint = c.baseURL
|
||||
selectedEndpoint := req.Target.Endpoint
|
||||
if strings.TrimSpace(selectedEndpoint) == "" {
|
||||
selectedEndpoint = c.baseURL
|
||||
}
|
||||
if endpoint == "" {
|
||||
return nil, fmt.Errorf("%w: endpoint is required", ErrInvalidRequest)
|
||||
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(selectedEndpoint)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: invalid endpoint: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
endpoint, err = url.JoinPath(endpoint, defaults.OpenAIChatCompletionsPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: invalid endpoint path: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
endpoint = strings.TrimRight(endpoint, "/") + defaults.OpenAIChatCompletionsPath
|
||||
|
||||
wireReq, err := openAIChatRequestFromGenerateRequest(req, c.defaultModel)
|
||||
if err != nil {
|
||||
@@ -130,7 +150,7 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
||||
|
||||
httpResp, err := httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrRequestFailed, err)
|
||||
return nil, &requestFailedError{cause: err}
|
||||
}
|
||||
defer httpResp.Body.Close()
|
||||
|
||||
@@ -138,10 +158,13 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(httpResp.Body, 4096))
|
||||
return nil, fmt.Errorf("%w: status=%d", ErrUnexpectedStatus, httpResp.StatusCode)
|
||||
}
|
||||
if httpResp.ContentLength > maxOpenAIChatResponseBytes {
|
||||
return nil, openAIChatResponseTooLargeError()
|
||||
}
|
||||
|
||||
var wireResp openAIChatResponse
|
||||
if err := json.NewDecoder(httpResp.Body).Decode(&wireResp); err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to decode response: %v", ErrMalformedResponse, err)
|
||||
wireResp, err := decodeOpenAIChatResponse(httpResp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if len(wireResp.Choices) == 0 {
|
||||
@@ -164,6 +187,46 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
||||
}, nil
|
||||
}
|
||||
|
||||
func decodeOpenAIChatResponse(body io.Reader) (openAIChatResponse, error) {
|
||||
limited := &io.LimitedReader{
|
||||
R: body,
|
||||
N: maxOpenAIChatResponseBytes + 1,
|
||||
}
|
||||
decoder := json.NewDecoder(limited)
|
||||
|
||||
var response openAIChatResponse
|
||||
if err := decoder.Decode(&response); err != nil {
|
||||
if limited.N == 0 {
|
||||
return openAIChatResponse{}, openAIChatResponseTooLargeError()
|
||||
}
|
||||
return openAIChatResponse{}, fmt.Errorf("%w: failed to decode response", ErrMalformedResponse)
|
||||
}
|
||||
if limited.N == 0 {
|
||||
return openAIChatResponse{}, openAIChatResponseTooLargeError()
|
||||
}
|
||||
|
||||
var trailing any
|
||||
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
|
||||
if limited.N == 0 {
|
||||
return openAIChatResponse{}, openAIChatResponseTooLargeError()
|
||||
}
|
||||
return openAIChatResponse{}, fmt.Errorf("%w: response contains trailing data", ErrMalformedResponse)
|
||||
}
|
||||
if limited.N == 0 {
|
||||
return openAIChatResponse{}, openAIChatResponseTooLargeError()
|
||||
}
|
||||
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func openAIChatResponseTooLargeError() error {
|
||||
return fmt.Errorf(
|
||||
"%w: response exceeds %d-byte limit",
|
||||
ErrMalformedResponse,
|
||||
maxOpenAIChatResponseBytes,
|
||||
)
|
||||
}
|
||||
|
||||
func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultModel string) (openAIChatRequest, error) {
|
||||
model := strings.TrimSpace(req.Target.Model)
|
||||
if model == "" {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -2,7 +2,6 @@ package builtin
|
||||
|
||||
import (
|
||||
"embed"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
)
|
||||
@@ -15,17 +14,3 @@ var assets embed.FS
|
||||
func NewRepository() profile.Repository {
|
||||
return profile.NewFSRepository(assets, assetRoot)
|
||||
}
|
||||
|
||||
func NewRepositoryWithPrimary(primary profile.Repository) profile.Repository {
|
||||
if primary == nil {
|
||||
return NewRepository()
|
||||
}
|
||||
return profile.NewOverlayRepository(primary, NewRepository())
|
||||
}
|
||||
|
||||
func NewRepositoryWithDirectory(dir string) profile.Repository {
|
||||
if strings.TrimSpace(dir) == "" {
|
||||
return NewRepository()
|
||||
}
|
||||
return NewRepositoryWithPrimary(profile.NewFilesystemRepository(dir))
|
||||
}
|
||||
|
||||
@@ -2,14 +2,11 @@ package builtin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io/fs"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/backend"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
@@ -91,53 +88,3 @@ func loadBuiltInProfileIDs(t *testing.T) map[string]string {
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func TestRepositoryWithPrimaryUsesPrimaryBeforeBuiltIns(t *testing.T) {
|
||||
repo := NewRepositoryWithPrimary(staticProfileRepo{
|
||||
profiles: map[string]string{"mistral-small-3": "custom-model"},
|
||||
})
|
||||
|
||||
p, err := repo.GetProfile(context.Background(), "mistral-small-3")
|
||||
if err != nil {
|
||||
t.Fatalf("expected profile to load, got %v", err)
|
||||
}
|
||||
if p.Model != "custom-model" {
|
||||
t.Fatalf("expected primary profile to override built-in, got %+v", p)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRepositoryWithPrimaryFallsBackToBuiltIns(t *testing.T) {
|
||||
repo := NewRepositoryWithPrimary(staticProfileRepo{})
|
||||
|
||||
p, err := repo.GetProfile(context.Background(), "mistral-small-3")
|
||||
if err != nil {
|
||||
t.Fatalf("expected built-in profile to load, got %v", err)
|
||||
}
|
||||
if p.ID != "mistral-small-3" {
|
||||
t.Fatalf("unexpected profile: %+v", p)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRepositoryWithPrimaryDoesNotFallBackAfterPrimaryError(t *testing.T) {
|
||||
repo := NewRepositoryWithPrimary(staticProfileRepo{err: profile.ErrInvalidProfile})
|
||||
|
||||
_, err := repo.GetProfile(context.Background(), "mistral-small-3")
|
||||
if !errors.Is(err, profile.ErrInvalidProfile) {
|
||||
t.Fatalf("expected primary error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
type staticProfileRepo struct {
|
||||
profiles map[string]string
|
||||
err error
|
||||
}
|
||||
|
||||
func (r staticProfileRepo) GetProfile(_ context.Context, id string) (*domain.ExecutionProfile, error) {
|
||||
if r.err != nil {
|
||||
return nil, r.err
|
||||
}
|
||||
if model, ok := r.profiles[id]; ok {
|
||||
return &domain.ExecutionProfile{ID: id, Endpoint: "http://primary/v1", Model: model}, nil
|
||||
}
|
||||
return nil, profile.ErrProfileNotFound
|
||||
}
|
||||
|
||||
@@ -5,13 +5,14 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/filecatalog"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
@@ -73,7 +74,8 @@ func (r *overlayRepository) GetProfile(ctx context.Context, id string) (*domain.
|
||||
}
|
||||
|
||||
func loadProfile(ctx context.Context, fsys fs.FS, root string, id string) (*domain.ExecutionProfile, error) {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
id = strings.TrimSpace(id)
|
||||
if id == "" {
|
||||
return nil, fmt.Errorf("%w: profile id is required", ErrInvalidProfile)
|
||||
}
|
||||
if fsys == nil {
|
||||
@@ -94,42 +96,50 @@ func loadProfile(ctx context.Context, fsys fs.FS, root string, id string) (*doma
|
||||
}
|
||||
|
||||
relPath := filecatalog.DisplayPath(root, fullPath)
|
||||
fileMatch := filecatalog.Stem(path.Base(fullPath)) == id
|
||||
data, err := fs.ReadFile(fsys, fullPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read profile file %s: %w", relPath, err)
|
||||
}
|
||||
metadata := readProfileFileMetadata(data)
|
||||
idMatch := fileMatch || metadata.id == id
|
||||
metadata, metadataErr := readProfileFileMetadata(data)
|
||||
idMatch := metadata.matchesID(id)
|
||||
if metadataErr != nil {
|
||||
if idMatch {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, metadataErr)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if metadata.hasRawAPIKey {
|
||||
if idMatch {
|
||||
return nil, fmt.Errorf("%w: %s", ErrRawAPIKeyNotAllowed, relPath)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
var prof domain.ExecutionProfile
|
||||
decoder := yaml.NewDecoder(bytes.NewReader(data))
|
||||
decoder.KnownFields(true)
|
||||
if err := decoder.Decode(&prof); err != nil {
|
||||
if idMatch {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, err)
|
||||
}
|
||||
if !idMatch {
|
||||
continue
|
||||
}
|
||||
|
||||
prof, err := decodeProfile(data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, err)
|
||||
}
|
||||
|
||||
prof.ID = strings.TrimSpace(prof.ID)
|
||||
if prof.ID != id {
|
||||
continue
|
||||
}
|
||||
prof.BackendID = strings.TrimSpace(prof.BackendID)
|
||||
if err := validateProfile(&prof); err != nil {
|
||||
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 errors.Is(err, ErrRawAPIKeyNotAllowed) {
|
||||
return nil, fmt.Errorf("%w: %s", err, relPath)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidProfile, relPath, err)
|
||||
}
|
||||
matches = append(matches, profileMatch{
|
||||
profile: &prof,
|
||||
profile: prof,
|
||||
path: relPath,
|
||||
})
|
||||
}
|
||||
@@ -155,15 +165,36 @@ type profileMatch struct {
|
||||
}
|
||||
|
||||
type profileFileMetadata struct {
|
||||
id string
|
||||
ids []string
|
||||
hasRawAPIKey bool
|
||||
}
|
||||
|
||||
func readProfileFileMetadata(data []byte) profileFileMetadata {
|
||||
func readProfileFileMetadata(data []byte) (profileFileMetadata, error) {
|
||||
decoder := yaml.NewDecoder(bytes.NewReader(data))
|
||||
var node yaml.Node
|
||||
if err := yaml.NewDecoder(bytes.NewReader(data)).Decode(&node); err != nil {
|
||||
return profileFileMetadata{}
|
||||
if err := decoder.Decode(&node); err != nil {
|
||||
return profileFileMetadata{}, err
|
||||
}
|
||||
metadata := profileMetadataFromNode(&node)
|
||||
documentCount := 1
|
||||
for {
|
||||
var trailing yaml.Node
|
||||
err := decoder.Decode(&trailing)
|
||||
if errors.Is(err, io.EOF) {
|
||||
if documentCount == 1 {
|
||||
return metadata, nil
|
||||
}
|
||||
return metadata, errors.New("profile file must contain exactly one YAML document")
|
||||
}
|
||||
if err != nil {
|
||||
return metadata, err
|
||||
}
|
||||
documentCount++
|
||||
metadata.merge(profileMetadataFromNode(&trailing))
|
||||
}
|
||||
}
|
||||
|
||||
func profileMetadataFromNode(node *yaml.Node) profileFileMetadata {
|
||||
if node.Kind != yaml.DocumentNode || len(node.Content) == 0 {
|
||||
return profileFileMetadata{}
|
||||
}
|
||||
@@ -178,7 +209,7 @@ func readProfileFileMetadata(data []byte) profileFileMetadata {
|
||||
value := mapping.Content[i+1]
|
||||
switch key.Value {
|
||||
case "id":
|
||||
metadata.id = strings.TrimSpace(value.Value)
|
||||
metadata.ids = append(metadata.ids, strings.TrimSpace(value.Value))
|
||||
case "api_key":
|
||||
metadata.hasRawAPIKey = true
|
||||
}
|
||||
@@ -186,29 +217,68 @@ func readProfileFileMetadata(data []byte) profileFileMetadata {
|
||||
return metadata
|
||||
}
|
||||
|
||||
func validateProfile(p *domain.ExecutionProfile) error {
|
||||
func (m profileFileMetadata) matchesID(id string) bool {
|
||||
for _, candidate := range m.ids {
|
||||
if candidate == id {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (m *profileFileMetadata) merge(other profileFileMetadata) {
|
||||
m.ids = append(m.ids, other.ids...)
|
||||
m.hasRawAPIKey = m.hasRawAPIKey || other.hasRawAPIKey
|
||||
}
|
||||
|
||||
func decodeProfile(data []byte) (*domain.ExecutionProfile, error) {
|
||||
var prof domain.ExecutionProfile
|
||||
decoder := yaml.NewDecoder(bytes.NewReader(data))
|
||||
decoder.KnownFields(true)
|
||||
if err := decoder.Decode(&prof); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := requireYAMLStreamEnd(decoder); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &prof, nil
|
||||
}
|
||||
|
||||
func requireYAMLStreamEnd(decoder *yaml.Decoder) error {
|
||||
var trailing yaml.Node
|
||||
err := decoder.Decode(&trailing)
|
||||
if errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
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")
|
||||
}
|
||||
if strings.TrimSpace(p.BackendID) == "" && strings.TrimSpace(p.Endpoint) == "" {
|
||||
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")
|
||||
}
|
||||
|
||||
if p.Temperature < 0 || p.Temperature > 2 {
|
||||
return errors.New("temperature must be between 0 and 2")
|
||||
}
|
||||
if p.MaxTokens < 0 {
|
||||
return errors.New("max_tokens must be greater than or equal to 0")
|
||||
}
|
||||
if p.TopP < 0 || p.TopP > 1 {
|
||||
return errors.New("top_p must be between 0 and 1")
|
||||
}
|
||||
if p.TimeoutSeconds < 0 {
|
||||
return errors.New("timeout_seconds must be greater than or equal to 0")
|
||||
}
|
||||
|
||||
return nil
|
||||
return domain.ValidateExecutionTargetSettings(domain.ExecutionTarget{
|
||||
Temperature: p.Temperature,
|
||||
MaxTokens: p.MaxTokens,
|
||||
TopP: p.TopP,
|
||||
TimeoutSeconds: p.TimeoutSeconds,
|
||||
})
|
||||
}
|
||||
|
||||
58
internal/profile/repository_benchmark_test.go
Normal file
58
internal/profile/repository_benchmark_test.go
Normal file
@@ -0,0 +1,58 @@
|
||||
package profile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
)
|
||||
|
||||
func BenchmarkProfileRepositoryLookup(b *testing.B) {
|
||||
for _, size := range []int{10, 1000} {
|
||||
b.Run(fmt.Sprintf("catalog-%d", size), func(b *testing.B) {
|
||||
files := fstest.MapFS{
|
||||
"target.yaml": profileMapFile(`
|
||||
id: target
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: target-model
|
||||
extra_params:
|
||||
selected: true
|
||||
`),
|
||||
}
|
||||
metadataNames := []string{"target.yaml"}
|
||||
for i := 1; i < size; i++ {
|
||||
name := fmt.Sprintf("profile-%04d.yaml", i)
|
||||
files[name] = profileMapFile(fmt.Sprintf(`
|
||||
id: profile-%04d
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated-model
|
||||
temperature: 0.5
|
||||
max_tokens: 500
|
||||
extra_params:
|
||||
provider:
|
||||
order:
|
||||
- first
|
||||
- second
|
||||
`, i))
|
||||
metadataNames = append(metadataNames, name)
|
||||
}
|
||||
|
||||
fsys := &recordingProfileFS{FS: files}
|
||||
repo := NewFSRepository(fsys, ".")
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
if _, err := repo.GetProfile(context.Background(), "target"); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
|
||||
for _, name := range metadataNames {
|
||||
if got := fsys.openCount(name); got != b.N {
|
||||
b.Fatalf("metadata %q opens = %d, want %d", name, got, b.N)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -2,11 +2,13 @@ package profile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
|
||||
@@ -61,7 +63,7 @@ func TestFilesystemRepository_GetProfile(t *testing.T) {
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "backend only", connection: "backend: ' openrouter '", wantBackend: "openrouter"},
|
||||
{name: "endpoint only", connection: "endpoint: http://localhost:8000/v1", wantEndpoint: "http://localhost:8000/v1"},
|
||||
{name: "endpoint only", connection: "endpoint: ' https://localhost:8000/nested/v1 '", wantEndpoint: "https://localhost:8000/nested/v1"},
|
||||
{name: "both", connection: "backend: openrouter\nendpoint: http://localhost:8000/v1", wantBackend: "openrouter", wantEndpoint: "http://localhost:8000/v1"},
|
||||
{name: "neither", wantErr: true},
|
||||
{name: "blank backend", connection: "backend: ' '", wantErr: true},
|
||||
@@ -126,63 +128,6 @@ temperature: 0.1
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("valid profile with JSON-compatible extra params", func(t *testing.T) {
|
||||
writeProfileTestFile(t, filepath.Join(tmpDir, "json-extra-params.yaml"), `
|
||||
id: json-extra-params
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: nested-model
|
||||
extra_params:
|
||||
string_value: enabled
|
||||
number_value: 42
|
||||
boolean_value: true
|
||||
object_value:
|
||||
nested: value
|
||||
count: 2
|
||||
array_value:
|
||||
- first
|
||||
- 3
|
||||
- false
|
||||
`)
|
||||
|
||||
p, err := repo.GetProfile(ctx, "json-extra-params")
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
|
||||
var got map[string]any
|
||||
encoded, err := json.Marshal(p.ExtraParams)
|
||||
if err != nil {
|
||||
t.Fatalf("expected extra_params to marshal as JSON, got %v", err)
|
||||
}
|
||||
if err := json.Unmarshal(encoded, &got); err != nil {
|
||||
t.Fatalf("expected extra_params JSON to decode, got %v", err)
|
||||
}
|
||||
|
||||
if got["string_value"] != "enabled" {
|
||||
t.Fatalf("unexpected string extra param: %#v", got["string_value"])
|
||||
}
|
||||
if got["number_value"] != float64(42) {
|
||||
t.Fatalf("unexpected number extra param: %#v", got["number_value"])
|
||||
}
|
||||
if got["boolean_value"] != true {
|
||||
t.Fatalf("unexpected boolean extra param: %#v", got["boolean_value"])
|
||||
}
|
||||
objectValue, ok := got["object_value"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("expected object extra param, got %#v", got["object_value"])
|
||||
}
|
||||
if objectValue["nested"] != "value" || objectValue["count"] != float64(2) {
|
||||
t.Fatalf("unexpected object extra param: %#v", objectValue)
|
||||
}
|
||||
arrayValue, ok := got["array_value"].([]any)
|
||||
if !ok {
|
||||
t.Fatalf("expected array extra param, got %#v", got["array_value"])
|
||||
}
|
||||
if len(arrayValue) != 3 || arrayValue[0] != "first" || arrayValue[1] != float64(3) || arrayValue[2] != false {
|
||||
t.Fatalf("unexpected array extra param: %#v", arrayValue)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("duplicate profile IDs fail as ambiguous", func(t *testing.T) {
|
||||
writeProfileTestFile(t, filepath.Join(tmpDir, "duplicate-profile-a.yaml"), `
|
||||
id: duplicate-profile
|
||||
@@ -245,10 +190,10 @@ api_key: secret
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid yaml", func(t *testing.T) {
|
||||
t.Run("unidentifiable invalid yaml is unrelated", func(t *testing.T) {
|
||||
_, err := repo.GetProfile(ctx, "invalid_yaml")
|
||||
if !errors.Is(err, ErrInvalidYAML) {
|
||||
t.Fatalf("expected ErrInvalidYAML, got %v", err)
|
||||
if !errors.Is(err, ErrProfileNotFound) {
|
||||
t.Fatalf("expected ErrProfileNotFound, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -274,14 +219,14 @@ api_key: secret
|
||||
})
|
||||
|
||||
t.Run("unknown field", func(t *testing.T) {
|
||||
_, err := repo.GetProfile(ctx, "unknown_field")
|
||||
_, err := repo.GetProfile(ctx, "unknown-field")
|
||||
if !errors.Is(err, ErrInvalidYAML) {
|
||||
t.Fatalf("expected ErrInvalidYAML for strict decode unknown field, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("raw api_key rejected", func(t *testing.T) {
|
||||
_, err := repo.GetProfile(ctx, "raw_api_key")
|
||||
_, err := repo.GetProfile(ctx, "raw-api-key")
|
||||
if !errors.Is(err, ErrRawAPIKeyNotAllowed) {
|
||||
t.Fatalf("expected ErrRawAPIKeyNotAllowed, got %v", err)
|
||||
}
|
||||
@@ -406,6 +351,513 @@ model: second
|
||||
})
|
||||
}
|
||||
|
||||
func TestProfileRepositoriesRejectInvalidEndpoints(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
endpoint string
|
||||
withBackend bool
|
||||
}{
|
||||
{name: "relative", endpoint: "/v1"},
|
||||
{name: "missing host", endpoint: "https:///v1"},
|
||||
{name: "unsupported scheme", endpoint: "ftp://provider.example/v1"},
|
||||
{name: "user information", endpoint: "https://user@provider.example/v1"},
|
||||
{name: "query", endpoint: "https://provider.example/v1?mode=chat"},
|
||||
{name: "fragment", endpoint: "https://provider.example/v1#chat"},
|
||||
{name: "backend with invalid override", endpoint: "/v1", withBackend: true},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
backend := ""
|
||||
if tc.withBackend {
|
||||
backend = "backend: openrouter\n"
|
||||
}
|
||||
repo := NewFSRepository(fstest.MapFS{
|
||||
"profiles/invalid.yaml": profileMapFile(fmt.Sprintf(
|
||||
"id: invalid-endpoint\nmodel: model\n%sendpoint: %q\n",
|
||||
backend,
|
||||
tc.endpoint,
|
||||
)),
|
||||
}, "profiles")
|
||||
|
||||
_, err := repo.GetProfile(context.Background(), "invalid-endpoint")
|
||||
if !errors.Is(err, ErrInvalidProfile) {
|
||||
t.Fatalf("expected ErrInvalidProfile, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileRepositoriesValidateExtraParams(t *testing.T) {
|
||||
const validProfile = `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: model
|
||||
extra_params:
|
||||
string_value: enabled
|
||||
object_value:
|
||||
nested: true
|
||||
array_value:
|
||||
- first
|
||||
- 3
|
||||
`
|
||||
tests := []struct {
|
||||
name string
|
||||
definition string
|
||||
wantErr bool
|
||||
diagnostics []string
|
||||
}{
|
||||
{name: "valid nested values", definition: validProfile},
|
||||
{
|
||||
name: "empty key",
|
||||
definition: `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: model
|
||||
extra_params:
|
||||
"": value
|
||||
`,
|
||||
wantErr: true,
|
||||
diagnostics: []string{"extra_params", "key must not be empty"},
|
||||
},
|
||||
{
|
||||
name: "non-finite value",
|
||||
definition: `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: model
|
||||
extra_params:
|
||||
invalid: .nan
|
||||
`,
|
||||
wantErr: true,
|
||||
diagnostics: []string{"extra_params.invalid", "must be finite"},
|
||||
},
|
||||
{
|
||||
name: "nested non-finite value",
|
||||
definition: `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: model
|
||||
extra_params:
|
||||
outer:
|
||||
invalid: .inf
|
||||
`,
|
||||
wantErr: true,
|
||||
diagnostics: []string{"extra_params.outer.invalid", "must be finite"},
|
||||
},
|
||||
{
|
||||
name: "unsupported decoded value",
|
||||
definition: `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: model
|
||||
extra_params:
|
||||
timestamp: 2026-08-11T12:34:56Z
|
||||
`,
|
||||
wantErr: true,
|
||||
diagnostics: []string{"extra_params.timestamp", "unsupported JSON value type"},
|
||||
},
|
||||
{
|
||||
name: "excessive nesting",
|
||||
definition: deeplyNestedExtraParamsProfile(101),
|
||||
wantErr: true,
|
||||
diagnostics: []string{"extra_params", "JSON container depth limit exceeded"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, source := range profileRepositorySources() {
|
||||
for _, tc := range tests {
|
||||
t.Run(source.name+"/"+tc.name, func(t *testing.T) {
|
||||
repo := source.newRepository(t, map[string]string{"selected.yaml": tc.definition})
|
||||
got, err := repo.GetProfile(context.Background(), "selected-profile")
|
||||
if tc.wantErr {
|
||||
if !errors.Is(err, ErrInvalidProfile) {
|
||||
t.Fatalf("expected ErrInvalidProfile, got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "selected.yaml") {
|
||||
t.Fatalf("expected source path in error, got %v", err)
|
||||
}
|
||||
for _, diagnostic := range tc.diagnostics {
|
||||
if !strings.Contains(err.Error(), diagnostic) {
|
||||
t.Fatalf("expected error to contain %q, got %v", diagnostic, err)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("load valid profile: %v", err)
|
||||
}
|
||||
if got.ExtraParams["string_value"] != "enabled" {
|
||||
t.Fatalf("unexpected copied extra params: %#v", got.ExtraParams)
|
||||
}
|
||||
objectValue, objectOK := got.ExtraParams["object_value"].(map[string]any)
|
||||
arrayValue, arrayOK := got.ExtraParams["array_value"].([]any)
|
||||
if !objectOK || objectValue["nested"] != true ||
|
||||
!arrayOK || len(arrayValue) != 2 || arrayValue[0] != "first" || arrayValue[1] != 3 {
|
||||
t.Fatalf("unexpected copied nested extra params: %#v", got.ExtraParams)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileRepositoriesSelectCanonicalYAMLID(t *testing.T) {
|
||||
const validProfile = `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: selected-model
|
||||
`
|
||||
tests := []struct {
|
||||
name string
|
||||
files map[string]string
|
||||
wantErr error
|
||||
diagnostics []string
|
||||
}{
|
||||
{
|
||||
name: "same stem unknown field with different id is unrelated",
|
||||
files: map[string]string{
|
||||
"selected-profile.yaml": `
|
||||
id: unrelated-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated
|
||||
unknown: true
|
||||
`,
|
||||
"valid.yaml": validProfile,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "same stem unidentifiable yaml is unrelated",
|
||||
files: map[string]string{
|
||||
"selected-profile.yaml": "id: [",
|
||||
"valid.yaml": validProfile,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "same stem raw key with different id is unrelated",
|
||||
files: map[string]string{
|
||||
"selected-profile.yaml": `
|
||||
id: unrelated-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated
|
||||
api_key: secret
|
||||
`,
|
||||
"valid.yaml": validProfile,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "leading and trailing whitespace is normalized",
|
||||
files: map[string]string{
|
||||
"padded.yaml": `
|
||||
id: " selected-profile "
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: selected-model
|
||||
`,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "blank id is unrelated",
|
||||
files: map[string]string{
|
||||
"selected-profile.yaml": `
|
||||
id: " "
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated
|
||||
`,
|
||||
},
|
||||
wantErr: ErrProfileNotFound,
|
||||
},
|
||||
{
|
||||
name: "normalized duplicates are ambiguous",
|
||||
files: map[string]string{
|
||||
"first.yaml": validProfile,
|
||||
"nested/second.yaml": `
|
||||
id: " selected-profile "
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: duplicate
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidProfile,
|
||||
diagnostics: []string{"duplicate execution profile id", "first.yaml", "nested/second.yaml"},
|
||||
},
|
||||
{
|
||||
name: "selected unknown field is authoritative",
|
||||
files: map[string]string{
|
||||
"malformed.yaml": `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: selected-model
|
||||
unknown: true
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidYAML,
|
||||
diagnostics: []string{"malformed.yaml"},
|
||||
},
|
||||
{
|
||||
name: "selected raw key is authoritative",
|
||||
files: map[string]string{
|
||||
"insecure.yaml": `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: selected-model
|
||||
api_key: secret
|
||||
`,
|
||||
},
|
||||
wantErr: ErrRawAPIKeyNotAllowed,
|
||||
diagnostics: []string{"insecure.yaml"},
|
||||
},
|
||||
{
|
||||
name: "selected identity in an additional document is authoritative",
|
||||
files: map[string]string{
|
||||
"additional-document.yaml": `
|
||||
---
|
||||
---
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: selected-model
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidYAML,
|
||||
diagnostics: []string{"additional-document.yaml", "exactly one YAML document"},
|
||||
},
|
||||
}
|
||||
|
||||
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.wantErr != nil {
|
||||
if !errors.Is(err, tc.wantErr) {
|
||||
t.Fatalf("expected %v, got %v", tc.wantErr, err)
|
||||
}
|
||||
for _, diagnostic := range tc.diagnostics {
|
||||
if !strings.Contains(err.Error(), diagnostic) {
|
||||
t.Fatalf("expected error to contain %q, got %v", diagnostic, err)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("load selected profile: %v", err)
|
||||
}
|
||||
if got.ID != "selected-profile" || got.Model != "selected-model" {
|
||||
t.Fatalf("unexpected selected profile: %+v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileRepositoriesRequireOneYAMLDocument(t *testing.T) {
|
||||
const profile = `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: selected-model
|
||||
`
|
||||
tests := []struct {
|
||||
name string
|
||||
suffix string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "comments and trailing whitespace", suffix: "\n# trailing comment\n\n"},
|
||||
{name: "second populated document", suffix: "\n---\nid: another\n", wantErr: true},
|
||||
{name: "second empty document", suffix: "\n---\n", wantErr: true},
|
||||
{name: "malformed trailing yaml", suffix: "\n---\n[", wantErr: true},
|
||||
{name: "raw key in trailing document", suffix: "\n---\napi_key: secret\n", wantErr: true},
|
||||
}
|
||||
|
||||
for _, source := range profileRepositorySources() {
|
||||
for _, tc := range tests {
|
||||
t.Run(source.name+"/"+tc.name, func(t *testing.T) {
|
||||
repo := source.newRepository(t, map[string]string{"definition.yaml": profile + tc.suffix})
|
||||
got, err := repo.GetProfile(context.Background(), "selected-profile")
|
||||
if tc.wantErr {
|
||||
if !errors.Is(err, ErrInvalidYAML) {
|
||||
t.Fatalf("expected ErrInvalidYAML, got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "definition.yaml") {
|
||||
t.Fatalf("expected source path in error, got %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("load one-document profile: %v", err)
|
||||
}
|
||||
if got.ID != "selected-profile" {
|
||||
t.Fatalf("unexpected profile: %+v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileRepositoriesPreserveOverlayFallbackRules(t *testing.T) {
|
||||
fallback := staticProfileRepo{profiles: map[string]*domain.ExecutionProfile{
|
||||
"selected-profile": {ID: "selected-profile", Endpoint: "http://fallback", Model: "fallback-model"},
|
||||
}}
|
||||
tests := []struct {
|
||||
name string
|
||||
files map[string]string
|
||||
wantModel string
|
||||
wantErr error
|
||||
}{
|
||||
{
|
||||
name: "same stem malformed different id falls back",
|
||||
files: map[string]string{
|
||||
"selected-profile.yaml": `
|
||||
id: unrelated-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated
|
||||
unknown: true
|
||||
`,
|
||||
},
|
||||
wantModel: "fallback-model",
|
||||
},
|
||||
{
|
||||
name: "blank id falls back",
|
||||
files: map[string]string{
|
||||
"selected-profile.yaml": `
|
||||
id: " "
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated
|
||||
`,
|
||||
},
|
||||
wantModel: "fallback-model",
|
||||
},
|
||||
{
|
||||
name: "selected malformed profile stops fallback",
|
||||
files: map[string]string{
|
||||
"other-name.yaml": `
|
||||
id: selected-profile
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: selected
|
||||
unknown: true
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidYAML,
|
||||
},
|
||||
}
|
||||
|
||||
for _, source := range profileRepositorySources() {
|
||||
for _, tc := range tests {
|
||||
t.Run(source.name+"/"+tc.name, func(t *testing.T) {
|
||||
primary := source.newRepository(t, tc.files)
|
||||
got, err := NewOverlayRepository(primary, fallback).GetProfile(context.Background(), "selected-profile")
|
||||
if tc.wantErr != nil {
|
||||
if !errors.Is(err, tc.wantErr) {
|
||||
t.Fatalf("expected %v, got %v", tc.wantErr, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("load fallback profile: %v", err)
|
||||
}
|
||||
if got.Model != tc.wantModel {
|
||||
t.Fatalf("model = %q, want %q", got.Model, tc.wantModel)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileRepositoryReadsSourcesFreshOnEveryLookup(t *testing.T) {
|
||||
newSource := func() (*recordingProfileFS, Repository) {
|
||||
fsys := &recordingProfileFS{FS: fstest.MapFS{
|
||||
"target.yaml": profileMapFile(`
|
||||
id: target
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: target-model
|
||||
`),
|
||||
"unrelated.yaml": profileMapFile(`
|
||||
id: unrelated
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated-model
|
||||
`),
|
||||
}}
|
||||
return fsys, NewFSRepository(fsys, ".")
|
||||
}
|
||||
|
||||
t.Run("selected source", func(t *testing.T) {
|
||||
fsys, repo := newSource()
|
||||
for lookup := 1; lookup <= 2; lookup++ {
|
||||
got, err := repo.GetProfile(context.Background(), "target")
|
||||
if err != nil {
|
||||
t.Fatalf("lookup %d: %v", lookup, err)
|
||||
}
|
||||
if got.Model != "target-model" {
|
||||
t.Fatalf("lookup %d model = %q", lookup, got.Model)
|
||||
}
|
||||
for _, name := range []string{"target.yaml", "unrelated.yaml"} {
|
||||
if count := fsys.openCount(name); count != lookup {
|
||||
t.Fatalf("%s opens after lookup %d = %d, want %d", name, lookup, count, lookup)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("overlay fallthrough", func(t *testing.T) {
|
||||
primaryFS := &recordingProfileFS{FS: fstest.MapFS{
|
||||
"unrelated.yaml": profileMapFile(`
|
||||
id: unrelated
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: unrelated-model
|
||||
`),
|
||||
}}
|
||||
fallbackFS, fallback := newSource()
|
||||
repo := NewOverlayRepository(NewFSRepository(primaryFS, "."), fallback)
|
||||
|
||||
for lookup := 1; lookup <= 2; lookup++ {
|
||||
got, err := repo.GetProfile(context.Background(), "target")
|
||||
if err != nil {
|
||||
t.Fatalf("lookup %d: %v", lookup, err)
|
||||
}
|
||||
if got.Model != "target-model" {
|
||||
t.Fatalf("lookup %d model = %q", lookup, got.Model)
|
||||
}
|
||||
if count := primaryFS.openCount("unrelated.yaml"); count != lookup {
|
||||
t.Fatalf("primary opens after lookup %d = %d, want %d", lookup, count, lookup)
|
||||
}
|
||||
if count := fallbackFS.openCount("target.yaml"); count != lookup {
|
||||
t.Fatalf("fallback opens after lookup %d = %d, want %d", lookup, count, lookup)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestProfileRepositoriesRejectInvalidExecutionSettings(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("operating-system filesystem", func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeProfileTestFile(t, filepath.Join(dir, "invalid.yaml"), `
|
||||
id: invalid
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: model
|
||||
temperature: .nan
|
||||
`)
|
||||
|
||||
_, err := NewFilesystemRepository(dir).GetProfile(ctx, "invalid")
|
||||
if !errors.Is(err, ErrInvalidProfile) {
|
||||
t.Fatalf("expected ErrInvalidProfile, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("fs.FS", func(t *testing.T) {
|
||||
repo := NewFSRepository(fstest.MapFS{
|
||||
"profiles/invalid.yaml": profileMapFile(`
|
||||
id: invalid
|
||||
endpoint: http://localhost:8000/v1
|
||||
model: model
|
||||
top_p: .inf
|
||||
`),
|
||||
}, "profiles")
|
||||
|
||||
_, err := repo.GetProfile(ctx, "invalid")
|
||||
if !errors.Is(err, ErrInvalidProfile) {
|
||||
t.Fatalf("expected ErrInvalidProfile, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestOverlayRepository(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
primaryProfile := &domain.ExecutionProfile{ID: "shared", Endpoint: "http://primary", Model: "primary"}
|
||||
@@ -495,6 +947,77 @@ func TestOverlayRepository(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
type profileRepositorySource struct {
|
||||
name string
|
||||
newRepository func(t *testing.T, files map[string]string) Repository
|
||||
}
|
||||
|
||||
type recordingProfileFS struct {
|
||||
fs.FS
|
||||
mu sync.Mutex
|
||||
opened []string
|
||||
}
|
||||
|
||||
func (f *recordingProfileFS) Open(name string) (fs.File, error) {
|
||||
f.mu.Lock()
|
||||
f.opened = append(f.opened, name)
|
||||
f.mu.Unlock()
|
||||
return f.FS.Open(name)
|
||||
}
|
||||
|
||||
func (f *recordingProfileFS) openCount(name string) int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
count := 0
|
||||
for _, opened := range f.opened {
|
||||
if opened == name {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func profileRepositorySources() []profileRepositorySource {
|
||||
return []profileRepositorySource{
|
||||
{
|
||||
name: "operating system",
|
||||
newRepository: func(t *testing.T, files map[string]string) Repository {
|
||||
t.Helper()
|
||||
root := t.TempDir()
|
||||
for name, content := range files {
|
||||
filePath := filepath.Join(root, filepath.FromSlash(name))
|
||||
if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil {
|
||||
t.Fatalf("create profile directory: %v", err)
|
||||
}
|
||||
writeProfileTestFile(t, filePath, content)
|
||||
}
|
||||
return NewFilesystemRepository(root)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "filesystem",
|
||||
newRepository: func(t *testing.T, files map[string]string) Repository {
|
||||
t.Helper()
|
||||
fsys := make(fstest.MapFS, len(files))
|
||||
for name, content := range files {
|
||||
fsys[name] = profileMapFile(content)
|
||||
}
|
||||
return NewFSRepository(fsys, ".")
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func deeplyNestedExtraParamsProfile(depth int) string {
|
||||
var definition strings.Builder
|
||||
definition.WriteString("id: selected-profile\nendpoint: http://localhost:8000/v1\nmodel: model\nextra_params:\n")
|
||||
for level := 0; level < depth; level++ {
|
||||
fmt.Fprintf(&definition, "%slevel_%d:\n", strings.Repeat(" ", level+1), level)
|
||||
}
|
||||
fmt.Fprintf(&definition, "%svalue: true\n", strings.Repeat(" ", depth+1))
|
||||
return definition.String()
|
||||
}
|
||||
|
||||
func profileMapFile(content string) *fstest.MapFile {
|
||||
return &fstest.MapFile{Data: []byte(strings.TrimLeft(content, "\n"))}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"text/template"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
@@ -18,6 +19,8 @@ var (
|
||||
ErrInvalidMessageRole = errors.New("invalid or empty message role")
|
||||
)
|
||||
|
||||
const artifactTextChunkSize = 64 * 1024
|
||||
|
||||
type goRenderer struct{}
|
||||
|
||||
func NewGoRenderer() Renderer {
|
||||
@@ -25,11 +28,13 @@ func NewGoRenderer() Renderer {
|
||||
}
|
||||
|
||||
func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefinition, inputs map[string]*domain.Artifact, vars map[string]string) (*domain.RenderedPrompt, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if definition == nil {
|
||||
return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure)
|
||||
}
|
||||
|
||||
// 1. Verify required inputs
|
||||
for _, in := range definition.Inputs {
|
||||
if !in.Required {
|
||||
continue
|
||||
@@ -40,44 +45,54 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Setup template functions
|
||||
resolver := newArtifactTextResolver(ctx, inputs)
|
||||
funcs := template.FuncMap{
|
||||
"input": func(name string) (string, error) {
|
||||
art, ok := inputs[name]
|
||||
if !ok || art == nil {
|
||||
return "", fmt.Errorf("%w: %s", ErrUnknownInput, name)
|
||||
}
|
||||
return string(art.Body), nil
|
||||
},
|
||||
"input": resolver.resolve,
|
||||
}
|
||||
|
||||
sessionID, err := renderSessionID(definition.SessionID, funcs, vars)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sessionID, err := renderSessionID(ctx, definition.SessionID, funcs, vars)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var renderedMessages []domain.RenderedMessage
|
||||
renderedMessages := make([]domain.RenderedMessage, 0, len(definition.Templates))
|
||||
|
||||
for i, tmplMsg := range definition.Templates {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if tmplMsg.Role == "" {
|
||||
return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i)
|
||||
}
|
||||
|
||||
// Parse and execute template
|
||||
tmpl, err := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, err)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tmpl, parseErr := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if parseErr != nil {
|
||||
return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, parseErr)
|
||||
}
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := tmpl.Execute(&buf, vars); err != nil {
|
||||
return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, err)
|
||||
executeErr := tmpl.Execute(&buf, vars)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if executeErr != nil {
|
||||
return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, executeErr)
|
||||
}
|
||||
|
||||
renderedMessages = append(renderedMessages, domain.RenderedMessage{
|
||||
@@ -85,6 +100,13 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
|
||||
Content: buf.String(),
|
||||
CacheControl: cloneCacheControl(tmplMsg.CacheControl),
|
||||
})
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &domain.RenderedPrompt{
|
||||
@@ -93,15 +115,75 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
|
||||
}, nil
|
||||
}
|
||||
|
||||
func renderSessionID(raw string, funcs template.FuncMap, vars map[string]string) (string, error) {
|
||||
tmpl, err := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, err)
|
||||
type artifactTextResolver struct {
|
||||
ctx context.Context
|
||||
inputs map[string]*domain.Artifact
|
||||
textByName map[string]string
|
||||
}
|
||||
|
||||
func newArtifactTextResolver(ctx context.Context, inputs map[string]*domain.Artifact) *artifactTextResolver {
|
||||
return &artifactTextResolver{
|
||||
ctx: ctx,
|
||||
inputs: inputs,
|
||||
textByName: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
func (r *artifactTextResolver) resolve(name string) (string, error) {
|
||||
if err := r.ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
artifact, ok := r.inputs[name]
|
||||
if !ok || artifact == nil {
|
||||
return "", fmt.Errorf("%w: %s", ErrUnknownInput, name)
|
||||
}
|
||||
if text, ok := r.textByName[name]; ok {
|
||||
return text, nil
|
||||
}
|
||||
|
||||
var builder strings.Builder
|
||||
builder.Grow(len(artifact.Body))
|
||||
for start := 0; start < len(artifact.Body); start += artifactTextChunkSize {
|
||||
if err := r.ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
end := min(start+artifactTextChunkSize, len(artifact.Body))
|
||||
_, _ = builder.Write(artifact.Body[start:end])
|
||||
if err := r.ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
if err := r.ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
text := builder.String()
|
||||
r.textByName[name] = text
|
||||
return text, nil
|
||||
}
|
||||
|
||||
func renderSessionID(ctx context.Context, raw string, funcs template.FuncMap, vars map[string]string) (string, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
tmpl, parseErr := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if parseErr != nil {
|
||||
return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, parseErr)
|
||||
}
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := tmpl.Execute(&buf, vars); err != nil {
|
||||
return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, err)
|
||||
executeErr := tmpl.Execute(&buf, vars)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if executeErr != nil {
|
||||
return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, executeErr)
|
||||
}
|
||||
|
||||
sessionID, err := domain.NormalizeSessionID(buf.String())
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package prompt
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
@@ -228,6 +229,24 @@ func TestGoRenderer_Render(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("malformed rendered session id fails rendering", func(t *testing.T) {
|
||||
def := &domain.PromptDefinition{
|
||||
SessionID: "{{ .session_id }}",
|
||||
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
|
||||
Templates: []domain.PromptMessageTemplate{
|
||||
{Role: "system", Content: "Speak in a {{.tone}} tone."},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := renderer.Render(ctx, def, inputs, map[string]string{
|
||||
"tone": "concise",
|
||||
"session_id": "session" + string([]byte{0xff}),
|
||||
})
|
||||
if !errors.Is(err, ErrRenderFailure) {
|
||||
t.Fatalf("expected ErrRenderFailure, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("inserting required input artifact", func(t *testing.T) {
|
||||
def := &domain.PromptDefinition{
|
||||
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}},
|
||||
@@ -343,3 +362,186 @@ func TestGoRenderer_Render(t *testing.T) {
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestGoRendererCancellation(t *testing.T) {
|
||||
t.Run("before session parsing", func(t *testing.T) {
|
||||
definition := &domain.PromptDefinition{
|
||||
SessionID: "{{ malformed",
|
||||
Templates: []domain.PromptMessageTemplate{
|
||||
{Role: "user", Content: "not rendered"},
|
||||
},
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
result, err := NewGoRenderer().Render(ctx, definition, nil, nil)
|
||||
if result != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("result=%#v err=%v, want nil/context.Canceled", result, err)
|
||||
}
|
||||
if errors.Is(err, ErrInvalidTemplate) {
|
||||
t.Fatalf("pre-canceled render parsed the malformed session: %v", err)
|
||||
}
|
||||
|
||||
result, err = NewGoRenderer().Render(context.Background(), definition, nil, nil)
|
||||
if result != nil || !errors.Is(err, ErrInvalidTemplate) {
|
||||
t.Fatalf("active render result=%#v err=%v, want nil/ErrInvalidTemplate", result, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("during artifact text conversion", func(t *testing.T) {
|
||||
ctx := newCancelOnCheckContext(3)
|
||||
body := bytes.Repeat([]byte("x"), artifactTextChunkSize*2)
|
||||
original := append([]byte(nil), body...)
|
||||
resolver := newArtifactTextResolver(ctx, map[string]*domain.Artifact{
|
||||
"document": {Body: body},
|
||||
})
|
||||
|
||||
text, err := resolver.resolve("document")
|
||||
if text != "" || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("text length=%d err=%v, want empty/context.Canceled", len(text), err)
|
||||
}
|
||||
if _, published := resolver.textByName["document"]; published {
|
||||
t.Fatal("canceled conversion published partial artifact text")
|
||||
}
|
||||
if !bytes.Equal(body, original) {
|
||||
t.Fatal("resolver mutated the artifact body")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("after final message execution", func(t *testing.T) {
|
||||
definition := &domain.PromptDefinition{
|
||||
Templates: []domain.PromptMessageTemplate{
|
||||
{Role: "user", Content: "fully rendered"},
|
||||
},
|
||||
}
|
||||
counter := &checkCountingContext{Context: context.Background()}
|
||||
if _, err := NewGoRenderer().Render(counter, definition, nil, nil); err != nil {
|
||||
t.Fatalf("count render checkpoints: %v", err)
|
||||
}
|
||||
|
||||
// The final three checks occur after template execution, after the
|
||||
// message is assembled, and immediately before publication.
|
||||
ctx := newCancelOnCheckContext(counter.checks - 2)
|
||||
result, err := NewGoRenderer().Render(ctx, definition, nil, nil)
|
||||
if result != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("result=%#v err=%v, want nil/context.Canceled", result, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestGoRendererArtifactTextLifecycle(t *testing.T) {
|
||||
body := []byte{'a', 0xff, 'b', 0xfe}
|
||||
original := append([]byte(nil), body...)
|
||||
artifact := &domain.Artifact{Body: body}
|
||||
inputs := map[string]*domain.Artifact{"document": artifact}
|
||||
definition := &domain.PromptDefinition{
|
||||
Templates: []domain.PromptMessageTemplate{
|
||||
{Role: "user", Content: "{{input \"document\"}}|{{input \"document\"}}"},
|
||||
},
|
||||
}
|
||||
|
||||
first, err := NewGoRenderer().Render(context.Background(), definition, inputs, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("first render: %v", err)
|
||||
}
|
||||
wantFirst := append(append(append([]byte(nil), body...), '|'), body...)
|
||||
if !bytes.Equal([]byte(first.Messages[0].Content), wantFirst) {
|
||||
t.Fatalf("rendered bytes=%v, want %v", []byte(first.Messages[0].Content), wantFirst)
|
||||
}
|
||||
if !bytes.Equal(body, original) {
|
||||
t.Fatalf("renderer mutated artifact body: got %v want %v", body, original)
|
||||
}
|
||||
|
||||
body[0] = 'z'
|
||||
if bytes.Equal([]byte(first.Messages[0].Content), append(append(append([]byte(nil), body...), '|'), body...)) {
|
||||
t.Fatal("completed render aliases the artifact body")
|
||||
}
|
||||
second, err := NewGoRenderer().Render(context.Background(), definition, inputs, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("second render: %v", err)
|
||||
}
|
||||
wantSecond := append(append(append([]byte(nil), body...), '|'), body...)
|
||||
if !bytes.Equal([]byte(second.Messages[0].Content), wantSecond) {
|
||||
t.Fatalf("second render reused text from another call: got %v want %v", []byte(second.Messages[0].Content), wantSecond)
|
||||
}
|
||||
|
||||
nilInputs := map[string]*domain.Artifact{"document": nil}
|
||||
result, err := NewGoRenderer().Render(context.Background(), definition, nilInputs, nil)
|
||||
if result != nil || !errors.Is(err, ErrUnknownInput) || !errors.Is(err, ErrRenderFailure) {
|
||||
t.Fatalf("nil input result=%#v err=%v, want ErrUnknownInput and ErrRenderFailure", result, err)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkGoRendererArtifactReferences(b *testing.B) {
|
||||
body := bytes.Repeat([]byte("document content "), (artifactTextChunkSize*4)/len("document content "))
|
||||
inputs := map[string]*domain.Artifact{
|
||||
"document": {Body: body},
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
definition *domain.PromptDefinition
|
||||
}{
|
||||
{
|
||||
name: "one reference",
|
||||
definition: &domain.PromptDefinition{
|
||||
Templates: []domain.PromptMessageTemplate{
|
||||
{Role: "user", Content: "{{input \"document\"}}"},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "repeated across session and messages",
|
||||
definition: &domain.PromptDefinition{
|
||||
SessionID: "document-{{len (input \"document\")}}",
|
||||
Templates: []domain.PromptMessageTemplate{
|
||||
{Role: "system", Content: "{{input \"document\"}}"},
|
||||
{Role: "user", Content: "{{input \"document\"}} {{input \"document\"}}"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
b.Run(tc.name, func(b *testing.B) {
|
||||
renderer := NewGoRenderer()
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(body)))
|
||||
for range b.N {
|
||||
if _, err := renderer.Render(context.Background(), tc.definition, inputs, nil); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type checkCountingContext struct {
|
||||
context.Context
|
||||
checks int
|
||||
}
|
||||
|
||||
func (c *checkCountingContext) Err() error {
|
||||
c.checks++
|
||||
return c.Context.Err()
|
||||
}
|
||||
|
||||
type cancelOnCheckContext struct {
|
||||
context.Context
|
||||
cancel context.CancelFunc
|
||||
remaining int
|
||||
}
|
||||
|
||||
func newCancelOnCheckContext(checks int) *cancelOnCheckContext {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
return &cancelOnCheckContext{Context: ctx, cancel: cancel, remaining: checks}
|
||||
}
|
||||
|
||||
func (c *cancelOnCheckContext) Err() error {
|
||||
if c.Context.Err() == nil {
|
||||
c.remaining--
|
||||
if c.remaining == 0 {
|
||||
c.cancel()
|
||||
}
|
||||
}
|
||||
return c.Context.Err()
|
||||
}
|
||||
|
||||
98
internal/promptdef/content_source.go
Normal file
98
internal/promptdef/content_source.go
Normal file
@@ -0,0 +1,98 @@
|
||||
package promptdef
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/filecatalog"
|
||||
)
|
||||
|
||||
type contentSourceRoot interface {
|
||||
readContentFile(sourcePath string, contentFile string) (string, string, error)
|
||||
}
|
||||
|
||||
type osContentSourceRoot struct {
|
||||
root string
|
||||
sourcePathsRelative bool
|
||||
}
|
||||
|
||||
func (r osContentSourceRoot) readContentFile(sourcePath string, contentFile string) (string, string, error) {
|
||||
if strings.TrimSpace(contentFile) == "" {
|
||||
return "", "", fmt.Errorf("path is required")
|
||||
}
|
||||
if filepath.IsAbs(contentFile) {
|
||||
return "", "", fmt.Errorf("path %q must be relative", contentFile)
|
||||
}
|
||||
|
||||
root, err := filepath.Abs(r.root)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("resolve source root %q: %w", r.root, err)
|
||||
}
|
||||
canonicalRoot, err := filepath.EvalSymlinks(root)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("resolve source root %q: %w", r.root, err)
|
||||
}
|
||||
|
||||
promptPath := sourcePath
|
||||
if r.sourcePathsRelative && !filepath.IsAbs(promptPath) {
|
||||
promptPath = filepath.Join(root, filepath.FromSlash(promptPath))
|
||||
} else {
|
||||
promptPath, err = filepath.Abs(promptPath)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("resolve prompt source %q: %w", sourcePath, err)
|
||||
}
|
||||
}
|
||||
resolvedPath := filepath.Clean(filepath.Join(filepath.Dir(promptPath), contentFile))
|
||||
if !containsOSPath(root, resolvedPath) {
|
||||
return "", "", fmt.Errorf("path %q escapes source root %q", contentFile, r.root)
|
||||
}
|
||||
|
||||
canonicalPath, err := filepath.EvalSymlinks(resolvedPath)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
if !containsOSPath(canonicalRoot, canonicalPath) {
|
||||
return "", "", fmt.Errorf("path %q escapes source root %q", contentFile, r.root)
|
||||
}
|
||||
|
||||
body, err := os.ReadFile(canonicalPath)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return string(body), resolvedPath, nil
|
||||
}
|
||||
|
||||
type fsContentSourceRoot struct {
|
||||
fsys fs.FS
|
||||
root string
|
||||
}
|
||||
|
||||
func (r fsContentSourceRoot) readContentFile(sourcePath string, contentFile string) (string, string, error) {
|
||||
root := filecatalog.CleanFSRoot(r.root)
|
||||
cleanSourcePath := path.Clean(sourcePath)
|
||||
if cleanSourcePath == root {
|
||||
root = path.Dir(root)
|
||||
}
|
||||
|
||||
resolvedPath, _, err := filecatalog.ResolveFSPath(root, path.Dir(cleanSourcePath), contentFile)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
body, err := fs.ReadFile(r.fsys, resolvedPath)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return string(body), resolvedPath, nil
|
||||
}
|
||||
|
||||
func containsOSPath(root string, name string) bool {
|
||||
relative, err := filepath.Rel(root, name)
|
||||
if err != nil || filepath.IsAbs(relative) {
|
||||
return false
|
||||
}
|
||||
return relative != ".." && !strings.HasPrefix(relative, ".."+string(filepath.Separator))
|
||||
}
|
||||
@@ -5,14 +5,11 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/filecatalog"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
@@ -22,13 +19,8 @@ var (
|
||||
ErrInvalidPromptDefinition = errors.New("invalid prompt definition configuration")
|
||||
)
|
||||
|
||||
type filesystemRepository struct {
|
||||
dir string
|
||||
}
|
||||
|
||||
type fsRepository struct {
|
||||
fsys fs.FS
|
||||
root string
|
||||
type sourceRepository struct {
|
||||
source promptDefinitionSource
|
||||
}
|
||||
|
||||
type promptDefinitionFile struct {
|
||||
@@ -69,19 +61,47 @@ type promptOutputContractFile struct {
|
||||
}
|
||||
|
||||
func NewFilesystemRepository(dir string) Repository {
|
||||
return &filesystemRepository{dir: dir}
|
||||
return &sourceRepository{
|
||||
source: osPromptSource{
|
||||
root: dir,
|
||||
contentRoot: osContentSourceRoot{root: dir},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func NewFSRepository(fsys fs.FS, root string) Repository {
|
||||
return &fsRepository{fsys: fsys, root: root}
|
||||
return &sourceRepository{
|
||||
source: fsPromptSource{
|
||||
fsys: fsys,
|
||||
root: root,
|
||||
contentRoot: fsContentSourceRoot{fsys: fsys, root: root},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (r *filesystemRepository) GetPromptDefinition(ctx context.Context, id string, version string) (*domain.PromptDefinition, error) {
|
||||
// NewFileRepository constructs a repository for one operating-system prompt file.
|
||||
func NewFileRepository(fsys fs.FS, file string, sourceDir string) Repository {
|
||||
return &sourceRepository{
|
||||
source: fsPromptSource{
|
||||
fsys: fsys,
|
||||
root: file,
|
||||
contentRoot: osContentSourceRoot{
|
||||
root: sourceDir,
|
||||
sourcePathsRelative: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (r *sourceRepository) GetPromptDefinition(ctx context.Context, id string, version string) (*domain.PromptDefinition, error) {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidPromptDefinition)
|
||||
}
|
||||
if r == nil || r.source == nil {
|
||||
return nil, errors.New("failed to read prompt definition directory: source is nil")
|
||||
}
|
||||
|
||||
files, err := filecatalog.FindYAMLFiles(ctx, r.dir)
|
||||
files, err := r.source.findYAMLFiles(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read prompt definition directory: %w", err)
|
||||
}
|
||||
@@ -94,34 +114,25 @@ func (r *filesystemRepository) GetPromptDefinition(ctx context.Context, id strin
|
||||
default:
|
||||
}
|
||||
|
||||
relPath := filecatalog.RelativePath(r.dir, fullPath)
|
||||
fileMatch := filecatalog.Stem(filepath.Base(fullPath)) == id
|
||||
|
||||
raw, err := loadPromptDefinitionFile(fullPath)
|
||||
relPath := r.source.displayPath(fullPath)
|
||||
data, err := r.source.readDefinition(fullPath)
|
||||
if err != nil {
|
||||
if fileMatch || promptDefinitionFileHasID(fullPath, id) {
|
||||
return nil, fmt.Errorf("failed to read prompt definition file %s: %w", relPath, err)
|
||||
}
|
||||
raw, err := decodePromptDefinition(data)
|
||||
if err != nil {
|
||||
if promptDefinitionDataMatches(data, id, version) {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
def, err := normalizePromptDefinition(raw, fullPath)
|
||||
if err != nil {
|
||||
if fileMatch || strings.TrimSpace(raw.ID) == id {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidPromptDefinition, relPath, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if def.ID != id {
|
||||
continue
|
||||
}
|
||||
if version != "" && def.Version != version {
|
||||
if !promptDefinitionMatches(raw, id, version) {
|
||||
continue
|
||||
}
|
||||
matches = append(matches, promptDefinitionMatch{
|
||||
def: def,
|
||||
path: relPath,
|
||||
raw: raw,
|
||||
sourcePath: fullPath,
|
||||
path: relPath,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -137,130 +148,21 @@ func (r *filesystemRepository) GetPromptDefinition(ctx context.Context, id strin
|
||||
}
|
||||
|
||||
if len(matches) == 1 {
|
||||
return matches[0].def, nil
|
||||
match := matches[0]
|
||||
def, err := normalizePromptDefinition(match.raw, r.source, match.sourcePath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidPromptDefinition, match.path, err)
|
||||
}
|
||||
return def, nil
|
||||
}
|
||||
|
||||
return nil, ErrPromptDefinitionNotFound
|
||||
}
|
||||
|
||||
func (r *fsRepository) GetPromptDefinition(ctx context.Context, id string, version string) (*domain.PromptDefinition, error) {
|
||||
return loadPromptDefinition(ctx, r.fsys, r.root, id, version)
|
||||
}
|
||||
|
||||
type promptDefinitionMatch struct {
|
||||
def *domain.PromptDefinition
|
||||
path string
|
||||
}
|
||||
|
||||
func loadPromptDefinitionFile(path string) (*promptDefinitionFile, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read prompt definition file: %w", err)
|
||||
}
|
||||
|
||||
var raw promptDefinitionFile
|
||||
decoder := yaml.NewDecoder(bytes.NewReader(data))
|
||||
decoder.KnownFields(true)
|
||||
if err := decoder.Decode(&raw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &raw, nil
|
||||
}
|
||||
|
||||
func promptDefinitionFileHasID(path string, id string) bool {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
var raw struct {
|
||||
ID string `yaml:"id"`
|
||||
}
|
||||
if err := yaml.NewDecoder(bytes.NewReader(data)).Decode(&raw); err != nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(raw.ID) == id
|
||||
}
|
||||
|
||||
func loadPromptDefinition(ctx context.Context, fsys fs.FS, root string, id string, version string) (*domain.PromptDefinition, error) {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidPromptDefinition)
|
||||
}
|
||||
if fsys == nil {
|
||||
return nil, fmt.Errorf("failed to read prompt definition directory: filesystem is nil")
|
||||
}
|
||||
|
||||
files, err := filecatalog.FindFSYAMLFiles(ctx, fsys, root)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read prompt definition directory: %w", err)
|
||||
}
|
||||
cleanRoot := filecatalog.CleanFSRoot(root)
|
||||
rootInfo, err := fs.Stat(fsys, cleanRoot)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read prompt definition directory: %w", err)
|
||||
}
|
||||
|
||||
var matches []promptDefinitionMatch
|
||||
for _, fullPath := range files {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
relPath := filecatalog.DisplayPath(root, fullPath)
|
||||
fileMatch := filecatalog.Stem(path.Base(fullPath)) == id
|
||||
data, err := fs.ReadFile(fsys, fullPath)
|
||||
if err != nil {
|
||||
if fileMatch {
|
||||
return nil, fmt.Errorf("%w: %s: failed to read prompt definition file: %v", ErrInvalidYAML, relPath, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
raw, err := decodePromptDefinition(data)
|
||||
if err != nil {
|
||||
if fileMatch || promptDefinitionDataHasID(data, id) {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
def, err := normalizePromptDefinitionFromFS(raw, fsys, root, fullPath, rootInfo.IsDir())
|
||||
if err != nil {
|
||||
if fileMatch || strings.TrimSpace(raw.ID) == id {
|
||||
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidPromptDefinition, relPath, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if def.ID != id {
|
||||
continue
|
||||
}
|
||||
if version != "" && def.Version != version {
|
||||
continue
|
||||
}
|
||||
matches = append(matches, promptDefinitionMatch{
|
||||
def: def,
|
||||
path: relPath,
|
||||
})
|
||||
}
|
||||
|
||||
if len(matches) > 1 {
|
||||
paths := make([]string, 0, len(matches))
|
||||
for _, match := range matches {
|
||||
paths = append(paths, match.path)
|
||||
}
|
||||
if version != "" {
|
||||
return nil, fmt.Errorf("%w: duplicate prompt definition id %q version %q found in: %s", ErrInvalidPromptDefinition, id, version, strings.Join(paths, ", "))
|
||||
}
|
||||
return nil, fmt.Errorf("%w: duplicate prompt definition id %q found in: %s", ErrInvalidPromptDefinition, id, strings.Join(paths, ", "))
|
||||
}
|
||||
|
||||
if len(matches) == 1 {
|
||||
return matches[0].def, nil
|
||||
}
|
||||
|
||||
return nil, ErrPromptDefinitionNotFound
|
||||
raw *promptDefinitionFile
|
||||
sourcePath string
|
||||
path string
|
||||
}
|
||||
|
||||
func decodePromptDefinition(data []byte) (*promptDefinitionFile, error) {
|
||||
@@ -270,59 +172,44 @@ func decodePromptDefinition(data []byte) (*promptDefinitionFile, error) {
|
||||
if err := decoder.Decode(&raw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var additional yaml.Node
|
||||
if err := decoder.Decode(&additional); err != io.EOF {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, errors.New("prompt definition file must contain exactly one YAML document")
|
||||
}
|
||||
return &raw, nil
|
||||
}
|
||||
|
||||
func promptDefinitionDataHasID(data []byte, id string) bool {
|
||||
func promptDefinitionDataMatches(data []byte, id string, version string) bool {
|
||||
var raw struct {
|
||||
ID string `yaml:"id"`
|
||||
ID string `yaml:"id"`
|
||||
Version string `yaml:"version"`
|
||||
}
|
||||
if err := yaml.NewDecoder(bytes.NewReader(data)).Decode(&raw); err != nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(raw.ID) == id
|
||||
return promptSelectorMatches(raw.ID, raw.Version, id, version)
|
||||
}
|
||||
|
||||
func normalizePromptDefinition(raw *promptDefinitionFile, sourcePath string) (*domain.PromptDefinition, error) {
|
||||
promptDir := filepath.Dir(sourcePath)
|
||||
return normalizePromptDefinitionWithContent(raw, func(contentFile string) (string, string, error) {
|
||||
resolvedPath := strings.TrimSpace(contentFile)
|
||||
if !filepath.IsAbs(resolvedPath) {
|
||||
resolvedPath = filepath.Join(promptDir, resolvedPath)
|
||||
}
|
||||
resolvedPath = filepath.Clean(resolvedPath)
|
||||
|
||||
body, err := os.ReadFile(resolvedPath)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return string(body), resolvedPath, nil
|
||||
})
|
||||
func promptDefinitionMatches(raw *promptDefinitionFile, id string, version string) bool {
|
||||
if raw == nil {
|
||||
return false
|
||||
}
|
||||
return promptSelectorMatches(raw.ID, raw.Version, id, version)
|
||||
}
|
||||
|
||||
func normalizePromptDefinitionFromFS(raw *promptDefinitionFile, fsys fs.FS, root string, sourcePath string, rootIsDir bool) (*domain.PromptDefinition, error) {
|
||||
promptDir := path.Dir(sourcePath)
|
||||
return normalizePromptDefinitionWithContent(raw, func(contentFile string) (string, string, error) {
|
||||
var resolvedPath string
|
||||
if rootIsDir {
|
||||
var err error
|
||||
resolvedPath, _, err = filecatalog.ResolveFSPath(root, promptDir, contentFile)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
} else {
|
||||
resolvedPath = strings.TrimSpace(contentFile)
|
||||
if !path.IsAbs(resolvedPath) {
|
||||
resolvedPath = path.Join(promptDir, resolvedPath)
|
||||
}
|
||||
resolvedPath = strings.TrimPrefix(path.Clean(resolvedPath), "/")
|
||||
}
|
||||
func promptSelectorMatches(rawID string, rawVersion string, id string, version string) bool {
|
||||
if strings.TrimSpace(rawID) != id {
|
||||
return false
|
||||
}
|
||||
return version == "" || strings.TrimSpace(rawVersion) == version
|
||||
}
|
||||
|
||||
body, err := fs.ReadFile(fsys, resolvedPath)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return string(body), resolvedPath, nil
|
||||
func normalizePromptDefinition(raw *promptDefinitionFile, sourceRoot contentSourceRoot, sourcePath string) (*domain.PromptDefinition, error) {
|
||||
return normalizePromptDefinitionWithContent(raw, func(contentFile string) (string, string, error) {
|
||||
return sourceRoot.readContentFile(sourcePath, contentFile)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -402,17 +289,14 @@ func normalizePromptDefinitionWithContent(raw *promptDefinitionFile, readContent
|
||||
})
|
||||
}
|
||||
|
||||
if !isValidOutputFormat(raw.Output.Format) {
|
||||
return nil, fmt.Errorf("invalid output format: %q", raw.Output.Format)
|
||||
outputContract := domain.OutputContract{
|
||||
Format: raw.Output.Format,
|
||||
ValidationMode: raw.Output.ValidationMode,
|
||||
SchemaPath: strings.TrimSpace(raw.Output.SchemaPath),
|
||||
RepairAttempts: raw.Output.RepairAttempts,
|
||||
}
|
||||
if !isValidValidationMode(raw.Output.ValidationMode) {
|
||||
return nil, fmt.Errorf("invalid validation mode: %q", raw.Output.ValidationMode)
|
||||
}
|
||||
if raw.Output.ValidationMode == domain.ValidationJSONSchema && strings.TrimSpace(raw.Output.SchemaPath) == "" {
|
||||
return nil, errors.New("output.schema_path is required when output.validation_mode is json_schema")
|
||||
}
|
||||
if raw.Output.RepairAttempts < 0 {
|
||||
return nil, errors.New("output.repair_attempts must be greater than or equal to 0")
|
||||
if err := domain.ValidateOutputContract(outputContract); err != nil {
|
||||
return nil, fmt.Errorf("output: %w", err)
|
||||
}
|
||||
|
||||
defaultProfile := ""
|
||||
@@ -432,12 +316,7 @@ func normalizePromptDefinitionWithContent(raw *promptDefinitionFile, readContent
|
||||
Inputs: inputs,
|
||||
Templates: templates,
|
||||
OutputFormat: raw.Output.Format,
|
||||
Validation: domain.OutputContract{
|
||||
Format: raw.Output.Format,
|
||||
ValidationMode: raw.Output.ValidationMode,
|
||||
SchemaPath: strings.TrimSpace(raw.Output.SchemaPath),
|
||||
RepairAttempts: raw.Output.RepairAttempts,
|
||||
},
|
||||
Validation: outputContract,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -464,21 +343,3 @@ func normalizeCacheControl(raw *cacheControlFile) (*domain.CacheControl, error)
|
||||
TTL: ttl,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func isValidOutputFormat(f domain.OutputFormat) bool {
|
||||
switch f {
|
||||
case domain.FormatText, domain.FormatMarkdown, domain.FormatJSON:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isValidValidationMode(m domain.ValidationMode) bool {
|
||||
switch m {
|
||||
case domain.ValidationNone, domain.ValidationBasic, domain.ValidationJSON, domain.ValidationJSONSchema:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,23 +3,51 @@ package promptdef
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
func TestFilesystemRepository_GetPromptDefinition(t *testing.T) {
|
||||
func TestPromptRepositoryDefinitionFixtures(t *testing.T) {
|
||||
sources := []struct {
|
||||
name string
|
||||
newRepository func(string) Repository
|
||||
contentPathsAreFull bool
|
||||
}{
|
||||
{
|
||||
name: "operating system",
|
||||
newRepository: NewFilesystemRepository,
|
||||
contentPathsAreFull: true,
|
||||
},
|
||||
{
|
||||
name: "filesystem",
|
||||
newRepository: func(root string) Repository {
|
||||
return NewFSRepository(os.DirFS(root), ".")
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, source := range sources {
|
||||
t.Run(source.name, func(t *testing.T) {
|
||||
testPromptRepositoryDefinitionFixtures(t, source.newRepository, source.contentPathsAreFull)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func testPromptRepositoryDefinitionFixtures(t *testing.T, newRepository func(string) Repository, contentPathsAreFull bool) {
|
||||
tmpDir := t.TempDir()
|
||||
if err := copyTree("testdata", tmpDir); err != nil {
|
||||
t.Fatalf("failed to copy testdata: %v", err)
|
||||
}
|
||||
|
||||
repo := NewFilesystemRepository(tmpDir)
|
||||
repo := newRepository(tmpDir)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("valid inline prompt", func(t *testing.T) {
|
||||
@@ -64,8 +92,8 @@ func TestFilesystemRepository_GetPromptDefinition(t *testing.T) {
|
||||
if p.Templates[1].ContentFile == "" {
|
||||
t.Fatal("expected ContentFile source metadata to be preserved")
|
||||
}
|
||||
if !filepath.IsAbs(p.Templates[1].ContentFile) {
|
||||
t.Fatalf("expected resolved content_file path to be absolute, got %q", p.Templates[1].ContentFile)
|
||||
if filepath.IsAbs(p.Templates[1].ContentFile) != contentPathsAreFull {
|
||||
t.Fatalf("unexpected content_file path representation: %q", p.Templates[1].ContentFile)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -287,20 +315,20 @@ output:
|
||||
targetErr error
|
||||
errSubstrs []string
|
||||
}{
|
||||
{name: "invalid YAML", id: "invalid_yaml", targetErr: ErrInvalidYAML},
|
||||
{name: "missing id", id: "missing_id", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"id is required"}},
|
||||
{name: "no messages", id: "no_messages", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"at least one message is required"}},
|
||||
{name: "both content and content_file", id: "both_content_and_content_file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"exactly one"}},
|
||||
{name: "neither content nor content_file", id: "neither_content_nor_content_file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"exactly one"}},
|
||||
{name: "missing content_file", id: "missing_content_file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"failed to read content_file"}},
|
||||
{name: "duplicate input names", id: "duplicate_input_names", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"duplicate input name"}},
|
||||
{name: "invalid validation mode", id: "invalid_validation_mode", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"invalid validation mode"}},
|
||||
{name: "json_schema without schema_path", id: "json_schema_without_schema_path", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"schema_path"}},
|
||||
{name: "unknown input field", id: "unknown_input_field", targetErr: ErrInvalidYAML, errSubstrs: []string{"field unknown_input_setting not found"}},
|
||||
{name: "empty cache control type", id: "empty_cache_control_type", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "type is required"}},
|
||||
{name: "unsupported cache control type", id: "unsupported_cache_control_type", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "unsupported type"}},
|
||||
{name: "unsupported cache control ttl", id: "unsupported_cache_control_ttl", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "unsupported ttl"}},
|
||||
{name: "unknown cache control field", id: "unknown_cache_control_field", targetErr: ErrInvalidYAML, errSubstrs: []string{"field unexpected not found"}},
|
||||
{name: "unidentifiable invalid YAML is unrelated", id: "invalid-yaml", targetErr: ErrPromptDefinitionNotFound},
|
||||
{name: "missing id is not selected by filename", id: "missing_id", targetErr: ErrPromptDefinitionNotFound},
|
||||
{name: "no messages", id: "no-messages", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"at least one message is required"}},
|
||||
{name: "both content and content_file", id: "both-content-and-content-file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"exactly one"}},
|
||||
{name: "neither content nor content_file", id: "neither-content-nor-content-file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"exactly one"}},
|
||||
{name: "missing content_file", id: "missing-content-file", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"failed to read content_file"}},
|
||||
{name: "duplicate input names", id: "duplicate-input-names", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"duplicate input name"}},
|
||||
{name: "invalid validation mode", id: "invalid-validation-mode", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"invalid validation mode"}},
|
||||
{name: "json_schema without schema_path", id: "json-schema-without-schema-path", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"schema_path"}},
|
||||
{name: "unknown input field", id: "unknown-input-field", targetErr: ErrInvalidYAML, errSubstrs: []string{"field unknown_input_setting not found"}},
|
||||
{name: "empty cache control type", id: "empty-cache-control-type", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "type is required"}},
|
||||
{name: "unsupported cache control type", id: "unsupported-cache-control-type", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "unsupported type"}},
|
||||
{name: "unsupported cache control ttl", id: "unsupported-cache-control-ttl", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"cache_control", "unsupported ttl"}},
|
||||
{name: "unknown cache control field", id: "unknown-cache-control-field", targetErr: ErrInvalidYAML, errSubstrs: []string{"field unexpected not found"}},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
@@ -396,7 +424,7 @@ output:
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
repo := NewFSRepository(fstest.MapFS{
|
||||
fsys := &recordingFS{FS: fstest.MapFS{
|
||||
"prompts/prompt.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: fs-escaped-prompt
|
||||
version: "1.0.0"
|
||||
@@ -409,7 +437,8 @@ output:
|
||||
repair_attempts: 0
|
||||
`)},
|
||||
"outside.tmpl": &fstest.MapFile{Data: []byte(`Outside root.`)},
|
||||
}, "prompts")
|
||||
}}
|
||||
repo := NewFSRepository(fsys, "prompts")
|
||||
|
||||
_, err := repo.GetPromptDefinition(context.Background(), "fs-escaped-prompt", "")
|
||||
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
||||
@@ -418,64 +447,611 @@ output:
|
||||
if !strings.Contains(err.Error(), tc.wantErr) {
|
||||
t.Fatalf("expected error to contain %q, got %v", tc.wantErr, err)
|
||||
}
|
||||
if fsys.wasOpened("outside.tmpl") {
|
||||
t.Fatal("rejected content path opened the outside file")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFSRepositoryRejectsDuplicatePromptIDs(t *testing.T) {
|
||||
repo := NewFSRepository(fstest.MapFS{
|
||||
"one.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: duplicate-fs-prompt
|
||||
version: "1.0.0"
|
||||
messages:
|
||||
- role: user
|
||||
content: First.
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
repair_attempts: 0
|
||||
`)},
|
||||
"nested/two.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: duplicate-fs-prompt
|
||||
version: "1.0.0"
|
||||
messages:
|
||||
- role: user
|
||||
content: Second.
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
repair_attempts: 0
|
||||
`)},
|
||||
}, ".")
|
||||
type recordingFS struct {
|
||||
fs.FS
|
||||
mu sync.Mutex
|
||||
opened []string
|
||||
}
|
||||
|
||||
_, err := repo.GetPromptDefinition(context.Background(), "duplicate-fs-prompt", "")
|
||||
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
||||
t.Fatalf("expected ErrInvalidPromptDefinition, got %v", err)
|
||||
func TestPromptRepositoryReturnsDefinitionReadFailures(t *testing.T) {
|
||||
readErr := errors.New("definition read failed")
|
||||
fsys := &definitionReadFailureFS{
|
||||
FS: fstest.MapFS{
|
||||
"prompts/target.yaml": &fstest.MapFile{Data: []byte("unread")},
|
||||
},
|
||||
target: "prompts/target.yaml",
|
||||
err: readErr,
|
||||
}
|
||||
if !strings.Contains(err.Error(), "one.yaml") || !strings.Contains(err.Error(), "nested/two.yaml") {
|
||||
t.Fatalf("expected duplicate paths in error, got %v", err)
|
||||
repo := NewFSRepository(fsys, "prompts")
|
||||
|
||||
definition, err := repo.GetPromptDefinition(context.Background(), "target", "1")
|
||||
if definition != nil || !errors.Is(err, readErr) {
|
||||
t.Fatalf("GetPromptDefinition() = (%#v, %v), want nil and definition read error", definition, err)
|
||||
}
|
||||
if errors.Is(err, ErrPromptDefinitionNotFound) {
|
||||
t.Fatalf("definition read error was classified as absence: %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "target.yaml") {
|
||||
t.Fatalf("definition read error lacks source context: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFSRepositoryRejectsUnknownYAMLFields(t *testing.T) {
|
||||
repo := NewFSRepository(fstest.MapFS{
|
||||
"not_named_like_id.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: strict-fs-prompt
|
||||
version: "1.0.0"
|
||||
unknown: true
|
||||
type definitionReadFailureFS struct {
|
||||
fs.FS
|
||||
target string
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *definitionReadFailureFS) Open(name string) (fs.File, error) {
|
||||
if name == f.target {
|
||||
return nil, f.err
|
||||
}
|
||||
return f.FS.Open(name)
|
||||
}
|
||||
|
||||
func (f *recordingFS) Open(name string) (fs.File, error) {
|
||||
f.mu.Lock()
|
||||
f.opened = append(f.opened, name)
|
||||
f.mu.Unlock()
|
||||
return f.FS.Open(name)
|
||||
}
|
||||
|
||||
func (f *recordingFS) wasOpened(name string) bool {
|
||||
return f.openCount(name) > 0
|
||||
}
|
||||
|
||||
func (f *recordingFS) openCount(name string) int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
count := 0
|
||||
for _, opened := range f.opened {
|
||||
if opened == name {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func TestPromptRepositorySelectionUsesYAMLMetadata(t *testing.T) {
|
||||
const validDefinition = `
|
||||
id: selected-prompt
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: Invalid.
|
||||
content: selected
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
repair_attempts: 0
|
||||
`)},
|
||||
}, ".")
|
||||
`
|
||||
tests := []struct {
|
||||
name string
|
||||
files map[string]string
|
||||
wantErr error
|
||||
diagnostics []string
|
||||
wantContent string
|
||||
}{
|
||||
{
|
||||
name: "same-stem strict error with different YAML ID is unrelated",
|
||||
files: map[string]string{
|
||||
"selected-prompt.yaml": `
|
||||
id: another-prompt
|
||||
version: "1"
|
||||
unknown: true
|
||||
`,
|
||||
"valid.yaml": validDefinition,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "unidentifiable same-stem YAML is unrelated",
|
||||
files: map[string]string{
|
||||
"selected-prompt.yaml": "id: [",
|
||||
"valid.yaml": validDefinition,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "same ID invalid different version is unrelated",
|
||||
files: map[string]string{
|
||||
"invalid-version.yaml": `
|
||||
id: selected-prompt
|
||||
version: "2"
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
"valid.yaml": validDefinition,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "selected content is resolved relative to its definition",
|
||||
files: map[string]string{
|
||||
"nested/selected.yaml": `
|
||||
id: selected-prompt
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: ./content/selected.tmpl
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
"nested/content/selected.tmpl": "selected from file",
|
||||
},
|
||||
wantContent: "selected from file",
|
||||
},
|
||||
{
|
||||
name: "duplicate selected definitions are ambiguous",
|
||||
files: map[string]string{
|
||||
"selected-a.yaml": validDefinition,
|
||||
"nested/selected-b.yaml": `
|
||||
id: selected-prompt
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: duplicate
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidPromptDefinition,
|
||||
diagnostics: []string{"duplicate prompt definition id", "selected-a.yaml", "nested/selected-b.yaml"},
|
||||
},
|
||||
{
|
||||
name: "selected strict error is authoritative",
|
||||
files: map[string]string{
|
||||
"selected-strict.yaml": `
|
||||
id: selected-prompt
|
||||
version: "1"
|
||||
unknown: true
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidYAML,
|
||||
diagnostics: []string{"selected-strict.yaml"},
|
||||
},
|
||||
{
|
||||
name: "selected semantic error is authoritative",
|
||||
files: map[string]string{
|
||||
"selected-invalid.yaml": `
|
||||
id: selected-prompt
|
||||
version: "1"
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidPromptDefinition,
|
||||
diagnostics: []string{"selected-invalid.yaml", "at least one message"},
|
||||
},
|
||||
{
|
||||
name: "selected content error includes definition context",
|
||||
files: map[string]string{
|
||||
"selected-missing-content.yaml": `
|
||||
id: selected-prompt
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: missing.tmpl
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
},
|
||||
wantErr: ErrInvalidPromptDefinition,
|
||||
diagnostics: []string{"selected-missing-content.yaml", "failed to read content_file", "missing.tmpl"},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := repo.GetPromptDefinition(context.Background(), "strict-fs-prompt", "")
|
||||
if !errors.Is(err, ErrInvalidYAML) {
|
||||
t.Fatalf("expected ErrInvalidYAML, got %v", err)
|
||||
for _, source := range promptRepositorySources() {
|
||||
for _, tc := range tests {
|
||||
t.Run(source.name+"/"+tc.name, func(t *testing.T) {
|
||||
repo := source.newRepository(t, tc.files)
|
||||
got, err := repo.GetPromptDefinition(context.Background(), "selected-prompt", "1")
|
||||
if tc.wantErr != nil {
|
||||
if !errors.Is(err, tc.wantErr) {
|
||||
t.Fatalf("expected %v, got %v", tc.wantErr, err)
|
||||
}
|
||||
for _, diagnostic := range tc.diagnostics {
|
||||
if !strings.Contains(err.Error(), diagnostic) {
|
||||
t.Fatalf("expected error to contain %q, got %v", diagnostic, err)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("load selected prompt: %v", err)
|
||||
}
|
||||
wantContent := tc.wantContent
|
||||
if wantContent == "" {
|
||||
wantContent = "selected"
|
||||
}
|
||||
if got.ID != "selected-prompt" || got.Version != "1" || got.Templates[0].Content != wantContent {
|
||||
t.Fatalf("unexpected selected definition: %+v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptRepositoryHonorsCancellation(t *testing.T) {
|
||||
for _, source := range promptRepositorySources() {
|
||||
t.Run(source.name, func(t *testing.T) {
|
||||
repo := source.newRepository(t, map[string]string{
|
||||
"definition.yaml": `
|
||||
id: cancelled-prompt
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: selected
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
})
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
_, err := repo.GetPromptDefinition(ctx, "cancelled-prompt", "1")
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptRepositoryRequiresOneYAMLDocument(t *testing.T) {
|
||||
const definition = `
|
||||
id: one-document
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: selected
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`
|
||||
tests := []struct {
|
||||
name string
|
||||
suffix string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "comments and trailing whitespace", suffix: "\n# trailing comment\n\n"},
|
||||
{name: "second populated document", suffix: "\n---\nid: another\n", wantErr: true},
|
||||
{name: "second empty document", suffix: "\n---\n", wantErr: true},
|
||||
{name: "malformed trailing YAML", suffix: "\n---\n[", wantErr: true},
|
||||
}
|
||||
|
||||
for _, source := range promptRepositorySources() {
|
||||
for _, tc := range tests {
|
||||
t.Run(source.name+"/"+tc.name, func(t *testing.T) {
|
||||
repo := source.newRepository(t, map[string]string{"definition.yaml": definition + tc.suffix})
|
||||
_, err := repo.GetPromptDefinition(context.Background(), "one-document", "1")
|
||||
if tc.wantErr {
|
||||
if !errors.Is(err, ErrInvalidYAML) {
|
||||
t.Fatalf("expected ErrInvalidYAML, got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "definition.yaml") {
|
||||
t.Fatalf("expected source path in error, got %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("load one document: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptRepositoryReadsOnlySelectedContent(t *testing.T) {
|
||||
fsys := &recordingFS{FS: fstest.MapFS{
|
||||
"target.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: target
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: target.tmpl
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
"target.tmpl": &fstest.MapFile{Data: []byte("selected")},
|
||||
"unrelated.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: unrelated
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: unrelated.tmpl
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
"unrelated.tmpl": &fstest.MapFile{Data: []byte("unrelated")},
|
||||
"other-version.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: target
|
||||
version: "2"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: other-version.tmpl
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
"other-version.tmpl": &fstest.MapFile{Data: []byte("other version")},
|
||||
}}
|
||||
repo := NewFSRepository(fsys, ".")
|
||||
|
||||
for lookup := 1; lookup <= 2; lookup++ {
|
||||
got, err := repo.GetPromptDefinition(context.Background(), "target", "1")
|
||||
if err != nil {
|
||||
t.Fatalf("lookup %d: %v", lookup, err)
|
||||
}
|
||||
if got.Templates[0].Content != "selected" {
|
||||
t.Fatalf("lookup %d content = %q", lookup, got.Templates[0].Content)
|
||||
}
|
||||
if count := fsys.openCount("target.tmpl"); count != lookup {
|
||||
t.Fatalf("selected content opens after lookup %d = %d, want %d", lookup, count, lookup)
|
||||
}
|
||||
for _, name := range []string{"unrelated.tmpl", "other-version.tmpl"} {
|
||||
if count := fsys.openCount(name); count != 0 {
|
||||
t.Fatalf("unselected content %q opened %d times", name, count)
|
||||
}
|
||||
}
|
||||
for _, name := range []string{"target.yaml", "unrelated.yaml", "other-version.yaml"} {
|
||||
if count := fsys.openCount(name); count != lookup {
|
||||
t.Fatalf("metadata %q opens after lookup %d = %d, want %d", name, lookup, count, lookup)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptDefinitionNormalizationRules(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
definition string
|
||||
wantErr bool
|
||||
wantDiagnostic string
|
||||
wantSchemaPath string
|
||||
}{
|
||||
{
|
||||
name: "missing version",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
messages:
|
||||
- role: user
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "version",
|
||||
},
|
||||
{
|
||||
name: "blank input name",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
version: "1"
|
||||
inputs:
|
||||
- name: " "
|
||||
messages:
|
||||
- role: user
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "input 0",
|
||||
},
|
||||
{
|
||||
name: "blank message role",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
version: "1"
|
||||
messages:
|
||||
- role: " "
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "role",
|
||||
},
|
||||
{
|
||||
name: "invalid output format",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: test
|
||||
output:
|
||||
format: binary
|
||||
validation_mode: none
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "format",
|
||||
},
|
||||
{
|
||||
name: "negative 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",
|
||||
},
|
||||
{
|
||||
name: "explicit blank default profile",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
version: "1"
|
||||
default_profile: " "
|
||||
messages:
|
||||
- role: user
|
||||
content: test
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`,
|
||||
wantErr: true,
|
||||
wantDiagnostic: "default_profile",
|
||||
},
|
||||
{
|
||||
name: "schema path normalization",
|
||||
definition: `
|
||||
id: normalization-rule
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: test
|
||||
output:
|
||||
format: json
|
||||
validation_mode: json_schema
|
||||
schema_path: ' schema.json '
|
||||
`,
|
||||
wantSchemaPath: "schema.json",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
repo := NewFSRepository(fstest.MapFS{
|
||||
"definition.yaml": &fstest.MapFile{Data: []byte(tt.definition)},
|
||||
}, ".")
|
||||
|
||||
got, err := repo.GetPromptDefinition(context.Background(), "normalization-rule", "")
|
||||
if tt.wantErr {
|
||||
if !errors.Is(err, ErrInvalidPromptDefinition) {
|
||||
t.Fatalf("expected ErrInvalidPromptDefinition, got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), tt.wantDiagnostic) {
|
||||
t.Fatalf("expected error containing %q, got %v", tt.wantDiagnostic, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("load prompt definition: %v", err)
|
||||
}
|
||||
if got.Validation.SchemaPath != tt.wantSchemaPath {
|
||||
t.Fatalf("schema path = %q, want %q", got.Validation.SchemaPath, tt.wantSchemaPath)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type promptRepositorySource struct {
|
||||
name string
|
||||
newRepository func(t *testing.T, files map[string]string) Repository
|
||||
}
|
||||
|
||||
func promptRepositorySources() []promptRepositorySource {
|
||||
return []promptRepositorySource{
|
||||
{
|
||||
name: "operating system",
|
||||
newRepository: func(t *testing.T, files map[string]string) Repository {
|
||||
t.Helper()
|
||||
root := t.TempDir()
|
||||
for name, content := range files {
|
||||
filePath := filepath.Join(root, filepath.FromSlash(name))
|
||||
if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil {
|
||||
t.Fatalf("create prompt directory: %v", err)
|
||||
}
|
||||
writePromptTestFile(t, filePath, content)
|
||||
}
|
||||
return NewFilesystemRepository(root)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "filesystem",
|
||||
newRepository: func(t *testing.T, files map[string]string) Repository {
|
||||
t.Helper()
|
||||
fsys := make(fstest.MapFS, len(files))
|
||||
for name, content := range files {
|
||||
fsys[name] = &fstest.MapFile{Data: []byte(strings.TrimLeft(content, "\n"))}
|
||||
}
|
||||
return NewFSRepository(fsys, ".")
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkPromptRepositoryLookup(b *testing.B) {
|
||||
for _, size := range []int{10, 1000} {
|
||||
b.Run(fmt.Sprintf("catalog-%d", size), func(b *testing.B) {
|
||||
files := fstest.MapFS{
|
||||
"target.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: target
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: target.tmpl
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
"target.tmpl": &fstest.MapFile{Data: []byte("selected")},
|
||||
}
|
||||
metadataNames := []string{"target.yaml"}
|
||||
contentNames := make([]string, 0, size-1)
|
||||
for i := 1; i < size; i++ {
|
||||
definitionName := fmt.Sprintf("prompt-%04d.yaml", i)
|
||||
contentName := fmt.Sprintf("prompt-%04d.tmpl", i)
|
||||
files[definitionName] = &fstest.MapFile{Data: []byte(fmt.Sprintf(`
|
||||
id: prompt-%04d
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: %s
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`, i, contentName))}
|
||||
files[contentName] = &fstest.MapFile{Data: []byte("unrelated")}
|
||||
metadataNames = append(metadataNames, definitionName)
|
||||
contentNames = append(contentNames, contentName)
|
||||
}
|
||||
|
||||
fsys := &recordingFS{FS: files}
|
||||
repo := NewFSRepository(fsys, ".")
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
if _, err := repo.GetPromptDefinition(context.Background(), "target", "1"); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
|
||||
if got := fsys.openCount("target.tmpl"); got != b.N {
|
||||
b.Fatalf("selected content opens = %d, want %d", got, b.N)
|
||||
}
|
||||
for _, name := range contentNames {
|
||||
if got := fsys.openCount(name); got != 0 {
|
||||
b.Fatalf("unrelated content %q opened %d times", name, got)
|
||||
}
|
||||
}
|
||||
for _, name := range metadataNames {
|
||||
if got := fsys.openCount(name); got != b.N {
|
||||
b.Fatalf("metadata %q opens = %d, want %d", name, got, b.N)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
63
internal/promptdef/source.go
Normal file
63
internal/promptdef/source.go
Normal file
@@ -0,0 +1,63 @@
|
||||
package promptdef
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io/fs"
|
||||
"os"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/filecatalog"
|
||||
)
|
||||
|
||||
type promptDefinitionSource interface {
|
||||
contentSourceRoot
|
||||
findYAMLFiles(context.Context) ([]string, error)
|
||||
readDefinition(string) ([]byte, error)
|
||||
displayPath(string) string
|
||||
}
|
||||
|
||||
type osPromptSource struct {
|
||||
root string
|
||||
contentRoot osContentSourceRoot
|
||||
}
|
||||
|
||||
func (s osPromptSource) findYAMLFiles(ctx context.Context) ([]string, error) {
|
||||
return filecatalog.FindYAMLFiles(ctx, s.root)
|
||||
}
|
||||
|
||||
func (s osPromptSource) readDefinition(name string) ([]byte, error) {
|
||||
return os.ReadFile(name)
|
||||
}
|
||||
|
||||
func (s osPromptSource) displayPath(name string) string {
|
||||
return filecatalog.RelativePath(s.root, name)
|
||||
}
|
||||
|
||||
func (s osPromptSource) readContentFile(sourcePath string, contentFile string) (string, string, error) {
|
||||
return s.contentRoot.readContentFile(sourcePath, contentFile)
|
||||
}
|
||||
|
||||
type fsPromptSource struct {
|
||||
fsys fs.FS
|
||||
root string
|
||||
contentRoot contentSourceRoot
|
||||
}
|
||||
|
||||
func (s fsPromptSource) findYAMLFiles(ctx context.Context) ([]string, error) {
|
||||
if s.fsys == nil {
|
||||
return nil, errors.New("filesystem is nil")
|
||||
}
|
||||
return filecatalog.FindFSYAMLFiles(ctx, s.fsys, s.root)
|
||||
}
|
||||
|
||||
func (s fsPromptSource) readDefinition(name string) ([]byte, error) {
|
||||
return fs.ReadFile(s.fsys, name)
|
||||
}
|
||||
|
||||
func (s fsPromptSource) displayPath(name string) string {
|
||||
return filecatalog.DisplayPath(s.root, name)
|
||||
}
|
||||
|
||||
func (s fsPromptSource) readContentFile(sourcePath string, contentFile string) (string, string, error) {
|
||||
return s.contentRoot.readContentFile(sourcePath, contentFile)
|
||||
}
|
||||
24
internal/usecase/capacity_error.go
Normal file
24
internal/usecase/capacity_error.go
Normal file
@@ -0,0 +1,24 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/capacity"
|
||||
)
|
||||
|
||||
// CapacityError identifies bounded admission rejected for one selected backend.
|
||||
type CapacityError struct {
|
||||
BackendID string
|
||||
}
|
||||
|
||||
func (e *CapacityError) Error() string {
|
||||
if e == nil || strings.TrimSpace(e.BackendID) == "" {
|
||||
return capacity.ErrCapacityExceeded.Error()
|
||||
}
|
||||
return fmt.Sprintf("backend %q admission: %v", e.BackendID, capacity.ErrCapacityExceeded)
|
||||
}
|
||||
|
||||
func (e *CapacityError) Unwrap() error {
|
||||
return capacity.ErrCapacityExceeded
|
||||
}
|
||||
83
internal/usecase/execution_settings_test.go
Normal file
83
internal/usecase/execution_settings_test.go
Normal file
@@ -0,0 +1,83 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
func TestRunnerPrepareExecutionRejectsInvalidExecutionSettings(t *testing.T) {
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
&fakeLLM{forbid: true},
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
|
||||
_, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Execution: &domain.ExecutionTargetOverride{TopP: float64Ptr(math.Inf(-1))},
|
||||
})
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerPrepareExecutionValidatesAndNormalizesRequestEndpoints(t *testing.T) {
|
||||
newRunner := func() *Runner {
|
||||
return NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
&fakeLLM{forbid: true},
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
}
|
||||
|
||||
invalidEndpoints := []string{
|
||||
"/v1",
|
||||
"https:///v1",
|
||||
"ftp://provider.example/v1",
|
||||
"https://user@provider.example/v1",
|
||||
"https://provider.example/v1?mode=chat",
|
||||
"https://provider.example/v1#chat",
|
||||
}
|
||||
for _, endpoint := range invalidEndpoints {
|
||||
t.Run(endpoint, func(t *testing.T) {
|
||||
_, err := newRunner().PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Execution: &domain.ExecutionTargetOverride{Endpoint: endpoint},
|
||||
})
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
prepared, err := newRunner().PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Execution: &domain.ExecutionTargetOverride{Endpoint: " https://provider.example/nested/v1 "},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare normalized endpoint: %v", err)
|
||||
}
|
||||
if got := prepared.Details().EffectiveModelParams.Endpoint; got != "https://provider.example/nested/v1" {
|
||||
t.Fatalf("effective endpoint = %q", got)
|
||||
}
|
||||
}
|
||||
19
internal/usecase/generation_request.go
Normal file
19
internal/usecase/generation_request.go
Normal file
@@ -0,0 +1,19 @@
|
||||
package usecase
|
||||
|
||||
import "gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
|
||||
func newGenerationRequest(
|
||||
prompt domain.RenderedPrompt,
|
||||
sessionID string,
|
||||
target domain.ExecutionTarget,
|
||||
targetPresence domain.ExecutionTargetPresence,
|
||||
structuredOutput *domain.StructuredOutputSpec,
|
||||
) domain.GenerateRequest {
|
||||
prompt.SessionID = sessionID
|
||||
return domain.GenerateRequest{
|
||||
Prompt: prompt,
|
||||
Target: target,
|
||||
TargetPresence: targetPresence,
|
||||
StructuredOutput: structuredOutput,
|
||||
}
|
||||
}
|
||||
185
internal/usecase/output_contract_test.go
Normal file
185
internal/usecase/output_contract_test.go
Normal file
@@ -0,0 +1,185 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
type outputContractTestCollaborators struct {
|
||||
artifacts *fakeArtifactReader
|
||||
renderer *fakeRenderer
|
||||
llm *fakeLLM
|
||||
validator *recordingValidationPreparer
|
||||
admitter *fakeRunAdmitter
|
||||
}
|
||||
|
||||
func newOutputContractTestRunner() (*Runner, outputContractTestCollaborators) {
|
||||
collaborators := outputContractTestCollaborators{
|
||||
artifacts: defaultArtifactReader(),
|
||||
renderer: defaultRenderer(),
|
||||
llm: &fakeLLM{forbid: true},
|
||||
validator: &recordingValidationPreparer{plan: &recordingPreparedValidation{}},
|
||||
admitter: &fakeRunAdmitter{},
|
||||
}
|
||||
return NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
collaborators.artifacts,
|
||||
collaborators.renderer,
|
||||
collaborators.llm,
|
||||
collaborators.validator,
|
||||
collaborators.admitter,
|
||||
), collaborators
|
||||
}
|
||||
|
||||
func TestRunnerPreparationNormalizesOutputContractConsistently(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
override domain.OutputContract
|
||||
want domain.OutputContract
|
||||
}{
|
||||
{
|
||||
name: "empty replacement format defaults to text",
|
||||
override: domain.OutputContract{ValidationMode: domain.ValidationNone},
|
||||
want: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone},
|
||||
},
|
||||
{
|
||||
name: "markdown basic replacement",
|
||||
override: domain.OutputContract{Format: domain.FormatMarkdown, ValidationMode: domain.ValidationBasic},
|
||||
want: domain.OutputContract{Format: domain.FormatMarkdown, ValidationMode: domain.ValidationBasic},
|
||||
},
|
||||
{
|
||||
name: "json replacement preserves non-schema fields",
|
||||
override: domain.OutputContract{
|
||||
Format: domain.FormatJSON,
|
||||
ValidationMode: domain.ValidationJSON,
|
||||
SchemaPath: "ignored.json",
|
||||
RepairAttempts: 2,
|
||||
},
|
||||
want: domain.OutputContract{
|
||||
Format: domain.FormatJSON,
|
||||
ValidationMode: domain.ValidationJSON,
|
||||
SchemaPath: "ignored.json",
|
||||
RepairAttempts: 2,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
runner, _ := newOutputContractTestRunner()
|
||||
req := domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Validation: &tt.override,
|
||||
}
|
||||
|
||||
prepared, err := runner.Prepare(context.Background(), req)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare: %v", err)
|
||||
}
|
||||
preparedExecution, err := runner.PrepareExecution(context.Background(), req)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
details := preparedExecution.Details()
|
||||
if details == nil {
|
||||
t.Fatal("prepared execution returned nil details")
|
||||
}
|
||||
if prepared.OutputContract != tt.want || details.OutputContract != tt.want {
|
||||
t.Fatalf("output contracts = (%+v, %+v), want %+v", prepared.OutputContract, details.OutputContract, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerPreparationRejectsInvalidOutputContractsBeforeCompletion(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
override domain.OutputContract
|
||||
}{
|
||||
{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: "json schema without path", override: domain.OutputContract{Format: domain.FormatJSON, ValidationMode: domain.ValidationJSONSchema}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
for _, operation := range []string{"Prepare", "PrepareExecution"} {
|
||||
t.Run(tt.name+"/"+operation, func(t *testing.T) {
|
||||
runner, collaborators := newOutputContractTestRunner()
|
||||
req := domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Validation: &tt.override,
|
||||
}
|
||||
|
||||
var err error
|
||||
switch operation {
|
||||
case "Prepare":
|
||||
var prepared *domain.PreparedRun
|
||||
prepared, err = runner.Prepare(context.Background(), req)
|
||||
if prepared != nil {
|
||||
t.Fatalf("expected no partial prepared run, got %+v", prepared)
|
||||
}
|
||||
case "PrepareExecution":
|
||||
var prepared *PreparedExecution
|
||||
prepared, err = runner.PrepareExecution(context.Background(), req)
|
||||
if prepared != nil {
|
||||
t.Fatalf("expected no partial prepared execution, got %+v", prepared)
|
||||
}
|
||||
default:
|
||||
t.Fatalf("unknown operation %q", operation)
|
||||
}
|
||||
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
assertOutputContractCompletionSkipped(t, collaborators)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunRejectsInvalidOutputContractBeforeAdmission(t *testing.T) {
|
||||
runner, collaborators := newOutputContractTestRunner()
|
||||
result, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Validation: &domain.OutputContract{
|
||||
Format: domain.OutputFormat("binary"),
|
||||
ValidationMode: domain.ValidationNone,
|
||||
},
|
||||
})
|
||||
if result != nil {
|
||||
t.Fatalf("expected no partial result, got %+v", result)
|
||||
}
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
assertOutputContractCompletionSkipped(t, collaborators)
|
||||
}
|
||||
|
||||
func assertOutputContractCompletionSkipped(t *testing.T, collaborators outputContractTestCollaborators) {
|
||||
t.Helper()
|
||||
if collaborators.artifacts.calls != 0 || collaborators.renderer.calls != 0 ||
|
||||
collaborators.validator.prepareCalls != 0 || collaborators.validator.directValidateCalls != 0 ||
|
||||
len(collaborators.admitter.backendIDs) != 0 || collaborators.llm.calls != 0 {
|
||||
t.Fatalf(
|
||||
"invalid output contract reached downstream work: artifacts=%d renderer=%d prepare_validation=%d validation=%d admissions=%d generation=%d",
|
||||
collaborators.artifacts.calls,
|
||||
collaborators.renderer.calls,
|
||||
collaborators.validator.prepareCalls,
|
||||
collaborators.validator.directValidateCalls,
|
||||
len(collaborators.admitter.backendIDs),
|
||||
collaborators.llm.calls,
|
||||
)
|
||||
}
|
||||
}
|
||||
245
internal/usecase/prepared_execution.go
Normal file
245
internal/usecase/prepared_execution.go
Normal file
@@ -0,0 +1,245 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/validate"
|
||||
)
|
||||
|
||||
type preparedExecutionState uint8
|
||||
|
||||
const (
|
||||
preparedExecutionReady preparedExecutionState = iota
|
||||
preparedExecutionClaimed
|
||||
preparedExecutionDiscarded
|
||||
)
|
||||
|
||||
// PreparedExecution owns one frozen, single-use runner execution.
|
||||
type PreparedExecution struct {
|
||||
owner *Runner
|
||||
mu sync.Mutex
|
||||
state preparedExecutionState
|
||||
details *domain.PreparedRun
|
||||
payload *preparedExecutionPayload
|
||||
}
|
||||
|
||||
type preparedExecutionPayload struct {
|
||||
prepared *domain.PreparedRun
|
||||
validation validate.PreparedValidation
|
||||
directKey string
|
||||
}
|
||||
|
||||
// PrepareExecution completes preparation without generation or admission and
|
||||
// returns a runner-bound, single-use execution.
|
||||
func (r *Runner) PrepareExecution(ctx context.Context, req domain.RunRequest) (*PreparedExecution, error) {
|
||||
state, err := r.resolvePreparation(ctx, req, time.Now().UTC())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
operation, err := r.completePreparation(ctx, req, state)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
executionSnapshot, err := clonePreparedRun(operation.run)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to copy prepared execution: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
details, err := clonePreparedRun(executionSnapshot)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to copy prepared execution details: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
return &PreparedExecution{
|
||||
owner: r,
|
||||
state: preparedExecutionReady,
|
||||
details: details,
|
||||
payload: &preparedExecutionPayload{
|
||||
prepared: executionSnapshot,
|
||||
validation: operation.validation,
|
||||
directKey: state.effectiveModel.APIKey,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Details returns a fresh credential-redacted copy of the prepared run.
|
||||
func (p *PreparedExecution) Details() *domain.PreparedRun {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
p.mu.Lock()
|
||||
detailsSnapshot := p.details
|
||||
p.mu.Unlock()
|
||||
|
||||
details, err := clonePreparedRun(detailsSnapshot)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return details
|
||||
}
|
||||
|
||||
// Discard invalidates an unclaimed execution and drops its private payload.
|
||||
func (p *PreparedExecution) Discard() {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
|
||||
p.mu.Lock()
|
||||
if p.state != preparedExecutionReady {
|
||||
p.mu.Unlock()
|
||||
return
|
||||
}
|
||||
p.state = preparedExecutionDiscarded
|
||||
payload := p.payload
|
||||
p.payload = nil
|
||||
p.mu.Unlock()
|
||||
|
||||
payload.clear()
|
||||
}
|
||||
|
||||
// RunPrepared claims and executes one prepared execution owned by this runner.
|
||||
func (r *Runner) RunPrepared(ctx context.Context, prepared *PreparedExecution) (*domain.RunResult, error) {
|
||||
payload, err := prepared.claim(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer payload.clear()
|
||||
|
||||
runID, err := newRunID()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create run id: %w", err)
|
||||
}
|
||||
start := time.Now().UTC()
|
||||
|
||||
target := payload.prepared.EffectiveModelParams
|
||||
if err := validateAPIKey(target.APIKeyEnv, payload.directKey, target.APIKeyRequired); err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
release, err := r.admitRun(ctx, target.BackendID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
return r.executePreparedRun(ctx, payload.prepared, payload.directKey, runID, start, func(
|
||||
ctx context.Context,
|
||||
artifact *domain.Artifact,
|
||||
attemptsUsed int,
|
||||
) (domain.ValidationResult, error) {
|
||||
result, validationErr := payload.validation.Validate(ctx, artifact)
|
||||
if validationErr != nil {
|
||||
return domain.ValidationResult{}, validationErr
|
||||
}
|
||||
result.RepairAttempts = attemptsUsed
|
||||
return result, nil
|
||||
})
|
||||
}
|
||||
|
||||
func (p *PreparedExecution) claim(owner *Runner) (*preparedExecutionPayload, error) {
|
||||
if p == nil || owner == nil || p.owner != owner {
|
||||
return nil, fmt.Errorf("%w: prepared execution does not belong to this runner", ErrInvalidRequest)
|
||||
}
|
||||
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if p.state != preparedExecutionReady || p.payload == nil {
|
||||
return nil, fmt.Errorf("%w: prepared execution is not ready", ErrInvalidRequest)
|
||||
}
|
||||
p.state = preparedExecutionClaimed
|
||||
payload := p.payload
|
||||
p.payload = nil
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func (p *preparedExecutionPayload) clear() {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
if p.prepared != nil {
|
||||
p.prepared.EffectiveModelParams.APIKey = ""
|
||||
}
|
||||
p.prepared = nil
|
||||
p.validation = nil
|
||||
p.directKey = ""
|
||||
}
|
||||
|
||||
type noOpPreparedValidation struct {
|
||||
contract domain.OutputContract
|
||||
}
|
||||
|
||||
func (p noOpPreparedValidation) Validate(
|
||||
ctx context.Context,
|
||||
_ *domain.Artifact,
|
||||
) (domain.ValidationResult, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return domain.ValidationResult{}, err
|
||||
}
|
||||
return domain.ValidationResult{
|
||||
Status: domain.ValidationSkipped,
|
||||
Mode: p.contract.ValidationMode,
|
||||
SchemaPath: p.contract.SchemaPath,
|
||||
RepairAttempts: p.contract.RepairAttempts,
|
||||
IsValid: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (noOpPreparedValidation) SchemaDocument() any {
|
||||
return nil
|
||||
}
|
||||
|
||||
func clonePreparedRun(source *domain.PreparedRun) (*domain.PreparedRun, error) {
|
||||
if source == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
copied := *source
|
||||
extraParams, err := jsonvalue.CopyMap(source.EffectiveModelParams.ExtraParams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
copied.EffectiveModelParams.ExtraParams = extraParams
|
||||
|
||||
if source.InputHashes != nil {
|
||||
copied.InputHashes = make(map[string]string, len(source.InputHashes))
|
||||
for name, hash := range source.InputHashes {
|
||||
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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if source.StructuredOutput != nil {
|
||||
structuredOutput := *source.StructuredOutput
|
||||
copied.StructuredOutput = &structuredOutput
|
||||
if source.StructuredOutput.JSONSchema != nil {
|
||||
jsonSchema := *source.StructuredOutput.JSONSchema
|
||||
copied.StructuredOutput.JSONSchema = &jsonSchema
|
||||
schema, copyErr := cloneJSONValue(source.StructuredOutput.JSONSchema.Schema)
|
||||
if copyErr != nil {
|
||||
return nil, copyErr
|
||||
}
|
||||
copied.StructuredOutput.JSONSchema.Schema = schema
|
||||
}
|
||||
}
|
||||
return &copied, nil
|
||||
}
|
||||
|
||||
func cloneJSONValue(source any) (any, error) {
|
||||
return jsonvalue.Copy(source)
|
||||
}
|
||||
606
internal/usecase/prepared_execution_test.go
Normal file
606
internal/usecase/prepared_execution_test.go
Normal file
@@ -0,0 +1,606 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/validate"
|
||||
)
|
||||
|
||||
type recordingPreparedValidation struct {
|
||||
contract domain.OutputContract
|
||||
schemaDocument any
|
||||
results []domain.ValidationResult
|
||||
errs []error
|
||||
artifacts []string
|
||||
}
|
||||
|
||||
func (p *recordingPreparedValidation) Validate(
|
||||
_ context.Context,
|
||||
artifact *domain.Artifact,
|
||||
) (domain.ValidationResult, error) {
|
||||
p.artifacts = append(p.artifacts, string(artifact.Body))
|
||||
index := len(p.artifacts) - 1
|
||||
if index < len(p.errs) && p.errs[index] != nil {
|
||||
return domain.ValidationResult{}, p.errs[index]
|
||||
}
|
||||
if len(p.results) == 0 {
|
||||
return domain.ValidationResult{
|
||||
Status: domain.ValidationPassed,
|
||||
Mode: p.contract.ValidationMode,
|
||||
IsValid: true,
|
||||
}, nil
|
||||
}
|
||||
if index >= len(p.results) {
|
||||
index = len(p.results) - 1
|
||||
}
|
||||
return p.results[index], nil
|
||||
}
|
||||
|
||||
func (p *recordingPreparedValidation) SchemaDocument() any {
|
||||
return p.schemaDocument
|
||||
}
|
||||
|
||||
type recordingValidationPreparer struct {
|
||||
plan *recordingPreparedValidation
|
||||
prepareErr error
|
||||
prepareCalls int
|
||||
directValidateCalls int
|
||||
}
|
||||
|
||||
type validationOnly struct{}
|
||||
|
||||
func (validationOnly) Validate(context.Context, *domain.Artifact, domain.OutputContract) (domain.ValidationResult, error) {
|
||||
return domain.ValidationResult{}, nil
|
||||
}
|
||||
|
||||
func (v *recordingValidationPreparer) Validate(
|
||||
context.Context,
|
||||
*domain.Artifact,
|
||||
domain.OutputContract,
|
||||
) (domain.ValidationResult, error) {
|
||||
v.directValidateCalls++
|
||||
return domain.ValidationResult{}, errors.New("live validation must not be used")
|
||||
}
|
||||
|
||||
func (v *recordingValidationPreparer) PrepareValidation(
|
||||
_ context.Context,
|
||||
contract domain.OutputContract,
|
||||
) (validate.PreparedValidation, error) {
|
||||
v.prepareCalls++
|
||||
if v.prepareErr != nil {
|
||||
return nil, v.prepareErr
|
||||
}
|
||||
v.plan.contract = contract
|
||||
return v.plan, nil
|
||||
}
|
||||
|
||||
func TestRunnerPrepareExecutionCompletesWithoutAdmissionOrGeneration(t *testing.T) {
|
||||
schemaDocument := map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"value": map[string]any{"type": "string"},
|
||||
"": map[string]any{"type": "boolean"},
|
||||
},
|
||||
}
|
||||
def := promptDef(domain.FormatJSON, domain.ValidationJSONSchema, 0)
|
||||
def.Validation.SchemaPath = "schema.json"
|
||||
reader := defaultArtifactReader()
|
||||
renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{
|
||||
SessionID: "prepared-session",
|
||||
Messages: []domain.RenderedMessage{{
|
||||
Role: "user",
|
||||
Content: "original message",
|
||||
}},
|
||||
}}
|
||||
llmClient := &fakeLLM{forbid: true}
|
||||
validator := &recordingValidationPreparer{
|
||||
plan: &recordingPreparedValidation{schemaDocument: schemaDocument},
|
||||
}
|
||||
admitter := &fakeRunAdmitter{}
|
||||
profile := defaultExecutionProfile()
|
||||
profile.ExtraParams = map[string]any{
|
||||
"metadata": map[string]any{"source": "original"},
|
||||
}
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: def},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": profile}},
|
||||
nil,
|
||||
reader,
|
||||
renderer,
|
||||
llmClient,
|
||||
validator,
|
||||
admitter,
|
||||
)
|
||||
|
||||
prepared, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
APIKey: "direct-test-key",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
defer prepared.Discard()
|
||||
|
||||
if validator.prepareCalls != 1 || validator.directValidateCalls != 0 {
|
||||
t.Fatalf(
|
||||
"validation calls=(prepare=%d direct=%d), want (1, 0)",
|
||||
validator.prepareCalls,
|
||||
validator.directValidateCalls,
|
||||
)
|
||||
}
|
||||
if reader.calls != 1 || renderer.calls != 1 {
|
||||
t.Fatalf("completion calls=(artifact=%d render=%d), want (1, 1)", reader.calls, renderer.calls)
|
||||
}
|
||||
if len(admitter.backendIDs) != 0 || llmClient.calls != 0 {
|
||||
t.Fatalf("prepare invoked execution collaborators: admission=%v generation=%d", admitter.backendIDs, llmClient.calls)
|
||||
}
|
||||
|
||||
first := prepared.Details()
|
||||
if first == nil {
|
||||
t.Fatal("prepared details are nil")
|
||||
}
|
||||
if first.EffectiveModelParams.APIKey != "" {
|
||||
t.Fatal("prepared details retained the direct API key")
|
||||
}
|
||||
if first.StructuredOutput == nil ||
|
||||
first.StructuredOutput.JSONSchema == nil ||
|
||||
!reflect.DeepEqual(first.StructuredOutput.JSONSchema.Schema, schemaDocument) {
|
||||
t.Fatalf("prepared details have unexpected structured output: %#v", first.StructuredOutput)
|
||||
}
|
||||
|
||||
first.Messages[0].Content = "caller mutation"
|
||||
first.InputHashes["input"] = "caller mutation"
|
||||
first.EffectiveModelParams.ExtraParams["metadata"].(map[string]any)["source"] = "caller mutation"
|
||||
first.StructuredOutput.JSONSchema.Schema.(map[string]any)["type"] = "string"
|
||||
renderer.rendered.Messages[0].Content = "source mutation"
|
||||
|
||||
second := prepared.Details()
|
||||
if second.Messages[0].Content != "original message" ||
|
||||
second.InputHashes["input"] == "caller mutation" ||
|
||||
second.EffectiveModelParams.ExtraParams["metadata"].(map[string]any)["source"] != "original" ||
|
||||
second.StructuredOutput.JSONSchema.Schema.(map[string]any)["type"] != "object" {
|
||||
t.Fatalf("details did not preserve an independent snapshot: %#v", second)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerPreparationRejectsExcessivelyDeepPreparedSchema(t *testing.T) {
|
||||
operations := []struct {
|
||||
name string
|
||||
run func(*Runner, domain.RunRequest) error
|
||||
}{
|
||||
{
|
||||
name: "Prepare",
|
||||
run: func(runner *Runner, request domain.RunRequest) error {
|
||||
_, err := runner.Prepare(context.Background(), request)
|
||||
return err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "Run",
|
||||
run: func(runner *Runner, request domain.RunRequest) error {
|
||||
_, err := runner.Run(context.Background(), request)
|
||||
return err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "PrepareExecution",
|
||||
run: func(runner *Runner, request domain.RunRequest) error {
|
||||
_, err := runner.PrepareExecution(context.Background(), request)
|
||||
return err
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, operation := range operations {
|
||||
t.Run(operation.name, func(t *testing.T) {
|
||||
def := promptDef(domain.FormatJSON, domain.ValidationJSONSchema, 0)
|
||||
def.Validation.SchemaPath = "schema.json"
|
||||
llmClient := &fakeLLM{forbid: true}
|
||||
validator := &recordingValidationPreparer{
|
||||
plan: &recordingPreparedValidation{schemaDocument: excessivelyDeepPreparedJSONValue()},
|
||||
}
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: def},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
validator,
|
||||
nil,
|
||||
)
|
||||
|
||||
err := operation.run(runner, domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if !errors.Is(err, ErrValidation) {
|
||||
t.Fatalf("expected ErrValidation, got %v", err)
|
||||
}
|
||||
if llmClient.calls != 0 {
|
||||
t.Fatalf("invalid prepared schema reached generation: %d calls", llmClient.calls)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func excessivelyDeepPreparedJSONValue() any {
|
||||
const clearlyUnsafeContainerDepth = 1_000
|
||||
var value any = true
|
||||
for level := 0; level < clearlyUnsafeContainerDepth; level++ {
|
||||
value = map[string]any{"child": value}
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func TestRunnerRunPreparedRechecksEnvironmentCredentialBeforeAdmission(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)
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunPreparedKeepsDirectCredentialOutOfMetadata(t *testing.T) {
|
||||
const directKey = "direct-prepared-test-key"
|
||||
|
||||
profile := defaultExecutionProfile()
|
||||
profile.APIKeyRequired = true
|
||||
client := &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(),
|
||||
client,
|
||||
&recordingValidationPreparer{plan: &recordingPreparedValidation{}},
|
||||
nil,
|
||||
)
|
||||
|
||||
prepared, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
APIKey: directKey,
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
if details := prepared.Details(); details.EffectiveModelParams.APIKey != "" {
|
||||
t.Fatal("prepared details retained direct credential")
|
||||
}
|
||||
|
||||
result, err := runner.RunPrepared(context.Background(), prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("run prepared: %v", err)
|
||||
}
|
||||
if client.lastReq.Target.APIKey != directKey {
|
||||
t.Fatal("generation did not receive direct credential")
|
||||
}
|
||||
if result.EffectiveModelParams.APIKey != "" {
|
||||
t.Fatal("run result retained direct credential")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunPreparedUsesFrozenValidationForInitialAndRepairOutputs(t *testing.T) {
|
||||
validator := &recordingValidationPreparer{
|
||||
plan: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{
|
||||
{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
},
|
||||
{
|
||||
Status: domain.ValidationPassed,
|
||||
Mode: domain.ValidationJSON,
|
||||
IsValid: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
repairer := &fakeRepairer{
|
||||
responses: []*domain.GenerateResponse{{
|
||||
Content: `{"repaired":true}`,
|
||||
Usage: domain.TokenUsage{
|
||||
PromptTokens: 2, CompletionTokens: 3, TotalTokens: 5,
|
||||
CachedTokens: 7, CacheWriteTokens: 11,
|
||||
},
|
||||
}},
|
||||
}
|
||||
admitter := &fakeRunAdmitter{}
|
||||
reader := defaultArtifactReader()
|
||||
renderer := defaultRenderer()
|
||||
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,
|
||||
},
|
||||
}},
|
||||
validator,
|
||||
repairer,
|
||||
admitter,
|
||||
)
|
||||
prepared, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
|
||||
result, err := runner.RunPrepared(context.Background(), prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("run prepared: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(validator.plan.artifacts, []string{`{"broken":true}`, `{"repaired":true}`}) {
|
||||
t.Fatalf("prepared validation artifacts=%#v", validator.plan.artifacts)
|
||||
}
|
||||
if validator.directValidateCalls != 0 || repairer.calls != 1 {
|
||||
t.Fatalf("validation/repair calls=(direct=%d repair=%d), want (0, 1)", validator.directValidateCalls, repairer.calls)
|
||||
}
|
||||
if result.Validation.Status != domain.ValidationPassed || result.Validation.RepairAttempts != 1 {
|
||||
t.Fatalf("unexpected repaired validation result: %+v", result.Validation)
|
||||
}
|
||||
wantUsage := domain.TokenUsage{
|
||||
PromptTokens: 15, CompletionTokens: 20, TotalTokens: 24,
|
||||
CachedTokens: 30, CacheWriteTokens: 40,
|
||||
}
|
||||
if result.Usage != wantUsage {
|
||||
t.Fatalf("prepared cumulative usage = %+v, want %+v", result.Usage, wantUsage)
|
||||
}
|
||||
if admitter.releaseCalls != 1 {
|
||||
t.Fatalf("admission releases=%d, want 1", admitter.releaseCalls)
|
||||
}
|
||||
if reader.calls != 1 || renderer.calls != 1 {
|
||||
t.Fatalf("execution reopened preparation sources: artifact=%d render=%d", reader.calls, renderer.calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunPreparedReleasesAdmissionAcrossExecutionErrors(t *testing.T) {
|
||||
generationFailure := errors.New("generation failed")
|
||||
validationFailure := errors.New("validation failed")
|
||||
repairFailure := errors.New("repair failed")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
generationErr error
|
||||
validation *recordingPreparedValidation
|
||||
repairer *fakeRepairer
|
||||
wantError error
|
||||
}{
|
||||
{
|
||||
name: "generation failure",
|
||||
generationErr: generationFailure,
|
||||
validation: &recordingPreparedValidation{},
|
||||
wantError: ErrLLMGenerate,
|
||||
},
|
||||
{
|
||||
name: "validation failure",
|
||||
validation: &recordingPreparedValidation{
|
||||
errs: []error{validationFailure},
|
||||
},
|
||||
wantError: ErrValidation,
|
||||
},
|
||||
{
|
||||
name: "repair failure",
|
||||
validation: &recordingPreparedValidation{
|
||||
results: []domain.ValidationResult{{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSON,
|
||||
Errors: []string{"invalid"},
|
||||
IsValid: false,
|
||||
}},
|
||||
},
|
||||
repairer: &fakeRepairer{err: repairFailure},
|
||||
wantError: ErrValidation,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
def := promptDef(domain.FormatJSON, domain.ValidationJSON, 1)
|
||||
if test.repairer == nil {
|
||||
def.Validation.RepairAttempts = 0
|
||||
}
|
||||
validator := &recordingValidationPreparer{plan: test.validation}
|
||||
admitter := &fakeRunAdmitter{}
|
||||
runner := NewRunnerWithRepairer(
|
||||
&fakePromptRepo{def: def},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
&fakeLLM{
|
||||
resp: &domain.GenerateResponse{Content: `{"value":true}`},
|
||||
err: test.generationErr,
|
||||
},
|
||||
validator,
|
||||
test.repairer,
|
||||
admitter,
|
||||
)
|
||||
prepared, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
|
||||
result, err := runner.RunPrepared(context.Background(), prepared)
|
||||
if result != nil || !errors.Is(err, test.wantError) {
|
||||
t.Fatalf("run prepared=(%+v, %v), want %v", result, err, test.wantError)
|
||||
}
|
||||
if len(admitter.backendIDs) != 1 || admitter.releaseCalls != 1 {
|
||||
t.Fatalf(
|
||||
"admission calls=%#v releases=%d, want one each",
|
||||
admitter.backendIDs,
|
||||
admitter.releaseCalls,
|
||||
)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerPreparedExecutionOwnershipUseAndDiscard(t *testing.T) {
|
||||
newRunner := func() *Runner {
|
||||
return NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
&fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}},
|
||||
&recordingValidationPreparer{plan: &recordingPreparedValidation{}},
|
||||
nil,
|
||||
)
|
||||
}
|
||||
request := domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
}
|
||||
|
||||
owner := newRunner()
|
||||
prepared, err := owner.PrepareExecution(context.Background(), request)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
if _, err := newRunner().RunPrepared(context.Background(), prepared); !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("foreign runner error=%v, want ErrInvalidRequest", err)
|
||||
}
|
||||
if _, err := owner.RunPrepared(context.Background(), prepared); err != nil {
|
||||
t.Fatalf("owner run prepared: %v", err)
|
||||
}
|
||||
if _, err := owner.RunPrepared(context.Background(), prepared); !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("second owner run error=%v, want ErrInvalidRequest", err)
|
||||
}
|
||||
if prepared.Details() == nil {
|
||||
t.Fatal("details unavailable after execution")
|
||||
}
|
||||
|
||||
discarded, err := owner.PrepareExecution(context.Background(), request)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare discarded execution: %v", err)
|
||||
}
|
||||
discarded.Discard()
|
||||
discarded.Discard()
|
||||
if _, err := owner.RunPrepared(context.Background(), discarded); !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("discarded execution error=%v, want ErrInvalidRequest", err)
|
||||
}
|
||||
if discarded.Details() == nil {
|
||||
t.Fatal("details unavailable after discard")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerPreparedExecutionWithoutValidatorSkipsValidation(t *testing.T) {
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
&fakeLLM{resp: &domain.GenerateResponse{Content: `{}`}},
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
prepared, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
|
||||
result, err := runner.RunPrepared(context.Background(), prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("run prepared: %v", err)
|
||||
}
|
||||
if result.Validation.Status != domain.ValidationSkipped || !result.Validation.IsValid {
|
||||
t.Fatalf("unexpected no-validator result: %+v", result.Validation)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerPrepareExecutionRequiresValidationPreparer(t *testing.T) {
|
||||
reader := defaultArtifactReader()
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationBasic, 0)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
reader,
|
||||
defaultRenderer(),
|
||||
&fakeLLM{forbid: true},
|
||||
validationOnly{},
|
||||
nil,
|
||||
)
|
||||
|
||||
prepared, err := runner.PrepareExecution(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if prepared != nil || !errors.Is(err, ErrValidation) {
|
||||
t.Fatalf("prepare execution=(%+v, %v), want ErrValidation", prepared, err)
|
||||
}
|
||||
if reader.calls != 0 {
|
||||
t.Fatalf("unsupported validator allowed completion, artifact calls=%d", reader.calls)
|
||||
}
|
||||
}
|
||||
106
internal/usecase/profile_inspection.go
Normal file
106
internal/usecase/profile_inspection.go
Normal file
@@ -0,0 +1,106 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
type resolvedProfileSelection struct {
|
||||
id string
|
||||
profile *domain.ExecutionProfile
|
||||
backend *domain.Backend
|
||||
}
|
||||
|
||||
func (r *Runner) resolveProfileSelection(
|
||||
ctx context.Context,
|
||||
profileID string,
|
||||
) (*resolvedProfileSelection, error) {
|
||||
normalizedID := strings.TrimSpace(profileID)
|
||||
if normalizedID == "" {
|
||||
return nil, fmt.Errorf("%w: profile id is required", ErrInvalidRequest)
|
||||
}
|
||||
if r == nil || r.profiles == nil {
|
||||
return nil, fmt.Errorf("%w: profile repository is not configured", ErrProfileLoad)
|
||||
}
|
||||
|
||||
selectedProfile, err := r.profiles.GetProfile(ctx, normalizedID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
|
||||
}
|
||||
if selectedProfile == nil {
|
||||
return nil, fmt.Errorf("%w: profile repository returned nil profile", ErrProfileLoad)
|
||||
}
|
||||
|
||||
profileValue := *selectedProfile
|
||||
profileValue.BackendID = strings.TrimSpace(profileValue.BackendID)
|
||||
|
||||
var selectedBackend *domain.Backend
|
||||
if profileValue.BackendID != "" {
|
||||
if r.backends == nil {
|
||||
return nil, fmt.Errorf("%w: backend %q cannot be resolved", ErrProfileLoad, profileValue.BackendID)
|
||||
}
|
||||
backendValue, err := r.backends.GetBackend(profileValue.BackendID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: backend %q: %w", ErrProfileLoad, profileValue.BackendID, err)
|
||||
}
|
||||
selectedBackend = &backendValue
|
||||
}
|
||||
|
||||
return &resolvedProfileSelection{
|
||||
id: normalizedID,
|
||||
profile: &profileValue,
|
||||
backend: selectedBackend,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeResolvedExecutionTarget(target domain.ExecutionTarget) (domain.ExecutionTarget, error) {
|
||||
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(target.Endpoint)
|
||||
if err != nil {
|
||||
return domain.ExecutionTarget{}, fmt.Errorf("execution endpoint: %w", err)
|
||||
}
|
||||
target.Endpoint = endpoint
|
||||
if strings.TrimSpace(target.Model) == "" {
|
||||
return domain.ExecutionTarget{}, errors.New("execution model is required")
|
||||
}
|
||||
if err := domain.ValidateExecutionTargetSettings(target); err != nil {
|
||||
return domain.ExecutionTarget{}, err
|
||||
}
|
||||
return target, nil
|
||||
}
|
||||
|
||||
// InspectProfile resolves one explicit profile without prompt or execution work.
|
||||
func (r *Runner) InspectProfile(
|
||||
ctx context.Context,
|
||||
profileID string,
|
||||
) (*domain.ProfileInspection, error) {
|
||||
normalizedID := strings.TrimSpace(profileID)
|
||||
if normalizedID == "" {
|
||||
return nil, fmt.Errorf("%w: profile id is required", ErrInvalidRequest)
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, ctx.Err())
|
||||
default:
|
||||
}
|
||||
|
||||
selection, err := r.resolveProfileSelection(ctx, normalizedID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
target, _ := resolveExecutionTarget(selection.backend, selection.profile, nil)
|
||||
target, err = normalizeResolvedExecutionTarget(target)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
|
||||
}
|
||||
target.APIKey = ""
|
||||
|
||||
return &domain.ProfileInspection{
|
||||
ProfileID: selection.id,
|
||||
EffectiveModelParams: target,
|
||||
APIKeyRequired: target.APIKeyRequired,
|
||||
}, nil
|
||||
}
|
||||
213
internal/usecase/profile_inspection_test.go
Normal file
213
internal/usecase/profile_inspection_test.go
Normal file
@@ -0,0 +1,213 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/defaults"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
)
|
||||
|
||||
type inspectionProfileRepository struct {
|
||||
profile *domain.ExecutionProfile
|
||||
err error
|
||||
calls int
|
||||
id string
|
||||
}
|
||||
|
||||
func (r *inspectionProfileRepository) GetProfile(
|
||||
_ context.Context,
|
||||
id string,
|
||||
) (*domain.ExecutionProfile, error) {
|
||||
r.calls++
|
||||
r.id = id
|
||||
if r.err != nil {
|
||||
return nil, r.err
|
||||
}
|
||||
return r.profile, nil
|
||||
}
|
||||
|
||||
type inspectionBackendResolver struct {
|
||||
backend domain.Backend
|
||||
err error
|
||||
calls int
|
||||
id string
|
||||
}
|
||||
|
||||
func (r *inspectionBackendResolver) GetBackend(id string) (domain.Backend, error) {
|
||||
r.calls++
|
||||
r.id = id
|
||||
if r.err != nil {
|
||||
return domain.Backend{}, r.err
|
||||
}
|
||||
return r.backend, nil
|
||||
}
|
||||
|
||||
func TestRunnerInspectProfileResolvesProfileAndBackendOnce(t *testing.T) {
|
||||
profiles := &inspectionProfileRepository{profile: &domain.ExecutionProfile{
|
||||
ID: "profile",
|
||||
BackendID: " backend ",
|
||||
Model: "profile-model",
|
||||
Temperature: 0.4,
|
||||
MaxTokens: 32,
|
||||
TimeoutSeconds: 45,
|
||||
ServiceTier: "priority",
|
||||
ReasoningEffort: "high",
|
||||
ExtraParams: map[string]any{
|
||||
"profile": "value",
|
||||
},
|
||||
}}
|
||||
backends := &inspectionBackendResolver{backend: domain.Backend{
|
||||
ID: "backend",
|
||||
Endpoint: "https://backend.example/v1",
|
||||
APIKeyEnv: "BACKEND_KEY",
|
||||
ExtraParams: map[string]any{
|
||||
"backend": "value",
|
||||
},
|
||||
}}
|
||||
runner := &Runner{profiles: profiles, backends: backends}
|
||||
|
||||
inspection, err := runner.InspectProfile(context.Background(), " profile ")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect profile: %v", err)
|
||||
}
|
||||
if profiles.calls != 1 || profiles.id != "profile" {
|
||||
t.Fatalf("profile lookup=(calls=%d id=%q), want one exact lookup", profiles.calls, profiles.id)
|
||||
}
|
||||
if backends.calls != 1 || backends.id != "backend" {
|
||||
t.Fatalf("backend lookup=(calls=%d id=%q), want one exact lookup", backends.calls, backends.id)
|
||||
}
|
||||
if profiles.profile.BackendID != " backend " {
|
||||
t.Fatalf("inspection mutated repository profile backend: %q", profiles.profile.BackendID)
|
||||
}
|
||||
|
||||
wantTarget := domain.ExecutionTarget{
|
||||
BackendID: "backend",
|
||||
Endpoint: "https://backend.example/v1",
|
||||
Model: "profile-model",
|
||||
Temperature: 0.4,
|
||||
MaxTokens: 32,
|
||||
TopP: defaults.ExecutionTargetDefault().TopP,
|
||||
TimeoutSeconds: 45,
|
||||
ServiceTier: "priority",
|
||||
ReasoningEffort: "high",
|
||||
APIKeyEnv: "BACKEND_KEY",
|
||||
ExtraParams: map[string]any{
|
||||
"profile": "value",
|
||||
},
|
||||
}
|
||||
if inspection.ProfileID != "profile" || inspection.APIKeyRequired ||
|
||||
!reflect.DeepEqual(inspection.EffectiveModelParams, wantTarget) {
|
||||
t.Fatalf("inspection=%#v, want profile=%q target=%#v", inspection, "profile", wantTarget)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerInspectProfileDoesNotNeedExecutionCollaboratorsOrCredentials(t *testing.T) {
|
||||
t.Setenv("PROMPTKIT_INSPECTION_TEST_KEY", "")
|
||||
profiles := &inspectionProfileRepository{profile: &domain.ExecutionProfile{
|
||||
ID: "endpoint-only",
|
||||
Endpoint: "https://profile.example/v1",
|
||||
Model: "profile-model",
|
||||
APIKeyEnv: "PROMPTKIT_INSPECTION_TEST_KEY",
|
||||
}}
|
||||
runner := &Runner{profiles: profiles}
|
||||
|
||||
inspection, err := runner.InspectProfile(context.Background(), "endpoint-only")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect endpoint-only profile: %v", err)
|
||||
}
|
||||
if inspection.EffectiveModelParams.BackendID != "" ||
|
||||
inspection.EffectiveModelParams.APIKeyEnv != "PROMPTKIT_INSPECTION_TEST_KEY" ||
|
||||
inspection.APIKeyRequired {
|
||||
t.Fatalf("unexpected endpoint-only inspection: %#v", inspection)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerInspectProfileDirectCredentialRequirementClearsBackendEnvironment(t *testing.T) {
|
||||
profiles := &inspectionProfileRepository{profile: &domain.ExecutionProfile{
|
||||
ID: "direct-key",
|
||||
BackendID: "backend",
|
||||
Model: "profile-model",
|
||||
APIKeyRequired: true,
|
||||
}}
|
||||
backends := &inspectionBackendResolver{backend: domain.Backend{
|
||||
ID: "backend",
|
||||
Endpoint: "https://backend.example/v1",
|
||||
APIKeyEnv: "BACKEND_KEY",
|
||||
}}
|
||||
|
||||
inspection, err := (&Runner{profiles: profiles, backends: backends}).InspectProfile(
|
||||
context.Background(),
|
||||
"direct-key",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("inspect direct-key profile: %v", err)
|
||||
}
|
||||
if !inspection.APIKeyRequired || inspection.EffectiveModelParams.APIKeyEnv != "" {
|
||||
t.Fatalf("credential requirement was not resolved exclusively: %#v", inspection)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerInspectProfileClassifiesFailuresWithoutRepositoryWorkAfterCancellation(t *testing.T) {
|
||||
t.Run("blank ID", func(t *testing.T) {
|
||||
profiles := &inspectionProfileRepository{}
|
||||
_, err := (&Runner{profiles: profiles}).InspectProfile(context.Background(), " \t ")
|
||||
if !errors.Is(err, ErrInvalidRequest) || profiles.calls != 0 {
|
||||
t.Fatalf("blank inspection=(%v, calls=%d), want invalid request without lookup", err, profiles.calls)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("canceled context", func(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
profiles := &inspectionProfileRepository{}
|
||||
_, err := (&Runner{profiles: profiles}).InspectProfile(ctx, "profile")
|
||||
if !errors.Is(err, ErrProfileLoad) || !errors.Is(err, context.Canceled) || profiles.calls != 0 {
|
||||
t.Fatalf("canceled inspection=(%v, calls=%d), want profile load and context identities without lookup", err, profiles.calls)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing profile", func(t *testing.T) {
|
||||
profiles := &inspectionProfileRepository{err: profile.ErrProfileNotFound}
|
||||
_, err := (&Runner{profiles: profiles}).InspectProfile(context.Background(), "missing")
|
||||
if !errors.Is(err, ErrProfileLoad) || !errors.Is(err, profile.ErrProfileNotFound) {
|
||||
t.Fatalf("missing profile error=%v, want profile load and not-found identities", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unknown backend", func(t *testing.T) {
|
||||
backendErr := errors.New("unknown backend")
|
||||
profiles := &inspectionProfileRepository{profile: &domain.ExecutionProfile{
|
||||
ID: "profile", BackendID: "backend", Model: "profile-model",
|
||||
}}
|
||||
backends := &inspectionBackendResolver{err: backendErr}
|
||||
_, err := (&Runner{profiles: profiles, backends: backends}).InspectProfile(context.Background(), "profile")
|
||||
if !errors.Is(err, ErrProfileLoad) || !errors.Is(err, backendErr) {
|
||||
t.Fatalf("unknown backend error=%v, want profile load and backend identities", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("defensive invalid dependencies", func(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
runner *Runner
|
||||
}{
|
||||
{name: "nil repository", runner: &Runner{}},
|
||||
{name: "nil profile", runner: &Runner{profiles: &inspectionProfileRepository{}}},
|
||||
{name: "invalid target", runner: &Runner{profiles: &inspectionProfileRepository{
|
||||
profile: &domain.ExecutionProfile{ID: "profile", Endpoint: "https://profile.example/v1"},
|
||||
}}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := tc.runner.InspectProfile(context.Background(), "profile")
|
||||
if !errors.Is(err, ErrProfileLoad) {
|
||||
t.Fatalf("inspection error=%v, want profile load", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
73
internal/usecase/prompt_inspection.go
Normal file
73
internal/usecase/prompt_inspection.go
Normal file
@@ -0,0 +1,73 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
type resolvedPromptDefinition struct {
|
||||
definition *domain.PromptDefinition
|
||||
hash string
|
||||
}
|
||||
|
||||
func (r *Runner) resolvePromptDefinition(
|
||||
ctx context.Context,
|
||||
promptID string,
|
||||
promptVersion string,
|
||||
) (*resolvedPromptDefinition, error) {
|
||||
if strings.TrimSpace(promptID) == "" {
|
||||
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest)
|
||||
}
|
||||
if r == nil || r.promptDefs == nil {
|
||||
return nil, fmt.Errorf("%w: prompt repository is not configured", ErrPromptLoad)
|
||||
}
|
||||
|
||||
definition, err := r.promptDefs.GetPromptDefinition(ctx, promptID, promptVersion)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrPromptLoad, err)
|
||||
}
|
||||
if definition == nil {
|
||||
return nil, fmt.Errorf("%w: prompt repository returned nil definition", ErrPromptLoad)
|
||||
}
|
||||
|
||||
hash, err := hashPromptDefinition(definition)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to hash prompt definition: %v", ErrPromptLoad, err)
|
||||
}
|
||||
return &resolvedPromptDefinition{definition: definition, hash: hash}, nil
|
||||
}
|
||||
|
||||
// InspectPrompt resolves one explicit prompt without execution work.
|
||||
func (r *Runner) InspectPrompt(
|
||||
ctx context.Context,
|
||||
promptID string,
|
||||
promptVersion string,
|
||||
) (*domain.PromptInspection, error) {
|
||||
if strings.TrimSpace(promptID) == "" {
|
||||
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest)
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, fmt.Errorf("%w: %w", ErrPromptLoad, ctx.Err())
|
||||
default:
|
||||
}
|
||||
|
||||
selection, err := r.resolvePromptDefinition(ctx, promptID, promptVersion)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inputs := make([]domain.PromptInput, len(selection.definition.Inputs))
|
||||
copy(inputs, selection.definition.Inputs)
|
||||
|
||||
return &domain.PromptInspection{
|
||||
PromptID: selection.definition.ID,
|
||||
PromptVersion: selection.definition.Version,
|
||||
PromptHash: selection.hash,
|
||||
DefaultProfileID: selection.definition.DefaultProfile,
|
||||
Inputs: inputs,
|
||||
OutputContract: selection.definition.Validation,
|
||||
}, nil
|
||||
}
|
||||
157
internal/usecase/prompt_inspection_test.go
Normal file
157
internal/usecase/prompt_inspection_test.go
Normal file
@@ -0,0 +1,157 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/promptdef"
|
||||
)
|
||||
|
||||
type inspectionPromptRepository struct {
|
||||
definition *domain.PromptDefinition
|
||||
err error
|
||||
calls int
|
||||
id string
|
||||
version string
|
||||
}
|
||||
|
||||
func (r *inspectionPromptRepository) GetPromptDefinition(
|
||||
_ context.Context,
|
||||
id string,
|
||||
version string,
|
||||
) (*domain.PromptDefinition, error) {
|
||||
r.calls++
|
||||
r.id = id
|
||||
r.version = version
|
||||
if r.err != nil {
|
||||
return nil, r.err
|
||||
}
|
||||
return r.definition, nil
|
||||
}
|
||||
|
||||
func TestRunnerInspectPromptResolvesOneDefinitionWithoutExecutionCollaborators(t *testing.T) {
|
||||
definition := &domain.PromptDefinition{
|
||||
ID: "normalized.prompt",
|
||||
Version: "1.2.3",
|
||||
DefaultProfile: "not-resolved",
|
||||
Inputs: []domain.PromptInput{
|
||||
{Name: "document", Required: true, ContentType: "text/plain", Description: "Source document."},
|
||||
{Name: "audience", ContentType: "text/plain", Description: "Intended reader."},
|
||||
},
|
||||
Validation: domain.OutputContract{
|
||||
Format: domain.FormatJSON,
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schemas/result.json",
|
||||
RepairAttempts: 2,
|
||||
},
|
||||
}
|
||||
repository := &inspectionPromptRepository{definition: definition}
|
||||
runner := &Runner{promptDefs: repository}
|
||||
|
||||
inspection, err := runner.InspectPrompt(context.Background(), " prompt-id ", " version ")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect prompt: %v", err)
|
||||
}
|
||||
wantHash, err := hashPromptDefinition(definition)
|
||||
if err != nil {
|
||||
t.Fatalf("hash prompt definition: %v", err)
|
||||
}
|
||||
if repository.calls != 1 || repository.id != " prompt-id " || repository.version != " version " {
|
||||
t.Fatalf("prompt lookup=(calls=%d id=%q version=%q), want one unchanged lookup", repository.calls, repository.id, repository.version)
|
||||
}
|
||||
if inspection.PromptID != definition.ID ||
|
||||
inspection.PromptVersion != definition.Version ||
|
||||
inspection.PromptHash != wantHash ||
|
||||
inspection.DefaultProfileID != definition.DefaultProfile ||
|
||||
!reflect.DeepEqual(inspection.Inputs, definition.Inputs) ||
|
||||
inspection.OutputContract != definition.Validation {
|
||||
t.Fatalf("inspection=%#v, want definition metadata", inspection)
|
||||
}
|
||||
|
||||
inspection.Inputs[0].Name = "changed"
|
||||
second, err := runner.InspectPrompt(context.Background(), " prompt-id ", " version ")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect prompt again: %v", err)
|
||||
}
|
||||
if definition.Inputs[0].Name != "document" || second.Inputs[0].Name != "document" {
|
||||
t.Fatalf("inspection input mutation escaped caller result: definition=%#v next=%#v", definition.Inputs, second.Inputs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerInspectPromptClassifiesFailuresWithoutRepositoryWorkAfterCancellation(t *testing.T) {
|
||||
t.Run("blank ID", func(t *testing.T) {
|
||||
repository := &inspectionPromptRepository{}
|
||||
_, err := (&Runner{promptDefs: repository}).InspectPrompt(context.Background(), " \t ", "version")
|
||||
if !errors.Is(err, ErrInvalidRequest) || repository.calls != 0 {
|
||||
t.Fatalf("blank inspection=(%v, calls=%d), want invalid request without lookup", err, repository.calls)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("canceled context", func(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
repository := &inspectionPromptRepository{}
|
||||
_, err := (&Runner{promptDefs: repository}).InspectPrompt(ctx, "prompt", "version")
|
||||
if !errors.Is(err, ErrPromptLoad) || !errors.Is(err, context.Canceled) || repository.calls != 0 {
|
||||
t.Fatalf("canceled inspection=(%v, calls=%d), want prompt load and context identities without lookup", err, repository.calls)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing prompt", func(t *testing.T) {
|
||||
repository := &inspectionPromptRepository{err: promptdef.ErrPromptDefinitionNotFound}
|
||||
_, err := (&Runner{promptDefs: repository}).InspectPrompt(context.Background(), "missing", "version")
|
||||
if !errors.Is(err, ErrPromptLoad) || !errors.Is(err, promptdef.ErrPromptDefinitionNotFound) {
|
||||
t.Fatalf("missing prompt error=%v, want prompt load and not-found identities", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("defensive prompt dependencies", func(t *testing.T) {
|
||||
var nilRunner *Runner
|
||||
cases := []struct {
|
||||
name string
|
||||
runner *Runner
|
||||
}{
|
||||
{name: "nil runner", runner: nilRunner},
|
||||
{name: "nil repository", runner: &Runner{}},
|
||||
{name: "nil definition", runner: &Runner{promptDefs: &inspectionPromptRepository{}}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := tc.runner.InspectPrompt(context.Background(), "prompt", "version")
|
||||
if !errors.Is(err, ErrPromptLoad) {
|
||||
t.Fatalf("inspection error=%v, want prompt load", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestRunnerPrepareUsesThePromptInspectionSelectionAndHash(t *testing.T) {
|
||||
definition := promptDef(domain.FormatMarkdown, domain.ValidationBasic, 0)
|
||||
repository := &fakePromptRepo{def: definition}
|
||||
runner := NewRunner(
|
||||
repository,
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
&fakeArtifactReader{},
|
||||
&fakeRenderer{rendered: &domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hello"}}}},
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
|
||||
inspection, err := runner.InspectPrompt(context.Background(), definition.ID, definition.Version)
|
||||
if err != nil {
|
||||
t.Fatalf("inspect prompt: %v", err)
|
||||
}
|
||||
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{PromptID: definition.ID, PromptVersion: definition.Version, ProfileID: "exec"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare prompt: %v", err)
|
||||
}
|
||||
if inspection.PromptHash != prepared.PromptHash {
|
||||
t.Fatalf("inspection hash=%q, preparation hash=%q", inspection.PromptHash, prepared.PromptHash)
|
||||
}
|
||||
}
|
||||
@@ -19,6 +19,7 @@ type RepairRequest struct {
|
||||
ValidationErrors []string
|
||||
SessionID string
|
||||
Target domain.ExecutionTarget
|
||||
TargetPresence domain.ExecutionTargetPresence
|
||||
StructuredOutput *domain.StructuredOutputSpec
|
||||
Attempt int
|
||||
MaxAttempts int
|
||||
@@ -44,7 +45,6 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
|
||||
}
|
||||
|
||||
prompt := domain.RenderedPrompt{
|
||||
SessionID: req.SessionID,
|
||||
Messages: []domain.RenderedMessage{
|
||||
{
|
||||
Role: "system",
|
||||
@@ -64,11 +64,13 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
|
||||
},
|
||||
}
|
||||
|
||||
resp, err := r.llm.Generate(ctx, domain.GenerateRequest{
|
||||
Prompt: prompt,
|
||||
Target: req.Target,
|
||||
StructuredOutput: req.StructuredOutput,
|
||||
})
|
||||
resp, err := r.llm.Generate(ctx, newGenerationRequest(
|
||||
prompt,
|
||||
req.SessionID,
|
||||
req.Target,
|
||||
req.TargetPresence,
|
||||
req.StructuredOutput,
|
||||
))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/capacity"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/defaults"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/prompt"
|
||||
@@ -71,6 +72,11 @@ type preparationState struct {
|
||||
start time.Time
|
||||
}
|
||||
|
||||
type preparedOperation struct {
|
||||
run *domain.PreparedRun
|
||||
validation validate.PreparedValidation
|
||||
}
|
||||
|
||||
func NewRunner(
|
||||
promptDefs promptdef.Repository,
|
||||
profiles profile.Repository,
|
||||
@@ -131,41 +137,62 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if r.admitter != nil {
|
||||
release, admitErr := r.admitter.Admit(ctx, state.effectiveModel.BackendID)
|
||||
if admitErr != nil {
|
||||
if errors.Is(admitErr, capacity.ErrCapacityExceeded) {
|
||||
return nil, fmt.Errorf(
|
||||
"backend %q admission: %w",
|
||||
state.effectiveModel.BackendID,
|
||||
admitErr,
|
||||
)
|
||||
}
|
||||
return nil, admitErr
|
||||
}
|
||||
defer release()
|
||||
}
|
||||
|
||||
prepared, err := r.completePreparation(ctx, req, state)
|
||||
release, err := r.admitRun(ctx, state.effectiveModel.BackendID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{SessionID: prepared.SessionID, Messages: prepared.Messages},
|
||||
Target: prepared.EffectiveModelParams,
|
||||
TargetPresence: prepared.TargetPresence,
|
||||
StructuredOutput: prepared.StructuredOutput,
|
||||
operation, err := r.completePreparation(ctx, req, state)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
directAPIKey := state.effectiveModel.APIKey
|
||||
|
||||
return r.executePreparedRun(ctx, operation.run, directAPIKey, runID, start, func(
|
||||
ctx context.Context,
|
||||
artifact *domain.Artifact,
|
||||
attemptsUsed int,
|
||||
) (domain.ValidationResult, error) {
|
||||
result, err := operation.validation.Validate(ctx, artifact)
|
||||
result.RepairAttempts = attemptsUsed
|
||||
return result, err
|
||||
})
|
||||
}
|
||||
|
||||
type preparedValidationFunc func(
|
||||
context.Context,
|
||||
*domain.Artifact,
|
||||
int,
|
||||
) (domain.ValidationResult, error)
|
||||
|
||||
func (r *Runner) executePreparedRun(
|
||||
ctx context.Context,
|
||||
prepared *domain.PreparedRun,
|
||||
directAPIKey string,
|
||||
runID string,
|
||||
start time.Time,
|
||||
validateArtifact preparedValidationFunc,
|
||||
) (*domain.RunResult, error) {
|
||||
executionTarget := prepared.EffectiveModelParams
|
||||
executionTarget.APIKey = directAPIKey
|
||||
genResp, err := r.llm.Generate(ctx, newGenerationRequest(
|
||||
domain.RenderedPrompt{Messages: prepared.Messages},
|
||||
prepared.SessionID,
|
||||
executionTarget,
|
||||
prepared.TargetPresence,
|
||||
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)
|
||||
}
|
||||
usage := genResp.Usage
|
||||
|
||||
outputArtifact := buildOutputArtifact(genResp.Content, prepared.OutputContract.Format)
|
||||
validationResult, err := r.validateOutput(ctx, &outputArtifact, prepared.OutputContract, 0)
|
||||
validationResult, err := validateArtifact(ctx, &outputArtifact, 0)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
|
||||
}
|
||||
@@ -179,7 +206,8 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
|
||||
PreviousOutput: genResp.Content,
|
||||
ValidationErrors: validationResult.Errors,
|
||||
SessionID: prepared.SessionID,
|
||||
Target: prepared.EffectiveModelParams,
|
||||
Target: executionTarget,
|
||||
TargetPresence: prepared.TargetPresence,
|
||||
StructuredOutput: prepared.StructuredOutput,
|
||||
Attempt: attemptsUsed,
|
||||
MaxAttempts: prepared.OutputContract.RepairAttempts,
|
||||
@@ -193,9 +221,10 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
|
||||
}
|
||||
|
||||
genResp = repairResp
|
||||
usage = addTokenUsage(usage, repairResp.Usage)
|
||||
outputArtifact = buildOutputArtifact(genResp.Content, prepared.OutputContract.Format)
|
||||
|
||||
validationResult, err = r.validateOutput(ctx, &outputArtifact, prepared.OutputContract, attemptsUsed)
|
||||
validationResult, err = validateArtifact(ctx, &outputArtifact, attemptsUsed)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
|
||||
}
|
||||
@@ -203,6 +232,7 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
|
||||
}
|
||||
|
||||
end := time.Now().UTC()
|
||||
executionTarget.APIKey = ""
|
||||
|
||||
return &domain.RunResult{
|
||||
RunID: runID,
|
||||
@@ -218,21 +248,35 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
|
||||
SelectedBackendID: prepared.SelectedBackendID,
|
||||
ModelName: prepared.EffectiveModelParams.Model,
|
||||
Endpoint: prepared.EffectiveModelParams.Endpoint,
|
||||
EffectiveModelParams: prepared.EffectiveModelParams,
|
||||
EffectiveModelParams: executionTarget,
|
||||
InputHashes: prepared.InputHashes,
|
||||
Usage: genResp.Usage,
|
||||
Usage: usage,
|
||||
StartTime: start,
|
||||
EndTime: end,
|
||||
Duration: end.Sub(start),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func addTokenUsage(total, next domain.TokenUsage) domain.TokenUsage {
|
||||
return domain.TokenUsage{
|
||||
PromptTokens: total.PromptTokens + next.PromptTokens,
|
||||
CompletionTokens: total.CompletionTokens + next.CompletionTokens,
|
||||
TotalTokens: total.TotalTokens + next.TotalTokens,
|
||||
CachedTokens: total.CachedTokens + next.CachedTokens,
|
||||
CacheWriteTokens: total.CacheWriteTokens + next.CacheWriteTokens,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.PreparedRun, error) {
|
||||
state, err := r.resolvePreparation(ctx, req, time.Now().UTC())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return r.completePreparation(ctx, req, state)
|
||||
operation, err := r.completePreparation(ctx, req, state)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return operation.run, nil
|
||||
}
|
||||
|
||||
func (r *Runner) resolvePreparation(
|
||||
@@ -248,13 +292,15 @@ func (r *Runner) resolvePreparation(
|
||||
return nil, fmt.Errorf("%w: session_id: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
def, err := r.promptDefs.GetPromptDefinition(ctx, req.PromptID, req.PromptVersion)
|
||||
promptSelection, err := r.resolvePromptDefinition(ctx, req.PromptID, req.PromptVersion)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrPromptLoad, err)
|
||||
return nil, err
|
||||
}
|
||||
promptDefinitionHash, err := hashPromptDefinition(def)
|
||||
def := promptSelection.definition
|
||||
promptDefinitionHash := promptSelection.hash
|
||||
effectiveContract, err := resolveOutputContract(def, req.Validation)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to hash prompt definition: %v", ErrPromptLoad, err)
|
||||
return nil, fmt.Errorf("%w: output contract: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
selectedProfileID := strings.TrimSpace(req.ProfileID)
|
||||
@@ -265,45 +311,26 @@ func (r *Runner) resolvePreparation(
|
||||
return nil, fmt.Errorf("%w: %w: profile id is required either in request or prompt default_profile", ErrInvalidRequest, ErrProfileRequired)
|
||||
}
|
||||
|
||||
execProfile, err := r.profiles.GetProfile(ctx, selectedProfileID)
|
||||
selection, err := r.resolveProfileSelection(ctx, selectedProfileID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var selectedBackend *domain.Backend
|
||||
if backendID := strings.TrimSpace(execProfile.BackendID); backendID != "" {
|
||||
execProfile.BackendID = backendID
|
||||
if r.backends == nil {
|
||||
return nil, fmt.Errorf("%w: backend %q cannot be resolved", ErrProfileLoad, backendID)
|
||||
}
|
||||
resolvedBackend, resolveErr := r.backends.GetBackend(backendID)
|
||||
if resolveErr != nil {
|
||||
return nil, fmt.Errorf("%w: backend %q: %w", ErrProfileLoad, backendID, resolveErr)
|
||||
}
|
||||
selectedBackend = &resolvedBackend
|
||||
}
|
||||
|
||||
effectiveModel, targetPresence, err := resolveExecutionTarget(selectedBackend, execProfile, req.Execution)
|
||||
effectiveModel, targetPresence := resolveExecutionTarget(selection.backend, selection.profile, req.Execution)
|
||||
effectiveModel.APIKey = req.APIKey
|
||||
effectiveModel, err = normalizeResolvedExecutionTarget(effectiveModel)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||
}
|
||||
effectiveModel.APIKey = req.APIKey
|
||||
if strings.TrimSpace(effectiveModel.Endpoint) == "" {
|
||||
return nil, fmt.Errorf("%w: execution endpoint is required", ErrInvalidRequest)
|
||||
}
|
||||
if strings.TrimSpace(effectiveModel.Model) == "" {
|
||||
return nil, fmt.Errorf("%w: execution model is required", ErrInvalidRequest)
|
||||
}
|
||||
if err := validateAPIKey(effectiveModel.APIKeyEnv, effectiveModel.APIKey, effectiveModel.APIKeyRequired); err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
effectiveContract := resolveOutputContract(def, req.Validation)
|
||||
return &preparationState{
|
||||
definition: def,
|
||||
directSessionID: directSessionID,
|
||||
promptDefinitionHash: promptDefinitionHash,
|
||||
selectedProfileID: selectedProfileID,
|
||||
selectedProfileID: selection.id,
|
||||
effectiveModel: effectiveModel,
|
||||
targetPresence: targetPresence,
|
||||
effectiveContract: effectiveContract,
|
||||
@@ -315,16 +342,76 @@ func (r *Runner) completePreparation(
|
||||
ctx context.Context,
|
||||
req domain.RunRequest,
|
||||
state *preparationState,
|
||||
) (*domain.PreparedRun, error) {
|
||||
structuredOutput, err := r.resolveStructuredOutput(
|
||||
ctx,
|
||||
) (*preparedOperation, error) {
|
||||
validationPlan, err := r.prepareValidation(ctx, state.effectiveContract)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
structuredOutput, err := r.structuredOutputFromValidationPlan(
|
||||
state.definition,
|
||||
state.effectiveContract,
|
||||
validationPlan,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
prepared, err := r.completePreparationWithStructuredOutput(ctx, req, state, structuredOutput)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &preparedOperation{run: prepared, validation: validationPlan}, nil
|
||||
}
|
||||
|
||||
func (r *Runner) prepareValidation(
|
||||
ctx context.Context,
|
||||
contract domain.OutputContract,
|
||||
) (validate.PreparedValidation, error) {
|
||||
if r.validator == nil {
|
||||
return noOpPreparedValidation{contract: contract}, nil
|
||||
}
|
||||
preparer, ok := r.validator.(validate.ValidationPreparer)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("%w: validator does not support prepared validation", ErrValidation)
|
||||
}
|
||||
plan, err := preparer.PrepareValidation(ctx, contract)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
|
||||
}
|
||||
if plan == nil {
|
||||
return nil, fmt.Errorf("%w: validator returned nil prepared validation", ErrValidation)
|
||||
}
|
||||
return plan, nil
|
||||
}
|
||||
|
||||
func (r *Runner) structuredOutputFromValidationPlan(
|
||||
def *domain.PromptDefinition,
|
||||
contract domain.OutputContract,
|
||||
plan validate.PreparedValidation,
|
||||
) (*domain.StructuredOutputSpec, error) {
|
||||
if contract.ValidationMode != domain.ValidationJSONSchema {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
schemaDocument := plan.SchemaDocument()
|
||||
if schemaDocument == nil {
|
||||
if r.validator == nil {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, fmt.Errorf("%w: prepared json_schema validation has no schema document", ErrValidation)
|
||||
}
|
||||
schemaDocument, err := jsonvalue.Copy(schemaDocument)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: invalid prepared json_schema schema document: %v", ErrValidation, err)
|
||||
}
|
||||
return structuredOutputSpec(def, schemaDocument), nil
|
||||
}
|
||||
|
||||
func (r *Runner) completePreparationWithStructuredOutput(
|
||||
ctx context.Context,
|
||||
req domain.RunRequest,
|
||||
state *preparationState,
|
||||
structuredOutput *domain.StructuredOutputSpec,
|
||||
) (*domain.PreparedRun, error) {
|
||||
resolvedInputs := make(map[string]*domain.Artifact, len(req.Inputs))
|
||||
inputHashes := make(map[string]string, len(req.Inputs))
|
||||
for name, ref := range req.Inputs {
|
||||
@@ -354,13 +441,15 @@ func (r *Runner) completePreparation(
|
||||
}
|
||||
|
||||
end := time.Now().UTC()
|
||||
effectiveModel := state.effectiveModel
|
||||
effectiveModel.APIKey = ""
|
||||
return &domain.PreparedRun{
|
||||
PromptID: state.definition.ID,
|
||||
PromptVersion: state.definition.Version,
|
||||
PromptHash: state.promptDefinitionHash,
|
||||
SelectedProfileID: state.selectedProfileID,
|
||||
SelectedBackendID: state.effectiveModel.BackendID,
|
||||
EffectiveModelParams: state.effectiveModel,
|
||||
EffectiveModelParams: effectiveModel,
|
||||
TargetPresence: state.targetPresence,
|
||||
OutputContract: state.effectiveContract,
|
||||
StructuredOutput: structuredOutput,
|
||||
@@ -374,29 +463,29 @@ func (r *Runner) completePreparation(
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *Runner) resolveStructuredOutput(ctx context.Context, def *domain.PromptDefinition, contract domain.OutputContract) (*domain.StructuredOutputSpec, error) {
|
||||
if contract.ValidationMode != domain.ValidationJSONSchema {
|
||||
return nil, nil
|
||||
func (r *Runner) admitRun(ctx context.Context, backendID string) (func(), error) {
|
||||
if r.admitter == nil {
|
||||
return func() {}, nil
|
||||
}
|
||||
|
||||
loader, ok := r.validator.(validate.SchemaDocumentLoader)
|
||||
if !ok || loader == nil {
|
||||
return nil, fmt.Errorf("%w: json_schema output requires schema document loader", ErrValidation)
|
||||
}
|
||||
|
||||
schemaDoc, err := loader.LoadSchemaDocument(ctx, contract.SchemaPath)
|
||||
release, err := r.admitter.Admit(ctx, backendID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: failed to load json schema for structured output: %v", ErrValidation, err)
|
||||
if errors.Is(err, capacity.ErrCapacityExceeded) && strings.TrimSpace(backendID) != "" {
|
||||
return nil, &CapacityError{BackendID: backendID}
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return release, nil
|
||||
}
|
||||
|
||||
func structuredOutputSpec(def *domain.PromptDefinition, schemaDocument any) *domain.StructuredOutputSpec {
|
||||
return &domain.StructuredOutputSpec{
|
||||
Type: domain.StructuredOutputJSONSchema,
|
||||
JSONSchema: &domain.StructuredOutputJSONSpec{
|
||||
Name: deriveStructuredSchemaName(def.ID, def.Version),
|
||||
Strict: true,
|
||||
Schema: schemaDoc,
|
||||
Schema: schemaDocument,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func deriveStructuredSchemaName(promptID string, promptVersion string) string {
|
||||
@@ -425,25 +514,6 @@ func deriveStructuredSchemaName(promptID string, promptVersion string) string {
|
||||
return name
|
||||
}
|
||||
|
||||
func (r *Runner) validateOutput(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract, attemptsUsed int) (domain.ValidationResult, error) {
|
||||
if r.validator == nil || contract.ValidationMode == domain.ValidationNone {
|
||||
return domain.ValidationResult{
|
||||
Status: domain.ValidationSkipped,
|
||||
Mode: contract.ValidationMode,
|
||||
SchemaPath: contract.SchemaPath,
|
||||
RepairAttempts: attemptsUsed,
|
||||
IsValid: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
res, err := r.validator.Validate(ctx, artifact, contract)
|
||||
if err != nil {
|
||||
return domain.ValidationResult{}, err
|
||||
}
|
||||
res.RepairAttempts = attemptsUsed
|
||||
return res, nil
|
||||
}
|
||||
|
||||
func (r *Runner) shouldAttemptRepair(contract domain.OutputContract, validationResult domain.ValidationResult) bool {
|
||||
if r.repairer == nil {
|
||||
return false
|
||||
@@ -499,7 +569,7 @@ func mergeExecutionTarget(base domain.ExecutionTarget, override domain.Execution
|
||||
return out
|
||||
}
|
||||
|
||||
func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence, error) {
|
||||
func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence) {
|
||||
out := base
|
||||
var presence domain.ExecutionTargetPresence
|
||||
if override.Endpoint != "" {
|
||||
@@ -509,30 +579,18 @@ func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.E
|
||||
out.Model = override.Model
|
||||
}
|
||||
if override.Temperature != nil {
|
||||
if *override.Temperature < 0 || *override.Temperature > 2 {
|
||||
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("temperature must be between 0 and 2")
|
||||
}
|
||||
out.Temperature = *override.Temperature
|
||||
presence.Temperature = true
|
||||
}
|
||||
if override.MaxTokens != nil {
|
||||
if *override.MaxTokens < 0 {
|
||||
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("max_tokens must be greater than or equal to 0")
|
||||
}
|
||||
out.MaxTokens = *override.MaxTokens
|
||||
presence.MaxTokens = true
|
||||
}
|
||||
if override.TopP != nil {
|
||||
if *override.TopP < 0 || *override.TopP > 1 {
|
||||
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("top_p must be between 0 and 1")
|
||||
}
|
||||
out.TopP = *override.TopP
|
||||
presence.TopP = true
|
||||
}
|
||||
if override.TimeoutSeconds != nil {
|
||||
if *override.TimeoutSeconds < 0 {
|
||||
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("timeout_seconds must be greater than or equal to 0")
|
||||
}
|
||||
out.TimeoutSeconds = *override.TimeoutSeconds
|
||||
presence.TimeoutSeconds = true
|
||||
}
|
||||
@@ -548,22 +606,18 @@ func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.E
|
||||
if len(override.ExtraParams) > 0 {
|
||||
out.ExtraParams = copyExtraParams(override.ExtraParams)
|
||||
}
|
||||
return out, presence, nil
|
||||
return out, presence
|
||||
}
|
||||
|
||||
func resolveExecutionTarget(backendValue *domain.Backend, profileValue *domain.ExecutionProfile, override *domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence, error) {
|
||||
func resolveExecutionTarget(backendValue *domain.Backend, profileValue *domain.ExecutionProfile, override *domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence) {
|
||||
out := defaults.ExecutionTargetDefault()
|
||||
out = mergeExecutionTarget(out, backendToTarget(backendValue))
|
||||
out = mergeExecutionTarget(out, executionProfileToTarget(profileValue))
|
||||
var presence domain.ExecutionTargetPresence
|
||||
if override != nil {
|
||||
var err error
|
||||
out, presence, err = mergeExecutionTargetOverride(out, *override)
|
||||
if err != nil {
|
||||
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, err
|
||||
}
|
||||
out, presence = mergeExecutionTargetOverride(out, *override)
|
||||
}
|
||||
return out, presence, nil
|
||||
return out, presence
|
||||
}
|
||||
|
||||
func validateAPIKey(apiKeyEnv string, apiKey string, apiKeyRequired bool) error {
|
||||
@@ -630,18 +684,21 @@ func copyExtraParams(src map[string]any) map[string]any {
|
||||
return cp
|
||||
}
|
||||
|
||||
func resolveOutputContract(def *domain.PromptDefinition, override *domain.OutputContract) domain.OutputContract {
|
||||
func resolveOutputContract(def *domain.PromptDefinition, override *domain.OutputContract) (domain.OutputContract, error) {
|
||||
contract := def.Validation
|
||||
if contract.Format == "" {
|
||||
contract.Format = def.OutputFormat
|
||||
}
|
||||
if override != nil {
|
||||
contract = *override
|
||||
if contract.Format == "" {
|
||||
contract.Format = domain.FormatText
|
||||
}
|
||||
}
|
||||
if contract.Format == "" {
|
||||
contract.Format = domain.FormatText
|
||||
if err := domain.ValidateOutputContract(contract); err != nil {
|
||||
return domain.OutputContract{}, err
|
||||
}
|
||||
return contract
|
||||
return contract, nil
|
||||
}
|
||||
|
||||
func hashRenderedPrompt(p domain.RenderedPrompt) string {
|
||||
|
||||
@@ -6,9 +6,11 @@ import (
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -189,16 +191,34 @@ func (f *fakeValidator) Validate(ctx context.Context, artifact *domain.Artifact,
|
||||
return f.result, nil
|
||||
}
|
||||
|
||||
func (f *fakeValidator) LoadSchemaDocument(ctx context.Context, schemaPath string) (any, error) {
|
||||
f.schemaLoads++
|
||||
f.schemaLoadPath = schemaPath
|
||||
if f.schemaErr != nil {
|
||||
return nil, f.schemaErr
|
||||
func (f *fakeValidator) PrepareValidation(_ context.Context, contract domain.OutputContract) (validate.PreparedValidation, error) {
|
||||
var schemaDocument any
|
||||
if contract.ValidationMode == domain.ValidationJSONSchema {
|
||||
f.schemaLoads++
|
||||
f.schemaLoadPath = contract.SchemaPath
|
||||
if f.schemaErr != nil {
|
||||
return nil, f.schemaErr
|
||||
}
|
||||
schemaDocument = f.schemaDoc
|
||||
if schemaDocument == nil {
|
||||
schemaDocument = map[string]any{"type": "object"}
|
||||
}
|
||||
}
|
||||
if f.schemaDoc != nil {
|
||||
return f.schemaDoc, nil
|
||||
}
|
||||
return map[string]any{"type": "object"}, nil
|
||||
return &fakePreparedValidator{validator: f, contract: contract, schemaDocument: schemaDocument}, nil
|
||||
}
|
||||
|
||||
type fakePreparedValidator struct {
|
||||
validator *fakeValidator
|
||||
contract domain.OutputContract
|
||||
schemaDocument any
|
||||
}
|
||||
|
||||
func (p *fakePreparedValidator) Validate(ctx context.Context, artifact *domain.Artifact) (domain.ValidationResult, error) {
|
||||
return p.validator.Validate(ctx, artifact, p.contract)
|
||||
}
|
||||
|
||||
func (p *fakePreparedValidator) SchemaDocument() any {
|
||||
return p.schemaDocument
|
||||
}
|
||||
|
||||
type fakeRepairer struct {
|
||||
@@ -208,6 +228,30 @@ type fakeRepairer struct {
|
||||
reqs []RepairRequest
|
||||
}
|
||||
|
||||
type sequenceLLM struct {
|
||||
responses []*domain.GenerateResponse
|
||||
requests []domain.GenerateRequest
|
||||
}
|
||||
|
||||
func (c *sequenceLLM) Generate(_ context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) {
|
||||
c.requests = append(c.requests, req)
|
||||
index := len(c.requests) - 1
|
||||
if index >= len(c.responses) {
|
||||
return nil, errors.New("no generation response configured")
|
||||
}
|
||||
return c.responses[index], nil
|
||||
}
|
||||
|
||||
type recordingRepairer struct {
|
||||
next OutputRepairer
|
||||
reqs []RepairRequest
|
||||
}
|
||||
|
||||
func (r *recordingRepairer) Repair(ctx context.Context, req RepairRequest) (*domain.GenerateResponse, error) {
|
||||
r.reqs = append(r.reqs, req)
|
||||
return r.next.Repair(ctx, req)
|
||||
}
|
||||
|
||||
type fakeRunAdmitter struct {
|
||||
backendIDs []string
|
||||
err error
|
||||
@@ -502,6 +546,35 @@ func TestRunnerDirectSessionResolution(t *testing.T) {
|
||||
t.Fatalf("invalid direct session invoked generation %d times", llmClient.calls)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("malformed direct value fails before loading or generation", func(t *testing.T) {
|
||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "unexpected"}}
|
||||
runner := NewRunner(
|
||||
promptRepo,
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
nil, nil)
|
||||
|
||||
_, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
SessionID: "session" + string([]byte{0xff}),
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
if promptRepo.lastID != "" {
|
||||
t.Fatalf("invalid direct session loaded prompt %q", promptRepo.lastID)
|
||||
}
|
||||
if llmClient.calls != 0 {
|
||||
t.Fatalf("invalid direct session invoked generation %d times", llmClient.calls)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestRunnerPrepareUsesPromptDefaultProfileWhenNoExplicitProfileID(t *testing.T) {
|
||||
@@ -704,16 +777,23 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRunnerPrepareInvalidRequestNumericOverridesFail(t *testing.T) {
|
||||
tests := []struct {
|
||||
type testCase struct {
|
||||
name string
|
||||
override *domain.ExecutionTargetOverride
|
||||
}{
|
||||
}
|
||||
tests := []testCase{
|
||||
{name: "temperature below range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(-0.1)}},
|
||||
{name: "temperature above range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(2.1)}},
|
||||
{name: "max tokens below range", override: &domain.ExecutionTargetOverride{MaxTokens: intPtr(-1)}},
|
||||
{name: "top p below range", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(-0.1)}},
|
||||
{name: "top p above range", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(1.1)}},
|
||||
{name: "timeout below range", override: &domain.ExecutionTargetOverride{TimeoutSeconds: intPtr(-1)}},
|
||||
{name: "temperature is not finite", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(math.NaN())}},
|
||||
{name: "top p is not finite", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(math.Inf(1))}},
|
||||
}
|
||||
if strconv.IntSize == 64 {
|
||||
durationLimit := int64(math.MaxInt64 / int64(time.Second))
|
||||
tests = append(tests, testCase{name: "timeout cannot be represented as a duration", override: &domain.ExecutionTargetOverride{TimeoutSeconds: intPtr(int(durationLimit) + 1)}})
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
@@ -1357,14 +1437,14 @@ func TestRunnerAdmissionUsesResolvedBackendIdentity(t *testing.T) {
|
||||
|
||||
func TestRunnerAdmissionFailureSkipsCompletionCollaborators(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
admissionError error
|
||||
wantBackendContext bool
|
||||
name string
|
||||
admissionError error
|
||||
wantCapacityType bool
|
||||
}{
|
||||
{
|
||||
name: "capacity exhausted",
|
||||
admissionError: capacity.ErrCapacityExceeded,
|
||||
wantBackendContext: true,
|
||||
name: "capacity exhausted",
|
||||
admissionError: capacity.ErrCapacityExceeded,
|
||||
wantCapacityType: true,
|
||||
},
|
||||
{
|
||||
name: "context canceled",
|
||||
@@ -1414,8 +1494,16 @@ func TestRunnerAdmissionFailureSkipsCompletionCollaborators(t *testing.T) {
|
||||
if errors.Is(err, ErrInvalidRequest) || errors.Is(err, ErrLLMGenerate) {
|
||||
t.Fatalf("admission error was recategorized: %v", err)
|
||||
}
|
||||
if tc.wantBackendContext && !strings.Contains(err.Error(), "custom") {
|
||||
t.Fatalf("capacity error lacks backend context: %v", err)
|
||||
var capacityErr *CapacityError
|
||||
if tc.wantCapacityType {
|
||||
if !errors.As(err, &capacityErr) {
|
||||
t.Fatalf("capacity error=%v, want internal typed identity", err)
|
||||
}
|
||||
if capacityErr.BackendID != "custom" {
|
||||
t.Fatalf("capacity backend ID=%q, want custom", capacityErr.BackendID)
|
||||
}
|
||||
} else if errors.As(err, &capacityErr) {
|
||||
t.Fatalf("non-capacity admission error exposed typed capacity identity: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(admitter.backendIDs, []string{"custom"}) {
|
||||
t.Fatalf("admitted backend IDs=%#v, want custom", admitter.backendIDs)
|
||||
@@ -1682,36 +1770,6 @@ func TestRunnerRunSelectedProfileBeatsBuiltInDefault(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunBuiltInDefaultsUsedWhenProfileOmitsOptionalFields(t *testing.T) {
|
||||
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"},
|
||||
}}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}
|
||||
runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil)
|
||||
|
||||
res, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
if res.EffectiveModelParams.Temperature != defaults.ExecutionDefaultTemperature {
|
||||
t.Fatalf("expected default temperature %v, got %v", defaults.ExecutionDefaultTemperature, res.EffectiveModelParams.Temperature)
|
||||
}
|
||||
if res.EffectiveModelParams.TopP != defaults.ExecutionDefaultTopP {
|
||||
t.Fatalf("expected default top_p %v, got %v", defaults.ExecutionDefaultTopP, res.EffectiveModelParams.TopP)
|
||||
}
|
||||
if res.EffectiveModelParams.MaxTokens != defaults.ExecutionDefaultMaxTokens {
|
||||
t.Fatalf("expected default max_tokens %d, got %d", defaults.ExecutionDefaultMaxTokens, res.EffectiveModelParams.MaxTokens)
|
||||
}
|
||||
if res.EffectiveModelParams.TimeoutSeconds != defaults.ExecutionDefaultTimeoutSeconds {
|
||||
t.Fatalf("expected default timeout_seconds %d, got %d", defaults.ExecutionDefaultTimeoutSeconds, res.EffectiveModelParams.TimeoutSeconds)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunAPIKeyEnvResolvesFromEnvironment(t *testing.T) {
|
||||
t.Setenv("PROMPTKIT_TEST_API_KEY", "secret")
|
||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
|
||||
@@ -1757,7 +1815,7 @@ func TestRunnerRunDirectAPIKeyBypassesMissingEnvAndReachesLLM(t *testing.T) {
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}
|
||||
runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil)
|
||||
|
||||
_, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
result, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
APIKey: directKey,
|
||||
@@ -1769,6 +1827,9 @@ func TestRunnerRunDirectAPIKeyBypassesMissingEnvAndReachesLLM(t *testing.T) {
|
||||
if llmClient.lastReq.Target.APIKey != directKey {
|
||||
t.Fatalf("expected direct API key to reach LLM request")
|
||||
}
|
||||
if result.EffectiveModelParams.APIKey != "" {
|
||||
t.Fatal("run result retained direct API key")
|
||||
}
|
||||
if llmClient.lastReq.Target.APIKeyEnv != "PROMPTKIT_MISSING_KEY" {
|
||||
t.Fatalf("expected api_key_env name to remain on target, got %q", llmClient.lastReq.Target.APIKeyEnv)
|
||||
}
|
||||
@@ -2033,50 +2094,255 @@ func TestRunnerRunValidationStillWorks(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunStructuredRepairRemainsBoundedAndUsesEffectiveModelSettings(t *testing.T) {
|
||||
repairer := &fakeRepairer{responses: []*domain.GenerateResponse{{Content: `{"broken":`}, {Content: `{"still":`}}}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: `{"initial":`}}
|
||||
func TestRunnerRepairStateMachine(t *testing.T) {
|
||||
failed := func(mode domain.ValidationMode, diagnostic string) domain.ValidationResult {
|
||||
return domain.ValidationResult{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: mode,
|
||||
Errors: []string{diagnostic},
|
||||
IsValid: false,
|
||||
}
|
||||
}
|
||||
passed := func(mode domain.ValidationMode) domain.ValidationResult {
|
||||
return domain.ValidationResult{
|
||||
Status: domain.ValidationPassed,
|
||||
Mode: mode,
|
||||
IsValid: true,
|
||||
}
|
||||
}
|
||||
responses := func(count int) []*domain.GenerateResponse {
|
||||
values := make([]*domain.GenerateResponse, count)
|
||||
for index := range values {
|
||||
unit := index + 1
|
||||
values[index] = &domain.GenerateResponse{
|
||||
Content: fmt.Sprintf(`{"candidate":%d}`, index),
|
||||
Usage: domain.TokenUsage{
|
||||
PromptTokens: unit,
|
||||
CompletionTokens: unit * 10,
|
||||
TotalTokens: unit * 100,
|
||||
CachedTokens: unit * 1000,
|
||||
CacheWriteTokens: unit * 10000,
|
||||
},
|
||||
}
|
||||
}
|
||||
return values
|
||||
}
|
||||
zeroOverrides := &domain.ExecutionTargetOverride{
|
||||
Temperature: float64Ptr(0),
|
||||
MaxTokens: intPtr(0),
|
||||
TopP: float64Ptr(0),
|
||||
TimeoutSeconds: intPtr(0),
|
||||
}
|
||||
allPresent := domain.ExecutionTargetPresence{
|
||||
Temperature: true,
|
||||
MaxTokens: true,
|
||||
TopP: true,
|
||||
TimeoutSeconds: true,
|
||||
}
|
||||
|
||||
runner := NewRunnerWithRepairer(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 1)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
|
||||
"exec": {ID: "exec", BackendID: "custom", Model: "profile-model", TimeoutSeconds: 55},
|
||||
}}, fakeBackendResolver{backends: map[string]domain.Backend{
|
||||
"custom": {ID: "custom", Endpoint: "http://backend/v1"},
|
||||
}},
|
||||
tests := []struct {
|
||||
name string
|
||||
mode domain.ValidationMode
|
||||
budget int
|
||||
validationResults []domain.ValidationResult
|
||||
responses []*domain.GenerateResponse
|
||||
execution *domain.ExecutionTargetOverride
|
||||
wantPresence domain.ExecutionTargetPresence
|
||||
wantRepairs int
|
||||
wantStatus domain.ValidationStatus
|
||||
structured bool
|
||||
}{
|
||||
{
|
||||
name: "initial success does not repair",
|
||||
mode: domain.ValidationJSON,
|
||||
budget: 3,
|
||||
validationResults: []domain.ValidationResult{passed(domain.ValidationJSON)},
|
||||
responses: responses(1),
|
||||
wantStatus: domain.ValidationPassed,
|
||||
},
|
||||
{
|
||||
name: "basic failure is ineligible despite budget",
|
||||
mode: domain.ValidationBasic,
|
||||
budget: 3,
|
||||
validationResults: []domain.ValidationResult{failed(domain.ValidationBasic, "empty output")},
|
||||
responses: responses(1),
|
||||
wantStatus: domain.ValidationFailed,
|
||||
},
|
||||
{
|
||||
name: "inherited numeric values remain absent",
|
||||
mode: domain.ValidationJSON,
|
||||
budget: 1,
|
||||
validationResults: []domain.ValidationResult{
|
||||
failed(domain.ValidationJSON, "initial syntax"),
|
||||
passed(domain.ValidationJSON),
|
||||
},
|
||||
responses: responses(2),
|
||||
wantRepairs: 1,
|
||||
wantStatus: domain.ValidationPassed,
|
||||
},
|
||||
{
|
||||
name: "explicit numeric zeros remain present",
|
||||
mode: domain.ValidationJSONSchema,
|
||||
budget: 1,
|
||||
execution: zeroOverrides,
|
||||
structured: true,
|
||||
validationResults: []domain.ValidationResult{
|
||||
failed(domain.ValidationJSONSchema, "initial schema mismatch"),
|
||||
passed(domain.ValidationJSONSchema),
|
||||
},
|
||||
responses: responses(2),
|
||||
wantPresence: allPresent,
|
||||
wantRepairs: 1,
|
||||
wantStatus: domain.ValidationPassed,
|
||||
},
|
||||
{
|
||||
name: "successful repair stops below larger budget",
|
||||
mode: domain.ValidationJSON,
|
||||
budget: 4,
|
||||
validationResults: []domain.ValidationResult{
|
||||
failed(domain.ValidationJSON, "candidate zero"),
|
||||
failed(domain.ValidationJSON, "candidate one"),
|
||||
passed(domain.ValidationJSON),
|
||||
},
|
||||
responses: responses(3),
|
||||
wantRepairs: 2,
|
||||
wantStatus: domain.ValidationPassed,
|
||||
},
|
||||
{
|
||||
name: "failed repairs exhaust exact larger budget",
|
||||
mode: domain.ValidationJSON,
|
||||
budget: 3,
|
||||
validationResults: []domain.ValidationResult{
|
||||
failed(domain.ValidationJSON, "candidate zero"),
|
||||
failed(domain.ValidationJSON, "candidate one"),
|
||||
failed(domain.ValidationJSON, "candidate two"),
|
||||
failed(domain.ValidationJSON, "candidate three"),
|
||||
},
|
||||
responses: responses(4),
|
||||
wantRepairs: 3,
|
||||
wantStatus: domain.ValidationFailed,
|
||||
},
|
||||
}
|
||||
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
validate.NewStandardValidator("."),
|
||||
repairer, nil)
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
definition := promptDef(domain.FormatJSON, tc.mode, tc.budget)
|
||||
plan := &recordingPreparedValidation{results: tc.validationResults}
|
||||
if tc.structured {
|
||||
definition.Validation.SchemaPath = "schema.json"
|
||||
plan.schemaDocument = map[string]any{"type": "object"}
|
||||
}
|
||||
validator := &recordingValidationPreparer{plan: plan}
|
||||
client := &sequenceLLM{responses: tc.responses}
|
||||
repairer := &recordingRepairer{next: NewDefaultOutputRepairer(client)}
|
||||
runner := NewRunnerWithRepairer(
|
||||
&fakePromptRepo{def: definition},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
|
||||
"exec": {ID: "exec", BackendID: "custom", Model: "profile-model"},
|
||||
}},
|
||||
fakeBackendResolver{backends: map[string]domain.Backend{
|
||||
"custom": {ID: "custom", Endpoint: "http://backend.example/v1"},
|
||||
}},
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
client,
|
||||
validator,
|
||||
repairer,
|
||||
nil,
|
||||
)
|
||||
|
||||
res, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Execution: &domain.ExecutionTargetOverride{Endpoint: "http://override/v1", Model: "override-model", TimeoutSeconds: intPtr(22)},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
if repairer.calls != 1 || res.Validation.RepairAttempts != 1 {
|
||||
t.Fatalf("expected one bounded repair, calls=%d attempts=%d", repairer.calls, res.Validation.RepairAttempts)
|
||||
}
|
||||
if len(repairer.reqs) != 1 {
|
||||
t.Fatalf("expected one repair request, got %d", len(repairer.reqs))
|
||||
}
|
||||
if repairer.reqs[0].Target.Endpoint != "http://override/v1" || repairer.reqs[0].Target.Model != "override-model" {
|
||||
t.Fatalf("expected repair to use effective target, got %+v", repairer.reqs[0].Target)
|
||||
}
|
||||
if repairer.reqs[0].Target.TimeoutSeconds != 22 {
|
||||
t.Fatalf("expected repair to use effective timeout, got %d", repairer.reqs[0].Target.TimeoutSeconds)
|
||||
}
|
||||
if llmClient.lastReq.Target.BackendID != "custom" ||
|
||||
repairer.reqs[0].Target.BackendID != "custom" ||
|
||||
res.SelectedBackendID != "custom" {
|
||||
t.Fatalf("expected backend identity in generation, repair, and result: generate=%q repair=%q result=%q",
|
||||
llmClient.lastReq.Target.BackendID, repairer.reqs[0].Target.BackendID, res.SelectedBackendID)
|
||||
result, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
SessionID: " repair-session ",
|
||||
APIKey: "direct-secret",
|
||||
Inputs: singleInputRef(),
|
||||
Execution: tc.execution,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("run: %v", err)
|
||||
}
|
||||
if len(client.requests) != tc.wantRepairs+1 || len(repairer.reqs) != tc.wantRepairs {
|
||||
t.Fatalf(
|
||||
"generation/repair calls = (%d, %d), want (%d, %d)",
|
||||
len(client.requests), len(repairer.reqs), tc.wantRepairs+1, tc.wantRepairs,
|
||||
)
|
||||
}
|
||||
if len(plan.artifacts) != tc.wantRepairs+1 {
|
||||
t.Fatalf("validation calls = %d, want %d", len(plan.artifacts), tc.wantRepairs+1)
|
||||
}
|
||||
|
||||
initialRequest := client.requests[0]
|
||||
if initialRequest.TargetPresence != tc.wantPresence {
|
||||
t.Fatalf("initial target presence = %+v, want %+v", initialRequest.TargetPresence, tc.wantPresence)
|
||||
}
|
||||
if initialRequest.Target.APIKey != "direct-secret" || initialRequest.Target.BackendID != "custom" ||
|
||||
initialRequest.Prompt.SessionID != "repair-session" {
|
||||
t.Fatalf("initial common request fields = %+v", initialRequest)
|
||||
}
|
||||
if initialRequest.Target.Endpoint != "http://backend.example/v1" ||
|
||||
initialRequest.Target.Model != "profile-model" ||
|
||||
initialRequest.Target.Temperature != 0 || initialRequest.Target.MaxTokens != 0 || initialRequest.Target.TopP != 0 {
|
||||
t.Fatalf("initial effective target = %+v", initialRequest.Target)
|
||||
}
|
||||
if tc.structured != (initialRequest.StructuredOutput != nil) {
|
||||
t.Fatalf("initial structured output = %+v, want present %v", initialRequest.StructuredOutput, tc.structured)
|
||||
}
|
||||
if tc.structured && (initialRequest.StructuredOutput.JSONSchema == nil ||
|
||||
initialRequest.StructuredOutput.JSONSchema.Name != "p_1") {
|
||||
t.Fatalf("initial JSON Schema metadata = %+v, want derived schema name p_1", initialRequest.StructuredOutput)
|
||||
}
|
||||
|
||||
for index, req := range repairer.reqs {
|
||||
if req.Attempt != index+1 || req.MaxAttempts != tc.budget || req.Mode != tc.mode {
|
||||
t.Fatalf("repair request %d progression = %+v", index, req)
|
||||
}
|
||||
if req.PreviousOutput != tc.responses[index].Content ||
|
||||
!reflect.DeepEqual(req.ValidationErrors, tc.validationResults[index].Errors) {
|
||||
t.Fatalf("repair request %d prior state = %+v", index, req)
|
||||
}
|
||||
if req.TargetPresence != tc.wantPresence || !reflect.DeepEqual(req.Target, initialRequest.Target) ||
|
||||
req.SessionID != initialRequest.Prompt.SessionID ||
|
||||
!reflect.DeepEqual(req.StructuredOutput, initialRequest.StructuredOutput) {
|
||||
t.Fatalf("repair request %d common fields drifted: %+v", index, req)
|
||||
}
|
||||
|
||||
generated := client.requests[index+1]
|
||||
if generated.TargetPresence != initialRequest.TargetPresence ||
|
||||
!reflect.DeepEqual(generated.Target, initialRequest.Target) ||
|
||||
generated.Prompt.SessionID != initialRequest.Prompt.SessionID ||
|
||||
!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)
|
||||
}
|
||||
}
|
||||
|
||||
lastResponse := tc.responses[tc.wantRepairs]
|
||||
if result.RawOutput != lastResponse.Content || string(result.Artifact.Body) != lastResponse.Content {
|
||||
t.Fatalf("final output = (%q, %q), want %q", result.RawOutput, result.Artifact.Body, lastResponse.Content)
|
||||
}
|
||||
if result.Validation.Status != tc.wantStatus || result.Validation.RepairAttempts != tc.wantRepairs {
|
||||
t.Fatalf("final validation = %+v, want status %q and %d repairs", result.Validation, tc.wantStatus, tc.wantRepairs)
|
||||
}
|
||||
if result.SelectedBackendID != "custom" || result.SessionID != "repair-session" ||
|
||||
result.EffectiveModelParams.APIKey != "" {
|
||||
t.Fatalf("result execution metadata = %+v", result)
|
||||
}
|
||||
|
||||
var wantUsage domain.TokenUsage
|
||||
for _, response := range tc.responses[:tc.wantRepairs+1] {
|
||||
wantUsage.PromptTokens += response.Usage.PromptTokens
|
||||
wantUsage.CompletionTokens += response.Usage.CompletionTokens
|
||||
wantUsage.TotalTokens += response.Usage.TotalTokens
|
||||
wantUsage.CachedTokens += response.Usage.CachedTokens
|
||||
wantUsage.CacheWriteTokens += response.Usage.CacheWriteTokens
|
||||
}
|
||||
if result.Usage != wantUsage {
|
||||
t.Fatalf("cumulative usage = %+v, want %+v", result.Usage, wantUsage)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2171,96 +2437,6 @@ func TestRunnerSchedulesInitialAndRepairGenerationThroughOneBackendPool(t *testi
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunRepairCarriesEffectiveSessionID(t *testing.T) {
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: `{"broken":`}}
|
||||
runner := NewRunnerWithRepairer(
|
||||
&fakePromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 1)},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
|
||||
"exec": {ID: "exec", Endpoint: "http://example.test/v1", Model: "model"},
|
||||
}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
validate.NewStandardValidator("."),
|
||||
NewDefaultOutputRepairer(llmClient), nil)
|
||||
|
||||
result, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
SessionID: " repair-session ",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
if llmClient.calls != 2 {
|
||||
t.Fatalf("expected initial generation and one repair, got %d calls", llmClient.calls)
|
||||
}
|
||||
if llmClient.lastReq.Prompt.SessionID != "repair-session" {
|
||||
t.Fatalf("expected repair generation to retain effective session, got %q", llmClient.lastReq.Prompt.SessionID)
|
||||
}
|
||||
if result.SessionID != "repair-session" {
|
||||
t.Fatalf("expected result to retain effective session, got %q", result.SessionID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerRunJSONSchemaRepairCarriesStructuredOutputSpec(t *testing.T) {
|
||||
def := promptDef(domain.FormatJSON, domain.ValidationJSONSchema, 1)
|
||||
def.Validation.SchemaPath = "events.schema.json"
|
||||
|
||||
validator := &fakeValidator{
|
||||
result: domain.ValidationResult{
|
||||
Status: domain.ValidationFailed,
|
||||
Mode: domain.ValidationJSONSchema,
|
||||
Errors: []string{"schema mismatch"},
|
||||
IsValid: false,
|
||||
},
|
||||
schemaDoc: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"events": map[string]any{"type": "array"},
|
||||
},
|
||||
},
|
||||
}
|
||||
repairer := &fakeRepairer{
|
||||
responses: []*domain.GenerateResponse{
|
||||
{Content: `{"events":[]}`},
|
||||
},
|
||||
}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: `{"events":[1]}`}}
|
||||
runner := NewRunnerWithRepairer(
|
||||
&fakePromptRepo{def: def},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, nil,
|
||||
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
validator,
|
||||
repairer, nil)
|
||||
|
||||
_, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
if llmClient.lastReq.StructuredOutput == nil || llmClient.lastReq.StructuredOutput.JSONSchema == nil {
|
||||
t.Fatalf("expected initial llm request to include structured output, got %+v", llmClient.lastReq.StructuredOutput)
|
||||
}
|
||||
if len(repairer.reqs) != 1 {
|
||||
t.Fatalf("expected one repair request, got %d", len(repairer.reqs))
|
||||
}
|
||||
if repairer.reqs[0].StructuredOutput == nil || repairer.reqs[0].StructuredOutput.JSONSchema == nil {
|
||||
t.Fatalf("expected repair request structured output, got %+v", repairer.reqs[0].StructuredOutput)
|
||||
}
|
||||
if repairer.reqs[0].StructuredOutput.JSONSchema.Name != "p_1" {
|
||||
t.Fatalf("expected derived schema name p_1, got %q", repairer.reqs[0].StructuredOutput.JSONSchema.Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecutionProfileToTargetPopulatesAllFieldsAndCopiesExtraParams(t *testing.T) {
|
||||
src := &domain.ExecutionProfile{
|
||||
ID: "exec",
|
||||
@@ -2319,10 +2495,7 @@ func TestResolveExecutionTargetProfileValuesPopulateAllSupportedFields(t *testin
|
||||
},
|
||||
}
|
||||
|
||||
target, presence, err := resolveExecutionTarget(nil, profileValue, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
target, presence := resolveExecutionTarget(nil, profileValue, nil)
|
||||
if presence != (domain.ExecutionTargetPresence{}) {
|
||||
t.Fatalf("expected no request override presence, got %+v", presence)
|
||||
}
|
||||
@@ -2374,10 +2547,7 @@ func TestResolveExecutionTargetRuntimeOverridesBeatProfileForAllOverrideableFiel
|
||||
},
|
||||
}
|
||||
|
||||
target, presence, err := resolveExecutionTarget(nil, profileValue, override)
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
target, presence := resolveExecutionTarget(nil, profileValue, override)
|
||||
if presence != (domain.ExecutionTargetPresence{Temperature: true, MaxTokens: true, TopP: true, TimeoutSeconds: true}) {
|
||||
t.Fatalf("unexpected override presence: %+v", presence)
|
||||
}
|
||||
@@ -2424,12 +2594,9 @@ func TestResolveExecutionTargetReasoningOverrideStates(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
target, _, err := resolveExecutionTarget(nil, profileValue, &domain.ExecutionTargetOverride{
|
||||
target, _ := resolveExecutionTarget(nil, profileValue, &domain.ExecutionTargetOverride{
|
||||
ReasoningEffort: tt.override,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("resolve execution target: %v", err)
|
||||
}
|
||||
if target.ReasoningEffort != tt.want {
|
||||
t.Fatalf("reasoning effort = %q, want %q", target.ReasoningEffort, tt.want)
|
||||
}
|
||||
@@ -2558,10 +2725,7 @@ func TestResolveExecutionTargetUsesBackendProfileAndRequestPrecedence(t *testing
|
||||
ExtraParams: map[string]any{"request": true},
|
||||
}
|
||||
|
||||
target, _, err := resolveExecutionTarget(backendValue, profileValue, override)
|
||||
if err != nil {
|
||||
t.Fatalf("resolve target: %v", err)
|
||||
}
|
||||
target, _ := resolveExecutionTarget(backendValue, profileValue, override)
|
||||
if target.BackendID != "custom" {
|
||||
t.Fatalf("endpoint override changed backend identity: %+v", target)
|
||||
}
|
||||
@@ -2572,12 +2736,9 @@ func TestResolveExecutionTargetUsesBackendProfileAndRequestPrecedence(t *testing
|
||||
t.Fatalf("expected whole-map request replacement, got %#v", target.ExtraParams)
|
||||
}
|
||||
|
||||
target, _, err = resolveExecutionTarget(backendValue, &domain.ExecutionProfile{
|
||||
target, _ = resolveExecutionTarget(backendValue, &domain.ExecutionProfile{
|
||||
ID: "exec", BackendID: "custom", Model: "profile-model",
|
||||
}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("resolve backend defaults: %v", err)
|
||||
}
|
||||
if target.Endpoint != backendValue.Endpoint ||
|
||||
target.APIKeyEnv != backendValue.APIKeyEnv ||
|
||||
!reflect.DeepEqual(target.ExtraParams, backendValue.ExtraParams) {
|
||||
|
||||
311
internal/validate/cancellation_test.go
Normal file
311
internal/validate/cancellation_test.go
Normal file
@@ -0,0 +1,311 @@
|
||||
package validate
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io/fs"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
"github.com/santhosh-tekuri/jsonschema/v6"
|
||||
)
|
||||
|
||||
func TestValidationCancellationBeforeWorkDoesNotOpenSchemaSource(t *testing.T) {
|
||||
source := &countingSchemaFS{FS: fstest.MapFS{
|
||||
"schema.json": {Data: []byte(`{"type":"object"}`)},
|
||||
}}
|
||||
validator := NewFSValidator(source, ".").(ValidationPreparer)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
plan, err := validator.PrepareValidation(ctx, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
if plan != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("PrepareValidation() = (%v, %v), want nil plan and context cancellation", plan, err)
|
||||
}
|
||||
if opens := source.opens.Load(); opens != 0 {
|
||||
t.Fatalf("schema source opens = %d, want 0", opens)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidationCancellationBetweenReferencedSchemaReadChunks(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
source := &controlledSchemaFS{
|
||||
FS: fstest.MapFS{
|
||||
"root.json": {Data: []byte(`{"$ref":"child.json"}`)},
|
||||
"child.json": {Data: []byte(`{"type":"string"}` + strings.Repeat(" ", schemaReadChunkSize*2))},
|
||||
},
|
||||
target: "child.json",
|
||||
cancel: cancel,
|
||||
}
|
||||
validator := NewFSValidator(source, ".").(ValidationPreparer)
|
||||
|
||||
plan, err := validator.PrepareValidation(ctx, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "root.json",
|
||||
})
|
||||
if plan != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("PrepareValidation() = (%v, %v), want nil plan and context cancellation", plan, err)
|
||||
}
|
||||
if reads := source.reads.Load(); reads != 1 {
|
||||
t.Fatalf("controlled child reads = %d, want 1", reads)
|
||||
}
|
||||
if closes := source.closes.Load(); closes != 1 {
|
||||
t.Fatalf("controlled child closes = %d, want 1", closes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCancellationAfterSynchronousCallWins(t *testing.T) {
|
||||
ctx := newCheckpointContext(2)
|
||||
value, err := decodeJSONValue(ctx, []byte(`{"value":1}`))
|
||||
if value != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("decodeJSONValue() = (%v, %v), want nil value and context cancellation", value, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompileCancellationAfterSynchronousCallWins(t *testing.T) {
|
||||
for _, dependencyErr := range []error{nil, errors.New("compile failed")} {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
compiled := &jsonschema.Schema{}
|
||||
schema, err := compileJSONSchema(ctx, "promptkit-schema:/root.json", func(string) (*jsonschema.Schema, error) {
|
||||
cancel()
|
||||
return compiled, dependencyErr
|
||||
})
|
||||
if schema != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("compileJSONSchema() = (%v, %v), want authoritative context cancellation", schema, err)
|
||||
}
|
||||
if dependencyErr != nil && errors.Is(err, dependencyErr) {
|
||||
t.Fatalf("compileJSONSchema() error = %v, dependency error should not win", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecutionCancellationAfterSynchronousCallWins(t *testing.T) {
|
||||
for _, dependencyErr := range []error{nil, errors.New("schema mismatch")} {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
executor := &controlledSchemaExecutor{cancel: cancel, err: dependencyErr}
|
||||
validationErrors, err := executeJSONSchema(ctx, executor, map[string]any{"ok": true})
|
||||
if validationErrors != nil || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("executeJSONSchema() = (%v, %v), want no result and context cancellation", validationErrors, err)
|
||||
}
|
||||
if calls := executor.calls.Load(); calls != 1 {
|
||||
t.Fatalf("schema execution calls = %d, want 1", calls)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancellationDoesNotDetachBlockedSchemaRead(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
source := &controlledSchemaFS{
|
||||
FS: fstest.MapFS{
|
||||
"schema.json": {Data: []byte(`{"type":"object"}`)},
|
||||
},
|
||||
target: "schema.json",
|
||||
started: started,
|
||||
release: release,
|
||||
}
|
||||
validator := NewFSValidator(source, ".").(ValidationPreparer)
|
||||
type outcome struct {
|
||||
plan PreparedValidation
|
||||
err error
|
||||
}
|
||||
result := make(chan outcome, 1)
|
||||
go func() {
|
||||
plan, err := validator.PrepareValidation(ctx, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
result <- outcome{plan: plan, err: err}
|
||||
}()
|
||||
|
||||
waitForSignal(t, started, "schema read to start")
|
||||
cancel()
|
||||
select {
|
||||
case got := <-result:
|
||||
t.Fatalf("blocked dependency returned before release: (%v, %v)", got.plan, got.err)
|
||||
default:
|
||||
}
|
||||
close(release)
|
||||
got := waitForOutcome(t, result)
|
||||
if got.plan != nil || !errors.Is(got.err, context.Canceled) {
|
||||
t.Fatalf("PrepareValidation() after release = (%v, %v), want nil plan and context cancellation", got.plan, got.err)
|
||||
}
|
||||
if reads := source.reads.Load(); reads != 1 {
|
||||
t.Fatalf("blocked reads = %d, want 1 completed read", reads)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancellationDoesNotDetachBlockedSchemaExecution(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
executor := &controlledSchemaExecutor{started: started, release: release}
|
||||
result := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := executeJSONSchema(ctx, executor, map[string]any{"ok": true})
|
||||
result <- err
|
||||
}()
|
||||
|
||||
waitForSignal(t, started, "schema execution to start")
|
||||
cancel()
|
||||
select {
|
||||
case err := <-result:
|
||||
t.Fatalf("blocked dependency returned before release: %v", err)
|
||||
default:
|
||||
}
|
||||
close(release)
|
||||
select {
|
||||
case err := <-result:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("executeJSONSchema() after release = %v, want context cancellation", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for released schema execution")
|
||||
}
|
||||
if calls := executor.calls.Load(); calls != 1 {
|
||||
t.Fatalf("schema execution calls = %d, want 1 completed call", calls)
|
||||
}
|
||||
}
|
||||
|
||||
type countingSchemaFS struct {
|
||||
fs.FS
|
||||
opens atomic.Int32
|
||||
}
|
||||
|
||||
func (f *countingSchemaFS) Open(name string) (fs.File, error) {
|
||||
f.opens.Add(1)
|
||||
return f.FS.Open(name)
|
||||
}
|
||||
|
||||
type controlledSchemaFS struct {
|
||||
fs.FS
|
||||
target string
|
||||
cancel context.CancelFunc
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
once sync.Once
|
||||
reads atomic.Int32
|
||||
closes atomic.Int32
|
||||
}
|
||||
|
||||
func (f *controlledSchemaFS) Open(name string) (fs.File, error) {
|
||||
file, err := f.FS.Open(name)
|
||||
if err != nil || name != f.target {
|
||||
return file, err
|
||||
}
|
||||
return &controlledSchemaFile{File: file, owner: f}, nil
|
||||
}
|
||||
|
||||
type controlledSchemaFile struct {
|
||||
fs.File
|
||||
owner *controlledSchemaFS
|
||||
}
|
||||
|
||||
func (f *controlledSchemaFile) Read(buffer []byte) (int, error) {
|
||||
f.owner.once.Do(func() {
|
||||
if f.owner.started != nil {
|
||||
close(f.owner.started)
|
||||
}
|
||||
if f.owner.release != nil {
|
||||
<-f.owner.release
|
||||
}
|
||||
})
|
||||
n, err := f.File.Read(buffer)
|
||||
f.owner.reads.Add(1)
|
||||
if f.owner.cancel != nil {
|
||||
f.owner.cancel()
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (f *controlledSchemaFile) Close() error {
|
||||
f.owner.closes.Add(1)
|
||||
return f.File.Close()
|
||||
}
|
||||
|
||||
type controlledSchemaExecutor struct {
|
||||
cancel context.CancelFunc
|
||||
err error
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
calls atomic.Int32
|
||||
}
|
||||
|
||||
func (e *controlledSchemaExecutor) Validate(any) error {
|
||||
e.calls.Add(1)
|
||||
if e.started != nil {
|
||||
close(e.started)
|
||||
}
|
||||
if e.release != nil {
|
||||
<-e.release
|
||||
}
|
||||
if e.cancel != nil {
|
||||
e.cancel()
|
||||
}
|
||||
return e.err
|
||||
}
|
||||
|
||||
type checkpointContext struct {
|
||||
context.Context
|
||||
mu sync.Mutex
|
||||
remaining int
|
||||
canceled bool
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
func newCheckpointContext(checksUntilCancel int) *checkpointContext {
|
||||
return &checkpointContext{
|
||||
Context: context.Background(),
|
||||
remaining: checksUntilCancel,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *checkpointContext) Done() <-chan struct{} {
|
||||
return c.done
|
||||
}
|
||||
|
||||
func (c *checkpointContext) Err() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.canceled {
|
||||
return context.Canceled
|
||||
}
|
||||
c.remaining--
|
||||
if c.remaining == 0 {
|
||||
c.canceled = true
|
||||
close(c.done)
|
||||
return context.Canceled
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func waitForSignal(t *testing.T, signal <-chan struct{}, description string) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-signal:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatalf("timed out waiting for %s", description)
|
||||
}
|
||||
}
|
||||
|
||||
func waitForOutcome[T any](t *testing.T, result <-chan T) T {
|
||||
t.Helper()
|
||||
select {
|
||||
case value := <-result:
|
||||
return value
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for released dependency")
|
||||
var zero T
|
||||
return zero
|
||||
}
|
||||
}
|
||||
@@ -1,15 +1,18 @@
|
||||
package validate
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"net/url"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
@@ -19,6 +22,8 @@ import (
|
||||
|
||||
const jsonSchemaDraft2020 = "https://json-schema.org/draft/2020-12/schema"
|
||||
|
||||
const schemaReadChunkSize = 64 * 1024
|
||||
|
||||
// StandardValidator provides basic, JSON, and JSON Schema output validation.
|
||||
type StandardValidator struct {
|
||||
schemaBaseDir string
|
||||
@@ -45,7 +50,121 @@ func (v *FSValidator) Validate(ctx context.Context, artifact *domain.Artifact, c
|
||||
return validateArtifact(ctx, artifact, contract, v.validateJSONSchema)
|
||||
}
|
||||
|
||||
type schemaValidatorFunc func(instance any, schemaPath string) ([]string, error)
|
||||
type preparedValidation struct {
|
||||
contract domain.OutputContract
|
||||
schemaDocument any
|
||||
schema schemaExecutor
|
||||
}
|
||||
|
||||
func (p *preparedValidation) Validate(ctx context.Context, artifact *domain.Artifact) (domain.ValidationResult, error) {
|
||||
return validateArtifact(ctx, artifact, p.contract, p.validateJSONSchema)
|
||||
}
|
||||
|
||||
func (p *preparedValidation) SchemaDocument() any {
|
||||
return p.schemaDocument
|
||||
}
|
||||
|
||||
func (p *preparedValidation) validateJSONSchema(ctx context.Context, instance any, _ string) ([]string, error) {
|
||||
if p.schema == nil {
|
||||
return nil, errors.New("prepared JSON schema is unavailable")
|
||||
}
|
||||
return executeJSONSchema(ctx, p.schema, instance)
|
||||
}
|
||||
|
||||
func (v *StandardValidator) PrepareValidation(ctx context.Context, contract domain.OutputContract) (PreparedValidation, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
prepared := &preparedValidation{contract: contract}
|
||||
if contract.ValidationMode != domain.ValidationJSONSchema {
|
||||
return prepared, nil
|
||||
}
|
||||
|
||||
resolvedSchemaPath, err := v.resolveSchemaPath(ctx, contract.SchemaPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
schemaDocument, err := loadJSONSchemaFile(ctx, resolvedSchemaPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", resolvedSchemaPath, err)
|
||||
}
|
||||
schemaRoot, err := v.schemaRoot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
compiler := newSchemaCompiler(standardSchemaLoader{ctx: ctx, root: schemaRoot})
|
||||
resourceURL := fileSchemaResourceURL(resolvedSchemaPath)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := compiler.AddResource(resourceURL.String(), schemaDocument); err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return nil, contextErr
|
||||
}
|
||||
return nil, fmt.Errorf("failed to register JSON schema %q: %w", resolvedSchemaPath, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
schema, err := compileJSONSchema(ctx, resourceURL.String(), compiler.Compile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", resolvedSchemaPath, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
prepared.schemaDocument = schemaDocument
|
||||
prepared.schema = schema
|
||||
return prepared, nil
|
||||
}
|
||||
|
||||
func (v *FSValidator) PrepareValidation(ctx context.Context, contract domain.OutputContract) (PreparedValidation, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
prepared := &preparedValidation{contract: contract}
|
||||
if contract.ValidationMode != domain.ValidationJSONSchema {
|
||||
return prepared, nil
|
||||
}
|
||||
|
||||
schemaName, schemaDocument, err := v.loadSchemaDocument(ctx, contract.SchemaPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resourceURL := fsSchemaResourceURL(schemaName)
|
||||
compiler := newSchemaCompiler(fsSchemaLoader{ctx: ctx, fsys: v.fsys, root: filecatalog.CleanFSRoot(v.root)})
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := compiler.AddResource(resourceURL.String(), schemaDocument); err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return nil, contextErr
|
||||
}
|
||||
return nil, fmt.Errorf("failed to register JSON schema %q: %w", schemaName, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
schema, err := compileJSONSchema(ctx, resourceURL.String(), compiler.Compile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", schemaName, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
prepared.schemaDocument = schemaDocument
|
||||
prepared.schema = schema
|
||||
return prepared, nil
|
||||
}
|
||||
|
||||
type schemaValidatorFunc func(ctx context.Context, instance any, schemaPath string) ([]string, error)
|
||||
|
||||
func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract, validateSchema schemaValidatorFunc) (domain.ValidationResult, error) {
|
||||
select {
|
||||
@@ -70,7 +189,11 @@ func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract d
|
||||
res.IsValid = true
|
||||
return res, nil
|
||||
case domain.ValidationBasic:
|
||||
if strings.TrimSpace(string(artifact.Body)) == "" {
|
||||
empty := strings.TrimSpace(string(artifact.Body)) == ""
|
||||
if err := ctx.Err(); err != nil {
|
||||
return domain.ValidationResult{}, err
|
||||
}
|
||||
if empty {
|
||||
res.Status = domain.ValidationFailed
|
||||
res.IsValid = false
|
||||
res.Errors = []string{"output is empty"}
|
||||
@@ -80,26 +203,32 @@ func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract d
|
||||
res.IsValid = true
|
||||
return res, nil
|
||||
case domain.ValidationJSON:
|
||||
_, jsonErr := parseJSON(artifact.Body)
|
||||
if jsonErr != nil {
|
||||
valid := json.Valid(artifact.Body)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return domain.ValidationResult{}, err
|
||||
}
|
||||
if !valid {
|
||||
res.Status = domain.ValidationFailed
|
||||
res.IsValid = false
|
||||
res.Errors = []string{fmt.Sprintf("invalid JSON: %v", jsonErr)}
|
||||
res.Errors = []string{"invalid JSON"}
|
||||
return res, nil
|
||||
}
|
||||
res.Status = domain.ValidationPassed
|
||||
res.IsValid = true
|
||||
return res, nil
|
||||
case domain.ValidationJSONSchema:
|
||||
instance, jsonErr := parseJSON(artifact.Body)
|
||||
instance, jsonErr := decodeJSONValue(ctx, artifact.Body)
|
||||
if jsonErr != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return domain.ValidationResult{}, contextErr
|
||||
}
|
||||
res.Status = domain.ValidationFailed
|
||||
res.IsValid = false
|
||||
res.Errors = []string{fmt.Sprintf("invalid JSON: %v", jsonErr)}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
validationErrors, err := validateSchema(instance, contract.SchemaPath)
|
||||
validationErrors, err := validateSchema(ctx, instance, contract.SchemaPath)
|
||||
if err != nil {
|
||||
return domain.ValidationResult{}, err
|
||||
}
|
||||
@@ -118,8 +247,8 @@ func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract d
|
||||
}
|
||||
}
|
||||
|
||||
func (v *StandardValidator) validateJSONSchema(instance any, schemaPath string) ([]string, error) {
|
||||
resolvedSchemaPath, err := v.resolveSchemaPath(schemaPath)
|
||||
func (v *StandardValidator) validateJSONSchema(ctx context.Context, instance any, schemaPath string) ([]string, error) {
|
||||
resolvedSchemaPath, err := v.resolveSchemaPath(ctx, schemaPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -128,20 +257,21 @@ func (v *StandardValidator) validateJSONSchema(instance any, schemaPath string)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
compiler := newSchemaCompiler(standardSchemaLoader{root: schemaRoot})
|
||||
schema, err := compiler.Compile(resolvedSchemaPath)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
compiler := newSchemaCompiler(standardSchemaLoader{ctx: ctx, root: schemaRoot})
|
||||
resourceURL := fileSchemaResourceURL(resolvedSchemaPath)
|
||||
schema, err := compileJSONSchema(ctx, resourceURL.String(), compiler.Compile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", resolvedSchemaPath, err)
|
||||
}
|
||||
|
||||
if err := schema.Validate(instance); err != nil {
|
||||
return []string{fmt.Sprintf("json schema validation failed: %v", err)}, nil
|
||||
}
|
||||
return nil, nil
|
||||
return executeJSONSchema(ctx, schema, instance)
|
||||
}
|
||||
|
||||
func (v *FSValidator) validateJSONSchema(instance any, schemaPath string) ([]string, error) {
|
||||
schemaName, schemaDoc, err := v.loadSchemaDocument(schemaPath)
|
||||
func (v *FSValidator) validateJSONSchema(ctx context.Context, instance any, schemaPath string) ([]string, error) {
|
||||
schemaName, schemaDoc, err := v.loadSchemaDocument(ctx, schemaPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -150,71 +280,60 @@ func (v *FSValidator) validateJSONSchema(instance any, schemaPath string) ([]str
|
||||
if err := validateSchemaDialect(schemaDoc); err != nil {
|
||||
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", schemaName, err)
|
||||
}
|
||||
compiler := newSchemaCompiler(fsSchemaLoader{fsys: v.fsys, root: filecatalog.CleanFSRoot(v.root)})
|
||||
if err := compiler.AddResource(resourceURL, schemaDoc); err != nil {
|
||||
compiler := newSchemaCompiler(fsSchemaLoader{ctx: ctx, fsys: v.fsys, root: filecatalog.CleanFSRoot(v.root)})
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := compiler.AddResource(resourceURL.String(), schemaDoc); err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return nil, contextErr
|
||||
}
|
||||
return nil, fmt.Errorf("failed to register JSON schema %q: %w", schemaName, err)
|
||||
}
|
||||
schema, err := compiler.Compile(resourceURL)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
schema, err := compileJSONSchema(ctx, resourceURL.String(), compiler.Compile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", schemaName, err)
|
||||
}
|
||||
|
||||
if err := schema.Validate(instance); err != nil {
|
||||
return []string{fmt.Sprintf("json schema validation failed: %v", err)}, nil
|
||||
}
|
||||
return nil, nil
|
||||
return executeJSONSchema(ctx, schema, instance)
|
||||
}
|
||||
|
||||
func parseJSON(body []byte) (any, error) {
|
||||
var v any
|
||||
if err := json.Unmarshal(body, &v); err != nil {
|
||||
func decodeJSONValue(ctx context.Context, body []byte) (any, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
decoder := json.NewDecoder(bytes.NewReader(body))
|
||||
decoder.UseNumber()
|
||||
|
||||
func (v *StandardValidator) LoadSchemaDocument(ctx context.Context, schemaPath string) (any, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
resolved, err := v.resolveSchemaPath(schemaPath)
|
||||
if err != nil {
|
||||
var value any
|
||||
decodeErr := decoder.Decode(&value)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
raw, err := os.ReadFile(resolved)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read schema file %q: %w", resolved, err)
|
||||
if decodeErr != nil {
|
||||
return nil, decodeErr
|
||||
}
|
||||
|
||||
var doc any
|
||||
if err := json.Unmarshal(raw, &doc); err != nil {
|
||||
return nil, fmt.Errorf("failed to decode JSON schema %q: %w", resolved, err)
|
||||
}
|
||||
if err := validateSchemaDialect(doc); err != nil {
|
||||
return nil, fmt.Errorf("failed to decode JSON schema %q: %w", resolved, err)
|
||||
}
|
||||
return doc, nil
|
||||
}
|
||||
|
||||
func (v *FSValidator) LoadSchemaDocument(ctx context.Context, schemaPath string) (any, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
_, doc, err := v.loadSchemaDocument(schemaPath)
|
||||
if err != nil {
|
||||
var trailing any
|
||||
trailingErr := decoder.Decode(&trailing)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return doc, nil
|
||||
if errors.Is(trailingErr, io.EOF) {
|
||||
return value, nil
|
||||
} else if trailingErr != nil {
|
||||
return nil, trailingErr
|
||||
}
|
||||
return nil, errors.New("multiple JSON values")
|
||||
}
|
||||
|
||||
func (v *StandardValidator) resolveSchemaPath(schemaPath string) (string, error) {
|
||||
func (v *StandardValidator) resolveSchemaPath(ctx context.Context, schemaPath string) (string, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(schemaPath) == "" {
|
||||
return "", errors.New("schema path is required for json_schema validation")
|
||||
}
|
||||
@@ -223,13 +342,25 @@ func (v *StandardValidator) resolveSchemaPath(schemaPath string) (string, error)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
resolved, err := containedFilesystemPath(root, schemaPath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, err := os.Stat(resolved); err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return "", contextErr
|
||||
}
|
||||
return "", fmt.Errorf("failed to access schema file %q: %w", resolved, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return resolved, nil
|
||||
}
|
||||
@@ -250,28 +381,39 @@ func (v *StandardValidator) schemaRoot() (string, error) {
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
func (v *FSValidator) loadSchemaDocument(schemaPath string) (string, any, error) {
|
||||
resolved, err := v.resolveSchemaPath(schemaPath)
|
||||
func (v *FSValidator) loadSchemaDocument(ctx context.Context, schemaPath string) (string, any, error) {
|
||||
resolved, err := v.resolveSchemaPath(ctx, schemaPath)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
raw, err := fs.ReadFile(v.fsys, resolved)
|
||||
raw, err := readSchemaFile(ctx, func() (fs.File, error) {
|
||||
return v.fsys.Open(resolved)
|
||||
})
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("failed to read schema file %q: %w", resolved, err)
|
||||
}
|
||||
|
||||
var doc any
|
||||
if err := json.Unmarshal(raw, &doc); err != nil {
|
||||
doc, err := decodeJSONValue(ctx, raw)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("failed to decode JSON schema %q: %w", resolved, err)
|
||||
}
|
||||
if err := validateSchemaDialect(doc); err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return "", nil, contextErr
|
||||
}
|
||||
return "", nil, fmt.Errorf("failed to decode JSON schema %q: %w", resolved, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return resolved, doc, nil
|
||||
}
|
||||
|
||||
func (v *FSValidator) resolveSchemaPath(schemaPath string) (string, error) {
|
||||
func (v *FSValidator) resolveSchemaPath(ctx context.Context, schemaPath string) (string, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(schemaPath) == "" {
|
||||
return "", errors.New("schema path is required for json_schema validation")
|
||||
}
|
||||
@@ -282,8 +424,14 @@ func (v *FSValidator) resolveSchemaPath(schemaPath string) (string, error) {
|
||||
cleanRoot := filecatalog.CleanFSRoot(v.root)
|
||||
rootInfo, err := fs.Stat(v.fsys, cleanRoot)
|
||||
if err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return "", contextErr
|
||||
}
|
||||
return "", fmt.Errorf("failed to access schema source %q: %w", cleanRoot, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var resolved string
|
||||
if rootInfo.IsDir() {
|
||||
@@ -297,15 +445,21 @@ func (v *FSValidator) resolveSchemaPath(schemaPath string) (string, error) {
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if cleanSchemaPath != path.Base(cleanRoot) {
|
||||
if cleanSchemaPath != strings.TrimSpace(path.Base(cleanRoot)) {
|
||||
return "", fmt.Errorf("schema path %q does not match schema file %q", cleanSchemaPath, path.Base(cleanRoot))
|
||||
}
|
||||
resolved = cleanRoot
|
||||
}
|
||||
|
||||
if _, err := fs.Stat(v.fsys, resolved); err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return "", contextErr
|
||||
}
|
||||
return "", fmt.Errorf("failed to access schema file %q: %w", resolved, err)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
@@ -321,8 +475,19 @@ func cleanSchemaFSPath(schemaPath string) (string, error) {
|
||||
return cleaned, nil
|
||||
}
|
||||
|
||||
func fsSchemaResourceURL(schemaName string) string {
|
||||
return "promptkit-schema:///" + strings.TrimPrefix(path.Clean(schemaName), "/")
|
||||
func fileSchemaResourceURL(schemaName string) *url.URL {
|
||||
filePath := filepath.ToSlash(schemaName)
|
||||
if runtime.GOOS == "windows" && !strings.HasPrefix(filePath, "/") {
|
||||
filePath = "/" + filePath
|
||||
}
|
||||
return &url.URL{Scheme: "file", Path: filePath}
|
||||
}
|
||||
|
||||
func fsSchemaResourceURL(schemaName string) *url.URL {
|
||||
return &url.URL{
|
||||
Scheme: "promptkit-schema",
|
||||
Path: "/" + strings.TrimPrefix(path.Clean(schemaName), "/"),
|
||||
}
|
||||
}
|
||||
|
||||
func newSchemaCompiler(loader jsonschema.URLLoader) *jsonschema.Compiler {
|
||||
@@ -332,6 +497,35 @@ func newSchemaCompiler(loader jsonschema.URLLoader) *jsonschema.Compiler {
|
||||
return compiler
|
||||
}
|
||||
|
||||
type schemaExecutor interface {
|
||||
Validate(instance any) error
|
||||
}
|
||||
|
||||
func compileJSONSchema(ctx context.Context, resourceURL string, compile func(string) (*jsonschema.Schema, error)) (*jsonschema.Schema, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
schema, compileErr := compile(resourceURL)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return schema, compileErr
|
||||
}
|
||||
|
||||
func executeJSONSchema(ctx context.Context, schema schemaExecutor, instance any) ([]string, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
validationErr := schema.Validate(instance)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validationErr != nil {
|
||||
return []string{fmt.Sprintf("json schema validation failed: %v", validationErr)}, nil
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func validateSchemaDialect(doc any) error {
|
||||
object, ok := doc.(map[string]any)
|
||||
if !ok {
|
||||
@@ -352,19 +546,33 @@ func validateSchemaDialect(doc any) error {
|
||||
}
|
||||
|
||||
type standardSchemaLoader struct {
|
||||
ctx context.Context
|
||||
root string
|
||||
}
|
||||
|
||||
func (l standardSchemaLoader) Load(resourceURL string) (any, error) {
|
||||
if err := l.ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parsed, err := url.Parse(resourceURL)
|
||||
if err != nil || parsed.Scheme != "file" || parsed.Host != "" || parsed.RawQuery != "" || parsed.Opaque != "" {
|
||||
return nil, fmt.Errorf("schema reference %q is not a contained file reference", resourceURL)
|
||||
}
|
||||
fileName, err := (jsonschema.FileLoader{}).ToFile(resourceURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("schema reference %q is not a contained file reference: %w", resourceURL, err)
|
||||
}
|
||||
resolved, err := containedFilesystemPath(l.root, fileName)
|
||||
if err != nil {
|
||||
if contextErr := l.ctx.Err(); contextErr != nil {
|
||||
return nil, contextErr
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return loadJSONSchemaFile(resolved)
|
||||
if err := l.ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return loadJSONSchemaFile(l.ctx, resolved)
|
||||
}
|
||||
|
||||
func containedFilesystemPath(root, name string) (string, error) {
|
||||
@@ -390,38 +598,94 @@ func containedFilesystemPath(root, name string) (string, error) {
|
||||
return candidate, nil
|
||||
}
|
||||
|
||||
func loadJSONSchemaFile(name string) (any, error) {
|
||||
raw, err := os.ReadFile(name)
|
||||
func loadJSONSchemaFile(ctx context.Context, name string) (any, error) {
|
||||
raw, err := readSchemaFile(ctx, func() (fs.File, error) {
|
||||
return os.Open(name)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var doc any
|
||||
if err := json.Unmarshal(raw, &doc); err != nil {
|
||||
doc, err := decodeJSONValue(ctx, raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateSchemaDialect(doc); err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return nil, contextErr
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return doc, nil
|
||||
}
|
||||
|
||||
func readSchemaFile(ctx context.Context, open func() (fs.File, error)) ([]byte, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
file, openErr := open()
|
||||
if err := ctx.Err(); err != nil {
|
||||
if file != nil {
|
||||
_ = file.Close()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if openErr != nil {
|
||||
if file != nil {
|
||||
_ = file.Close()
|
||||
}
|
||||
return nil, openErr
|
||||
}
|
||||
if file == nil {
|
||||
return nil, errors.New("schema source returned a nil file")
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
var contents []byte
|
||||
chunk := make([]byte, schemaReadChunkSize)
|
||||
for {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
n, readErr := file.Read(chunk)
|
||||
if n > 0 {
|
||||
contents = append(contents, chunk[:n]...)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if errors.Is(readErr, io.EOF) {
|
||||
return contents, nil
|
||||
}
|
||||
if readErr != nil {
|
||||
return nil, readErr
|
||||
}
|
||||
if n == 0 {
|
||||
return nil, io.ErrNoProgress
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type fsSchemaLoader struct {
|
||||
ctx context.Context
|
||||
fsys fs.FS
|
||||
root string
|
||||
}
|
||||
|
||||
func (l fsSchemaLoader) Load(resourceURL string) (any, error) {
|
||||
if err := l.ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parsed, err := url.Parse(resourceURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid schema reference %q: %w", resourceURL, err)
|
||||
}
|
||||
if parsed.Scheme != "promptkit-schema" || parsed.Host != "" {
|
||||
if parsed.Scheme != "promptkit-schema" || parsed.Host != "" || parsed.RawQuery != "" || parsed.Opaque != "" {
|
||||
return nil, fmt.Errorf("schema reference %q is not allowed", resourceURL)
|
||||
}
|
||||
name, err := url.PathUnescape(strings.TrimPrefix(parsed.Path, "/"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid schema reference %q: %w", resourceURL, err)
|
||||
}
|
||||
name := strings.TrimPrefix(parsed.Path, "/")
|
||||
name = path.Clean(name)
|
||||
if l.root == "." {
|
||||
if strings.HasPrefix(name, "../") || name == ".." {
|
||||
@@ -433,21 +697,35 @@ func (l fsSchemaLoader) Load(resourceURL string) (any, error) {
|
||||
|
||||
rootInfo, err := fs.Stat(l.fsys, l.root)
|
||||
if err != nil {
|
||||
if contextErr := l.ctx.Err(); contextErr != nil {
|
||||
return nil, contextErr
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if err := l.ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !rootInfo.IsDir() && name != l.root {
|
||||
return nil, fmt.Errorf("schema reference %q is outside the configured schema file", resourceURL)
|
||||
}
|
||||
|
||||
raw, err := fs.ReadFile(l.fsys, name)
|
||||
raw, err := readSchemaFile(l.ctx, func() (fs.File, error) {
|
||||
return l.fsys.Open(name)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var doc any
|
||||
if err := json.Unmarshal(raw, &doc); err != nil {
|
||||
doc, err := decodeJSONValue(l.ctx, raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateSchemaDialect(doc); err != nil {
|
||||
if contextErr := l.ctx.Err(); contextErr != nil {
|
||||
return nil, contextErr
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if err := l.ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return doc, nil
|
||||
|
||||
@@ -3,8 +3,11 @@ package validate
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -92,6 +95,51 @@ func TestStandardValidatorJSONFailure(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestJSONValidationChecksCompleteSyntaxWithoutChangingArtifact(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
wantValid bool
|
||||
}{
|
||||
{name: "ordinary object", body: `{"count":2,"ok":true}`, wantValid: true},
|
||||
{name: "integer at exact float boundary", body: `9007199254740992`, wantValid: true},
|
||||
{name: "integer beyond exact float boundary", body: `9007199254740993`, wantValid: true},
|
||||
{name: "large exponent", body: `1e400`, wantValid: true},
|
||||
{name: "precise decimal", body: `0.123456789012345678901234567890`, wantValid: true},
|
||||
{name: "surrounding whitespace", body: " \n [1,2,3] \t", wantValid: true},
|
||||
{name: "malformed document", body: `{"count":`, wantValid: false},
|
||||
{name: "trailing value", body: `1 2`, wantValid: false},
|
||||
}
|
||||
|
||||
validator := NewStandardValidator("")
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
body := []byte(tc.body)
|
||||
before := append([]byte(nil), body...)
|
||||
artifact := &domain.Artifact{Body: body}
|
||||
|
||||
result, err := validator.Validate(context.Background(), artifact, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSON,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("validate JSON: %v", err)
|
||||
}
|
||||
if result.IsValid != tc.wantValid {
|
||||
t.Fatalf("valid = %v, want %v; result=%+v", result.IsValid, tc.wantValid, result)
|
||||
}
|
||||
if tc.wantValid && result.Status != domain.ValidationPassed {
|
||||
t.Fatalf("status = %q, want %q", result.Status, domain.ValidationPassed)
|
||||
}
|
||||
if !tc.wantValid && result.Status != domain.ValidationFailed {
|
||||
t.Fatalf("status = %q, want %q", result.Status, domain.ValidationFailed)
|
||||
}
|
||||
if !reflect.DeepEqual(artifact.Body, before) {
|
||||
t.Fatalf("artifact body changed: got %q, want %q", artifact.Body, before)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorJSONSchemaSuccess(t *testing.T) {
|
||||
tmp := t.TempDir()
|
||||
schemaPath := filepath.Join(tmp, "schema.json")
|
||||
@@ -120,6 +168,68 @@ func TestStandardValidatorJSONSchemaSuccess(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorPreparedSchemaSurvivesSourceRemoval(t *testing.T) {
|
||||
tmp := t.TempDir()
|
||||
rootSchema := []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"title": "original root",
|
||||
"type": "object",
|
||||
"required": ["value"],
|
||||
"properties": {
|
||||
"value": {"$ref": "value.json"}
|
||||
}
|
||||
}`)
|
||||
rootPath := filepath.Join(tmp, "schema.json")
|
||||
referencePath := filepath.Join(tmp, "value.json")
|
||||
if err := os.WriteFile(rootPath, rootSchema, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(referencePath, []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"type": "integer",
|
||||
"minimum": 2
|
||||
}`), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
validator := NewStandardValidator(tmp)
|
||||
preparer, ok := validator.(ValidationPreparer)
|
||||
if !ok {
|
||||
t.Fatal("standard validator does not support validation preparation")
|
||||
}
|
||||
prepared, err := preparer.PrepareValidation(context.Background(), domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare validation: %v", err)
|
||||
}
|
||||
assertSchemaDocument(t, prepared.SchemaDocument(), rootSchema)
|
||||
|
||||
if err := os.Remove(rootPath); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.Remove(referencePath); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
valid, err := prepared.Validate(context.Background(), &domain.Artifact{Body: []byte(`{"value":3}`)})
|
||||
if err != nil {
|
||||
t.Fatalf("validate prepared artifact: %v", err)
|
||||
}
|
||||
if valid.Status != domain.ValidationPassed || !valid.IsValid {
|
||||
t.Fatalf("expected passed/valid, got status=%q valid=%v errors=%v", valid.Status, valid.IsValid, valid.Errors)
|
||||
}
|
||||
|
||||
invalid, err := prepared.Validate(context.Background(), &domain.Artifact{Body: []byte(`{"value":"changed"}`)})
|
||||
if err != nil {
|
||||
t.Fatalf("validate prepared artifact: %v", err)
|
||||
}
|
||||
if invalid.Status != domain.ValidationFailed || invalid.IsValid {
|
||||
t.Fatalf("expected failed/invalid, got status=%q valid=%v", invalid.Status, invalid.IsValid)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorJSONSchemaNestedSchemaPathSuccess(t *testing.T) {
|
||||
tmp := t.TempDir()
|
||||
nestedDir := filepath.Join(tmp, "dnd")
|
||||
@@ -222,55 +332,6 @@ func TestStandardValidatorJSONSchemaCompilationError(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorLoadSchemaDocumentSuccess(t *testing.T) {
|
||||
tmp := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(tmp, "schema.json"), []byte(`{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"}
|
||||
}
|
||||
}`), 0644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
v := NewStandardValidator(tmp)
|
||||
loader, ok := v.(SchemaDocumentLoader)
|
||||
if !ok {
|
||||
t.Fatal("standard validator must implement SchemaDocumentLoader")
|
||||
}
|
||||
|
||||
doc, err := loader.LoadSchemaDocument(context.Background(), "schema.json")
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
|
||||
obj, ok := doc.(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("expected object document, got %#v", doc)
|
||||
}
|
||||
if obj["type"] != "object" {
|
||||
t.Fatalf("expected schema type=object, got %#v", obj["type"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorLoadSchemaDocumentInvalidJSON(t *testing.T) {
|
||||
tmp := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(tmp, "schema.json"), []byte(`{`), 0644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
v := NewStandardValidator(tmp)
|
||||
loader, ok := v.(SchemaDocumentLoader)
|
||||
if !ok {
|
||||
t.Fatal("standard validator must implement SchemaDocumentLoader")
|
||||
}
|
||||
|
||||
_, err := loader.LoadSchemaDocument(context.Background(), "schema.json")
|
||||
if err == nil {
|
||||
t.Fatal("expected decode error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFSValidatorJSONSchemaSuccess(t *testing.T) {
|
||||
v := NewFSValidator(fstest.MapFS{
|
||||
"schemas/events.schema.json": &fstest.MapFile{Data: []byte(`{
|
||||
@@ -295,17 +356,128 @@ func TestFSValidatorJSONSchemaSuccess(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFSValidatorJSONSchemaRegistrationError(t *testing.T) {
|
||||
v := NewFSValidator(fstest.MapFS{
|
||||
"schemas/%zz.json": &fstest.MapFile{Data: []byte(`{"type":"object"}`)},
|
||||
}, "schemas")
|
||||
|
||||
res, err := v.Validate(context.Background(), &domain.Artifact{Body: []byte(`{}`)}, domain.OutputContract{
|
||||
func TestFSValidatorPreparedSchemaSurvivesSourceMutation(t *testing.T) {
|
||||
rootSchema := []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"title": "original root",
|
||||
"type": "object",
|
||||
"required": ["value"],
|
||||
"properties": {
|
||||
"value": {"$ref": "value.json"}
|
||||
}
|
||||
}`)
|
||||
fsys := fstest.MapFS{
|
||||
"schema.json": &fstest.MapFile{Data: rootSchema},
|
||||
"value.json": &fstest.MapFile{Data: []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"type": "integer",
|
||||
"minimum": 2
|
||||
}`)},
|
||||
}
|
||||
validator := NewFSValidator(fsys, ".")
|
||||
preparer, ok := validator.(ValidationPreparer)
|
||||
if !ok {
|
||||
t.Fatal("filesystem validator does not support validation preparation")
|
||||
}
|
||||
prepared, err := preparer.PrepareValidation(context.Background(), domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "%zz.json",
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "failed to register JSON schema") {
|
||||
t.Fatalf("expected schema registration error, got result=%#v error=%v", res, err)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare validation: %v", err)
|
||||
}
|
||||
assertSchemaDocument(t, prepared.SchemaDocument(), rootSchema)
|
||||
|
||||
fsys["schema.json"] = &fstest.MapFile{Data: []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"type": "string"
|
||||
}`)}
|
||||
fsys["value.json"] = &fstest.MapFile{Data: []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"type": "string"
|
||||
}`)}
|
||||
|
||||
valid, err := prepared.Validate(context.Background(), &domain.Artifact{Body: []byte(`{"value":3}`)})
|
||||
if err != nil {
|
||||
t.Fatalf("validate prepared artifact: %v", err)
|
||||
}
|
||||
if valid.Status != domain.ValidationPassed || !valid.IsValid {
|
||||
t.Fatalf("expected passed/valid, got status=%q valid=%v errors=%v", valid.Status, valid.IsValid, valid.Errors)
|
||||
}
|
||||
|
||||
invalid, err := prepared.Validate(context.Background(), &domain.Artifact{Body: []byte(`{"value":"changed"}`)})
|
||||
if err != nil {
|
||||
t.Fatalf("validate prepared artifact: %v", err)
|
||||
}
|
||||
if invalid.Status != domain.ValidationFailed || invalid.IsValid {
|
||||
t.Fatalf("expected failed/invalid, got status=%q valid=%v", invalid.Status, invalid.IsValid)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFSValidatorEscapesSchemaResourcePath(t *testing.T) {
|
||||
for _, schemaName := range []string{"%zz.json", "space name.json", "hash#.json", "query?.json", "rún.json"} {
|
||||
t.Run(schemaName, func(t *testing.T) {
|
||||
v := NewFSValidator(fstest.MapFS{
|
||||
"schemas/" + schemaName: &fstest.MapFile{Data: []byte(`{"type":"object"}`)},
|
||||
}, "schemas")
|
||||
|
||||
res, err := v.Validate(context.Background(), &domain.Artifact{Body: []byte(`{}`)}, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: schemaName,
|
||||
})
|
||||
if err != nil || !res.IsValid {
|
||||
t.Fatalf("validate schema %q: result=%#v error=%v", schemaName, res, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaReferencesPreserveEscapedFilenames(t *testing.T) {
|
||||
for _, name := range []string{"%2F.json", "space name.json", "hash#.json", "query?.json", "rún.json"} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rootName := "root-" + name
|
||||
childName := "child-" + name
|
||||
childPath := "nested/" + childName
|
||||
reference := (&url.URL{Path: childPath}).EscapedPath()
|
||||
rootSchema := []byte(`{"$ref":` + strconv.Quote(reference) + `}`)
|
||||
childSchema := []byte(`{"type":"integer","minimum":2}`)
|
||||
|
||||
t.Run("fs.FS", func(t *testing.T) {
|
||||
validator := NewFSValidator(fstest.MapFS{
|
||||
"schemas/" + rootName: &fstest.MapFile{Data: rootSchema},
|
||||
"schemas/" + childPath: &fstest.MapFile{Data: childSchema},
|
||||
}, "schemas")
|
||||
assertSchemaValidation(t, validator, rootName)
|
||||
})
|
||||
|
||||
t.Run("operating system files", func(t *testing.T) {
|
||||
if runtime.GOOS == "windows" && strings.ContainsAny(rootName+childName, `<>:"/\|?*`) {
|
||||
t.Skip("filename is not legal on Windows")
|
||||
}
|
||||
root := t.TempDir()
|
||||
if err := os.Mkdir(filepath.Join(root, "nested"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(root, rootName), rootSchema, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(root, filepath.FromSlash(childPath)), childSchema, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertSchemaValidation(t, NewStandardValidator(root), rootName)
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func assertSchemaValidation(t *testing.T, validator Validator, schemaPath string) {
|
||||
t.Helper()
|
||||
result, err := validator.Validate(context.Background(), &domain.Artifact{Body: []byte(`2`)}, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: schemaPath,
|
||||
})
|
||||
if err != nil || !result.IsValid {
|
||||
t.Fatalf("validate schema %q: result=%#v error=%v", schemaPath, result, err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -395,25 +567,6 @@ func TestFSValidatorSingleSchemaFileUsesBaseName(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFSValidatorLoadSchemaDocument(t *testing.T) {
|
||||
v := NewFSValidator(fstest.MapFS{
|
||||
"schemas/schema.json": &fstest.MapFile{Data: []byte(`{"type":"object"}`)},
|
||||
}, "schemas")
|
||||
loader, ok := v.(SchemaDocumentLoader)
|
||||
if !ok {
|
||||
t.Fatal("fs validator must implement SchemaDocumentLoader")
|
||||
}
|
||||
|
||||
doc, err := loader.LoadSchemaDocument(context.Background(), "schema.json")
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
obj, ok := doc.(map[string]any)
|
||||
if !ok || obj["type"] != "object" {
|
||||
t.Fatalf("unexpected schema document: %#v", doc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStandardValidatorJSONSchemaReferenceBoundaries(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(root, "child.json"), []byte(`{
|
||||
@@ -508,6 +661,225 @@ func TestFSValidatorJSONSchemaReferenceBoundaries(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestJSONSchemaNumericConstraintsRetainJSONPrecision(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
schema string
|
||||
instance string
|
||||
wantValid bool
|
||||
}{
|
||||
{
|
||||
name: "const distinguishes adjacent large integers",
|
||||
schema: `{"const":9007199254740993}`,
|
||||
instance: `9007199254740993`,
|
||||
wantValid: true,
|
||||
},
|
||||
{
|
||||
name: "const rejects adjacent large integer",
|
||||
schema: `{"const":9007199254740993}`,
|
||||
instance: `9007199254740992`,
|
||||
wantValid: false,
|
||||
},
|
||||
{
|
||||
name: "const accepts exponent beyond float range",
|
||||
schema: `{"const":1e400}`,
|
||||
instance: `1e400`,
|
||||
wantValid: true,
|
||||
},
|
||||
{
|
||||
name: "minimum accepts precise decimal boundary",
|
||||
schema: `{"type":"number","minimum":0.123456789012345678901234567890}`,
|
||||
instance: `0.123456789012345678901234567890`,
|
||||
wantValid: true,
|
||||
},
|
||||
{
|
||||
name: "minimum rejects lower precise decimal",
|
||||
schema: `{"type":"number","minimum":0.123456789012345678901234567890}`,
|
||||
instance: `0.123456789012345678901234567889`,
|
||||
wantValid: false,
|
||||
},
|
||||
{
|
||||
name: "maximum distinguishes adjacent large integers",
|
||||
schema: `{"type":"number","maximum":9007199254740992}`,
|
||||
instance: `9007199254740993`,
|
||||
wantValid: false,
|
||||
},
|
||||
{
|
||||
name: "multiple of accepts exact decimal multiple",
|
||||
schema: `{"type":"number","multipleOf":0.0000000000000000001}`,
|
||||
instance: `0.0000000000000000003`,
|
||||
wantValid: true,
|
||||
},
|
||||
{
|
||||
name: "multiple of rejects inexact decimal multiple",
|
||||
schema: `{"type":"number","multipleOf":0.0000000000000000001}`,
|
||||
instance: `0.00000000000000000031`,
|
||||
wantValid: false,
|
||||
},
|
||||
{
|
||||
name: "ordinary number remains supported",
|
||||
schema: `{"type":"number","minimum":1,"maximum":3}`,
|
||||
instance: `2`,
|
||||
wantValid: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, source := range jsonSchemaValidatorSources() {
|
||||
t.Run(source.name, func(t *testing.T) {
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
validator := source.new(t, []byte(tc.schema))
|
||||
result, err := validator.Validate(context.Background(), &domain.Artifact{Body: []byte(tc.instance)}, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("validate JSON Schema instance: %v", err)
|
||||
}
|
||||
if result.IsValid != tc.wantValid {
|
||||
t.Fatalf("valid = %v, want %v; result=%+v", result.IsValid, tc.wantValid, result)
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestJSONSchemaDecodingRequiresOneCompleteDocument(t *testing.T) {
|
||||
for _, source := range jsonSchemaValidatorSources() {
|
||||
t.Run(source.name, func(t *testing.T) {
|
||||
for _, schema := range []string{`{"type":`, `{} {}`} {
|
||||
validator := source.new(t, []byte(schema))
|
||||
_, err := validator.Validate(context.Background(), &domain.Artifact{Body: []byte(`1`)}, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatalf("schema %q: expected decoding error", schema)
|
||||
}
|
||||
}
|
||||
|
||||
validator := source.new(t, []byte(`{}`))
|
||||
for _, instance := range []string{`{"value":`, `1 2`} {
|
||||
result, err := validator.Validate(context.Background(), &domain.Artifact{Body: []byte(instance)}, domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("instance %q: expected completed validation, got %v", instance, err)
|
||||
}
|
||||
if result.Status != domain.ValidationFailed || result.IsValid || len(result.Errors) == 0 {
|
||||
t.Fatalf("instance %q: expected failed validation, got %+v", instance, result)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedSchemaDocumentsRetainExactNumbers(t *testing.T) {
|
||||
const schema = `{
|
||||
"const": 9007199254740993,
|
||||
"minimum": 0.123456789012345678901234567890,
|
||||
"maximum": 1e400,
|
||||
"multipleOf": 0.0000000000000000001
|
||||
}`
|
||||
want := map[string]string{
|
||||
"const": "9007199254740993",
|
||||
"minimum": "0.123456789012345678901234567890",
|
||||
"maximum": "1e400",
|
||||
"multipleOf": "0.0000000000000000001",
|
||||
}
|
||||
|
||||
for _, source := range jsonSchemaValidatorSources() {
|
||||
t.Run(source.name, func(t *testing.T) {
|
||||
validator := source.new(t, []byte(schema))
|
||||
preparer, ok := validator.(ValidationPreparer)
|
||||
if !ok {
|
||||
t.Fatal("validator does not support validation preparation")
|
||||
}
|
||||
prepared, err := preparer.PrepareValidation(context.Background(), domain.OutputContract{
|
||||
ValidationMode: domain.ValidationJSONSchema,
|
||||
SchemaPath: "schema.json",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare validation: %v", err)
|
||||
}
|
||||
document, ok := prepared.SchemaDocument().(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("schema document = %#v, want object", prepared.SchemaDocument())
|
||||
}
|
||||
for name, wantNumber := range want {
|
||||
got, ok := document[name].(json.Number)
|
||||
if !ok {
|
||||
t.Fatalf("schema field %q = %#v, want json.Number", name, document[name])
|
||||
}
|
||||
if got.String() != wantNumber {
|
||||
t.Fatalf("schema field %q = %q, want %q", name, got, wantNumber)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type jsonSchemaValidatorSource struct {
|
||||
name string
|
||||
new func(*testing.T, []byte) Validator
|
||||
}
|
||||
|
||||
func jsonSchemaValidatorSources() []jsonSchemaValidatorSource {
|
||||
return []jsonSchemaValidatorSource{
|
||||
{
|
||||
name: "operating system files",
|
||||
new: func(t *testing.T, schema []byte) Validator {
|
||||
t.Helper()
|
||||
root := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(root, "schema.json"), schema, 0o644); err != nil {
|
||||
t.Fatalf("write schema: %v", err)
|
||||
}
|
||||
return NewStandardValidator(root)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "fs.FS",
|
||||
new: func(t *testing.T, schema []byte) Validator {
|
||||
t.Helper()
|
||||
return NewFSValidator(fstest.MapFS{
|
||||
"schema.json": &fstest.MapFile{Data: schema},
|
||||
}, ".")
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkJSONValidation(b *testing.B) {
|
||||
largeArray := []byte(`[` + strings.Repeat(`12345678901234567890,`, 32*1024) + `0]`)
|
||||
benchmarks := []struct {
|
||||
name string
|
||||
body []byte
|
||||
}{
|
||||
{name: "scalar", body: []byte(`1e400`)},
|
||||
{name: "object", body: []byte(`{"name":"eris","count":9007199254740993,"enabled":true}`)},
|
||||
{name: "large array", body: largeArray},
|
||||
}
|
||||
validator := NewStandardValidator("")
|
||||
contract := domain.OutputContract{ValidationMode: domain.ValidationJSON}
|
||||
|
||||
for _, benchmark := range benchmarks {
|
||||
b.Run(benchmark.name, func(b *testing.B) {
|
||||
artifact := &domain.Artifact{Body: benchmark.body}
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(benchmark.body)))
|
||||
b.ResetTimer()
|
||||
for range b.N {
|
||||
result, err := validator.Validate(context.Background(), artifact, contract)
|
||||
if err != nil || !result.IsValid {
|
||||
b.Fatalf("validate JSON: result=%+v error=%v", result, err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestJSONSchemaDialectIsDraft2020(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -548,3 +920,15 @@ func TestJSONSchemaDialectIsDraft2020(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func assertSchemaDocument(t *testing.T, got any, expectedJSON []byte) {
|
||||
t.Helper()
|
||||
|
||||
expected, err := decodeJSONValue(context.Background(), expectedJSON)
|
||||
if err != nil {
|
||||
t.Fatalf("decode expected schema document: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(got, expected) {
|
||||
t.Fatalf("schema document mismatch:\n got: %#v\nwant: %#v", got, expected)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,15 +2,33 @@ package validate
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
// Validator validates the generated artifact based on the output contract.
|
||||
// Validation checks ctx around Promptkit-controlled work and synchronous
|
||||
// dependency calls. A dependency call already in progress cannot be preempted;
|
||||
// after it returns, cancellation takes precedence over its result.
|
||||
type Validator interface {
|
||||
Validate(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract) (domain.ValidationResult, error)
|
||||
}
|
||||
|
||||
// SchemaDocumentLoader loads JSON schema documents using validator path semantics.
|
||||
type SchemaDocumentLoader interface {
|
||||
LoadSchemaDocument(ctx context.Context, schemaPath string) (any, error)
|
||||
// PreparedValidation validates artifacts against one frozen output contract.
|
||||
// Its cancellation boundary is synchronous: Validate does not detach schema
|
||||
// execution, and an observed context error prevents publication of a result.
|
||||
type PreparedValidation interface {
|
||||
Validate(ctx context.Context, artifact *domain.Artifact) (domain.ValidationResult, error)
|
||||
// SchemaDocument returns the root JSON Schema document used for provider
|
||||
// structured output, or nil for non-schema modes. Returned internal
|
||||
// immutable state must not be mutated.
|
||||
SchemaDocument() any
|
||||
}
|
||||
|
||||
// ValidationPreparer freezes validation resources for one output contract.
|
||||
// Preparation reads schemas in context-checked chunks and checks ctx around
|
||||
// decoding and compilation. Filesystem and compiler calls remain synchronous,
|
||||
// so cancellation becomes authoritative when an in-progress call returns.
|
||||
type ValidationPreparer interface {
|
||||
PrepareValidation(ctx context.Context, contract domain.OutputContract) (PreparedValidation, error)
|
||||
}
|
||||
|
||||
142
json.go
142
json.go
@@ -2,9 +2,16 @@ package promptkit
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
minDurationMilliseconds = int64(time.Duration(math.MinInt64) / time.Millisecond)
|
||||
maxDurationMilliseconds = int64(time.Duration(math.MaxInt64) / time.Millisecond)
|
||||
)
|
||||
|
||||
// MarshalJSON implements json.Marshaler for PreparedRun. It uses RFC 3339
|
||||
// timestamps, integer duration_ms, and omits zero timing values.
|
||||
func (r PreparedRun) MarshalJSON() ([]byte, error) {
|
||||
@@ -21,38 +28,11 @@ func (r PreparedRun) MarshalJSON() ([]byte, error) {
|
||||
durationMS = &r.DurationMS
|
||||
}
|
||||
|
||||
return json.Marshal(struct {
|
||||
PromptID string `json:"prompt_id"`
|
||||
PromptVersion string `json:"prompt_version,omitempty"`
|
||||
PromptHash string `json:"prompt_hash,omitempty"`
|
||||
SelectedProfileID string `json:"selected_profile_id"`
|
||||
SelectedBackendID string `json:"selected_backend_id,omitempty"`
|
||||
EffectiveModelParams ExecutionTarget `json:"effective_model_params"`
|
||||
OutputContract OutputContract `json:"output_contract"`
|
||||
StructuredOutput *StructuredOutputSpec `json:"structured_output,omitempty"`
|
||||
InputHashes map[string]string `json:"input_hashes,omitempty"`
|
||||
SessionID string `json:"session_id,omitempty"`
|
||||
RenderedPromptHash string `json:"rendered_prompt_hash"`
|
||||
Messages []RenderedMessage `json:"messages"`
|
||||
StartTime *time.Time `json:"start_time,omitempty"`
|
||||
EndTime *time.Time `json:"end_time,omitempty"`
|
||||
DurationMS *int64 `json:"duration_ms,omitempty"`
|
||||
}{
|
||||
PromptID: r.PromptID,
|
||||
PromptVersion: r.PromptVersion,
|
||||
PromptHash: r.PromptHash,
|
||||
SelectedProfileID: r.SelectedProfileID,
|
||||
SelectedBackendID: r.SelectedBackendID,
|
||||
EffectiveModelParams: r.EffectiveModelParams,
|
||||
OutputContract: r.OutputContract,
|
||||
StructuredOutput: r.StructuredOutput,
|
||||
InputHashes: r.InputHashes,
|
||||
SessionID: r.SessionID,
|
||||
RenderedPromptHash: r.RenderedPromptHash,
|
||||
Messages: r.Messages,
|
||||
StartTime: startTime,
|
||||
EndTime: endTime,
|
||||
DurationMS: durationMS,
|
||||
return json.Marshal(preparedRunJSON{
|
||||
preparedRunJSONFields: preparedRunJSONFields(r),
|
||||
StartTime: startTime,
|
||||
EndTime: endTime,
|
||||
DurationMS: durationMS,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -74,89 +54,57 @@ func (r RunResult) MarshalJSON() ([]byte, error) {
|
||||
}
|
||||
|
||||
return json.Marshal(runResultJSON{
|
||||
RunID: r.RunID,
|
||||
Artifact: r.Artifact,
|
||||
RawOutput: r.RawOutput,
|
||||
Validation: r.Validation,
|
||||
PromptID: r.PromptID,
|
||||
PromptVersion: r.PromptVersion,
|
||||
PromptHash: r.PromptHash,
|
||||
SessionID: r.SessionID,
|
||||
RenderedPromptHash: r.RenderedPromptHash,
|
||||
SelectedProfileID: r.SelectedProfileID,
|
||||
SelectedBackendID: r.SelectedBackendID,
|
||||
ModelName: r.ModelName,
|
||||
Endpoint: r.Endpoint,
|
||||
EffectiveModelParams: r.EffectiveModelParams,
|
||||
InputHashes: r.InputHashes,
|
||||
Usage: r.Usage,
|
||||
StartTime: startTime,
|
||||
EndTime: endTime,
|
||||
DurationMS: durationMS,
|
||||
runResultJSONFields: runResultJSONFields(r),
|
||||
StartTime: startTime,
|
||||
EndTime: endTime,
|
||||
DurationMS: durationMS,
|
||||
})
|
||||
}
|
||||
|
||||
// UnmarshalJSON implements json.Unmarshaler for RunResult. It decodes
|
||||
// duration_ms into Duration with millisecond precision.
|
||||
// duration_ms into Duration with millisecond precision. A duration_ms outside
|
||||
// the range representable by time.Duration returns an error without changing
|
||||
// the receiver.
|
||||
func (r *RunResult) UnmarshalJSON(data []byte) error {
|
||||
var wire runResultJSON
|
||||
if err := json.Unmarshal(data, &wire); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*r = RunResult{
|
||||
RunID: wire.RunID,
|
||||
Artifact: wire.Artifact,
|
||||
RawOutput: wire.RawOutput,
|
||||
Validation: wire.Validation,
|
||||
PromptID: wire.PromptID,
|
||||
PromptVersion: wire.PromptVersion,
|
||||
PromptHash: wire.PromptHash,
|
||||
SessionID: wire.SessionID,
|
||||
RenderedPromptHash: wire.RenderedPromptHash,
|
||||
SelectedProfileID: wire.SelectedProfileID,
|
||||
SelectedBackendID: wire.SelectedBackendID,
|
||||
ModelName: wire.ModelName,
|
||||
Endpoint: wire.Endpoint,
|
||||
EffectiveModelParams: wire.EffectiveModelParams,
|
||||
InputHashes: wire.InputHashes,
|
||||
Usage: wire.Usage,
|
||||
Duration: time.Duration(valueOrZero(wire.DurationMS)) * time.Millisecond,
|
||||
result := RunResult(wire.runResultJSONFields)
|
||||
if wire.DurationMS != nil {
|
||||
if *wire.DurationMS < minDurationMilliseconds || *wire.DurationMS > maxDurationMilliseconds {
|
||||
return fmt.Errorf(
|
||||
"decode RunResult duration_ms: %d cannot be represented as time.Duration",
|
||||
*wire.DurationMS,
|
||||
)
|
||||
}
|
||||
result.Duration = time.Duration(*wire.DurationMS) * time.Millisecond
|
||||
}
|
||||
if wire.StartTime != nil {
|
||||
r.StartTime = *wire.StartTime
|
||||
result.StartTime = *wire.StartTime
|
||||
}
|
||||
if wire.EndTime != nil {
|
||||
r.EndTime = *wire.EndTime
|
||||
result.EndTime = *wire.EndTime
|
||||
}
|
||||
*r = result
|
||||
return nil
|
||||
}
|
||||
|
||||
type runResultJSON struct {
|
||||
RunID string `json:"run_id"`
|
||||
Artifact Artifact `json:"artifact"`
|
||||
RawOutput string `json:"raw_output"`
|
||||
Validation ValidationResult `json:"validation"`
|
||||
PromptID string `json:"prompt_id"`
|
||||
PromptVersion string `json:"prompt_version,omitempty"`
|
||||
PromptHash string `json:"prompt_hash,omitempty"`
|
||||
SessionID string `json:"session_id,omitempty"`
|
||||
RenderedPromptHash string `json:"rendered_prompt_hash"`
|
||||
SelectedProfileID string `json:"selected_profile_id"`
|
||||
SelectedBackendID string `json:"selected_backend_id,omitempty"`
|
||||
ModelName string `json:"model_name"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
EffectiveModelParams ExecutionTarget `json:"effective_model_params"`
|
||||
InputHashes map[string]string `json:"input_hashes,omitempty"`
|
||||
Usage TokenUsage `json:"usage"`
|
||||
StartTime *time.Time `json:"start_time,omitempty"`
|
||||
EndTime *time.Time `json:"end_time,omitempty"`
|
||||
DurationMS *int64 `json:"duration_ms,omitempty"`
|
||||
type preparedRunJSONFields PreparedRun
|
||||
|
||||
type preparedRunJSON struct {
|
||||
preparedRunJSONFields
|
||||
StartTime *time.Time `json:"start_time,omitempty"`
|
||||
EndTime *time.Time `json:"end_time,omitempty"`
|
||||
DurationMS *int64 `json:"duration_ms,omitempty"`
|
||||
}
|
||||
|
||||
func valueOrZero(value *int64) int64 {
|
||||
if value == nil {
|
||||
return 0
|
||||
}
|
||||
return *value
|
||||
type runResultJSONFields RunResult
|
||||
|
||||
type runResultJSON struct {
|
||||
runResultJSONFields
|
||||
StartTime *time.Time `json:"start_time,omitempty"`
|
||||
EndTime *time.Time `json:"end_time,omitempty"`
|
||||
DurationMS *int64 `json:"duration_ms,omitempty"`
|
||||
}
|
||||
|
||||
328
json_contract_test.go
Normal file
328
json_contract_test.go
Normal file
@@ -0,0 +1,328 @@
|
||||
package promptkit_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"math"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit"
|
||||
)
|
||||
|
||||
func TestPreparedRunJSONContractRoundTripsAllFields(t *testing.T) {
|
||||
value := fullyPopulatedPreparedRun()
|
||||
payload, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal PreparedRun: %v", err)
|
||||
}
|
||||
|
||||
object := decodeJSONObject(t, payload)
|
||||
assertJSONFields(t, object,
|
||||
"prompt_id",
|
||||
"prompt_version",
|
||||
"prompt_hash",
|
||||
"selected_profile_id",
|
||||
"selected_backend_id",
|
||||
"effective_model_params",
|
||||
"output_contract",
|
||||
"structured_output",
|
||||
"input_hashes",
|
||||
"session_id",
|
||||
"rendered_prompt_hash",
|
||||
"messages",
|
||||
"start_time",
|
||||
"end_time",
|
||||
"duration_ms",
|
||||
)
|
||||
|
||||
var messages []map[string]json.RawMessage
|
||||
if err := json.Unmarshal(object["messages"], &messages); err != nil {
|
||||
t.Fatalf("decode PreparedRun messages: %v", err)
|
||||
}
|
||||
if len(messages) != 2 {
|
||||
t.Fatalf("message count = %d, want 2", len(messages))
|
||||
}
|
||||
if _, ok := messages[0]["cache_control"]; !ok {
|
||||
t.Fatalf("first message omitted cache_control: %s", object["messages"])
|
||||
}
|
||||
if _, ok := messages[1]["cache_control"]; ok {
|
||||
t.Fatalf("second message included empty cache_control: %s", object["messages"])
|
||||
}
|
||||
|
||||
var decoded promptkit.PreparedRun
|
||||
if err := json.Unmarshal(payload, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal PreparedRun: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(decoded, value) {
|
||||
t.Fatalf("PreparedRun did not round trip:\ngot %#v\nwant %#v", decoded, value)
|
||||
}
|
||||
|
||||
zeroPayload, err := json.Marshal(promptkit.PreparedRun{})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal zero PreparedRun: %v", err)
|
||||
}
|
||||
zeroObject := decodeJSONObject(t, zeroPayload)
|
||||
for _, field := range []string{
|
||||
"prompt_version",
|
||||
"prompt_hash",
|
||||
"selected_backend_id",
|
||||
"structured_output",
|
||||
"input_hashes",
|
||||
"session_id",
|
||||
"start_time",
|
||||
"end_time",
|
||||
"duration_ms",
|
||||
} {
|
||||
if _, ok := zeroObject[field]; ok {
|
||||
t.Fatalf("zero PreparedRun included %q: %s", field, zeroPayload)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunResultJSONContractRoundTripsAllFields(t *testing.T) {
|
||||
value := fullyPopulatedRunResult()
|
||||
payload, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal RunResult: %v", err)
|
||||
}
|
||||
|
||||
object := decodeJSONObject(t, payload)
|
||||
assertJSONFields(t, object,
|
||||
"run_id",
|
||||
"artifact",
|
||||
"raw_output",
|
||||
"validation",
|
||||
"prompt_id",
|
||||
"prompt_version",
|
||||
"prompt_hash",
|
||||
"session_id",
|
||||
"rendered_prompt_hash",
|
||||
"selected_profile_id",
|
||||
"selected_backend_id",
|
||||
"model_name",
|
||||
"endpoint",
|
||||
"effective_model_params",
|
||||
"input_hashes",
|
||||
"usage",
|
||||
"start_time",
|
||||
"end_time",
|
||||
"duration_ms",
|
||||
)
|
||||
if _, ok := object["duration"]; ok {
|
||||
t.Fatalf("RunResult included nanosecond duration field: %s", payload)
|
||||
}
|
||||
|
||||
var decoded promptkit.RunResult
|
||||
if err := json.Unmarshal(payload, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal RunResult: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(decoded, value) {
|
||||
t.Fatalf("RunResult did not round trip:\ngot %#v\nwant %#v", decoded, value)
|
||||
}
|
||||
|
||||
zeroPayload, err := json.Marshal(promptkit.RunResult{})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal zero RunResult: %v", err)
|
||||
}
|
||||
zeroObject := decodeJSONObject(t, zeroPayload)
|
||||
for _, field := range []string{
|
||||
"prompt_version",
|
||||
"prompt_hash",
|
||||
"session_id",
|
||||
"selected_backend_id",
|
||||
"input_hashes",
|
||||
"start_time",
|
||||
"end_time",
|
||||
"duration_ms",
|
||||
"duration",
|
||||
} {
|
||||
if _, ok := zeroObject[field]; ok {
|
||||
t.Fatalf("zero RunResult included %q: %s", field, zeroPayload)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunResultJSONDurationMillisecondBoundaries(t *testing.T) {
|
||||
maxMilliseconds := int64(time.Duration(math.MaxInt64) / time.Millisecond)
|
||||
minMilliseconds := int64(time.Duration(math.MinInt64) / time.Millisecond)
|
||||
|
||||
for _, milliseconds := range []int64{minMilliseconds, maxMilliseconds} {
|
||||
t.Run(strconv.FormatInt(milliseconds, 10), func(t *testing.T) {
|
||||
var decoded promptkit.RunResult
|
||||
if err := json.Unmarshal(durationPayload(milliseconds), &decoded); err != nil {
|
||||
t.Fatalf("decode representable duration_ms: %v", err)
|
||||
}
|
||||
want := time.Duration(milliseconds) * time.Millisecond
|
||||
if decoded.Duration != want {
|
||||
t.Fatalf("Duration = %v, want %v", decoded.Duration, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
for _, milliseconds := range []int64{
|
||||
minMilliseconds - 1,
|
||||
maxMilliseconds + 1,
|
||||
math.MinInt64,
|
||||
math.MaxInt64,
|
||||
} {
|
||||
t.Run(strconv.FormatInt(milliseconds, 10), func(t *testing.T) {
|
||||
original := fullyPopulatedRunResult()
|
||||
decoded := original
|
||||
err := json.Unmarshal(durationPayload(milliseconds), &decoded)
|
||||
if err == nil || !strings.Contains(err.Error(), "duration_ms") || !strings.Contains(err.Error(), "time.Duration") {
|
||||
t.Fatalf("overflow error = %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(decoded, original) {
|
||||
t.Fatalf("failed decode partially updated receiver:\ngot %#v\nwant %#v", decoded, original)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func fullyPopulatedPreparedRun() promptkit.PreparedRun {
|
||||
start := time.Date(2026, time.August, 11, 12, 13, 14, 150_000_000, time.UTC)
|
||||
return promptkit.PreparedRun{
|
||||
PromptID: "prompt.prepared",
|
||||
PromptVersion: "2.1.0",
|
||||
PromptHash: "prompt-hash",
|
||||
SelectedProfileID: "profile-prepared",
|
||||
SelectedBackendID: "backend-prepared",
|
||||
EffectiveModelParams: jsonContractExecutionTarget(),
|
||||
OutputContract: promptkit.OutputContract{
|
||||
Format: promptkit.FormatJSON,
|
||||
ValidationMode: promptkit.ValidationJSONSchema,
|
||||
SchemaPath: "schemas/prepared.json",
|
||||
RepairAttempts: 2,
|
||||
},
|
||||
StructuredOutput: &promptkit.StructuredOutputSpec{
|
||||
Type: promptkit.StructuredOutputJSONSchema,
|
||||
JSONSchema: &promptkit.StructuredOutputJSONSpec{
|
||||
Name: "prepared_schema",
|
||||
Strict: true,
|
||||
Schema: map[string]any{
|
||||
"type": "object",
|
||||
"required": []any{"value"},
|
||||
"properties": map[string]any{
|
||||
"value": map[string]any{"minimum": float64(1)},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
InputHashes: map[string]string{"first": "hash-1", "second": "hash-2"},
|
||||
SessionID: "session-prepared",
|
||||
RenderedPromptHash: "rendered-hash",
|
||||
Messages: []promptkit.RenderedMessage{
|
||||
{
|
||||
Role: "system",
|
||||
Content: "System message",
|
||||
CacheControl: &promptkit.CacheControl{
|
||||
Type: promptkit.CacheControlEphemeral,
|
||||
TTL: "1h",
|
||||
},
|
||||
},
|
||||
{Role: "user", Content: "User message"},
|
||||
},
|
||||
StartTime: start,
|
||||
EndTime: start.Add(1501 * time.Millisecond),
|
||||
DurationMS: 1501,
|
||||
}
|
||||
}
|
||||
|
||||
func fullyPopulatedRunResult() promptkit.RunResult {
|
||||
start := time.Date(2026, time.August, 11, 15, 16, 17, 250_000_000, time.UTC)
|
||||
return promptkit.RunResult{
|
||||
RunID: "run-id",
|
||||
Artifact: promptkit.Artifact{
|
||||
Name: "result.json",
|
||||
ContentType: "application/json",
|
||||
Body: []byte(`{"value":2}`),
|
||||
URI: "memory://result.json",
|
||||
Size: 11,
|
||||
Hash: "artifact-hash",
|
||||
},
|
||||
RawOutput: `{"value":2}`,
|
||||
Validation: promptkit.ValidationResult{
|
||||
Status: promptkit.ValidationFailed,
|
||||
Mode: promptkit.ValidationJSONSchema,
|
||||
Errors: []string{"first error", "second error"},
|
||||
SchemaPath: "schemas/result.json",
|
||||
RepairAttempts: 2,
|
||||
IsValid: false,
|
||||
},
|
||||
PromptID: "prompt.result",
|
||||
PromptVersion: "3.2.1",
|
||||
PromptHash: "result-prompt-hash",
|
||||
SessionID: "session-result",
|
||||
RenderedPromptHash: "result-rendered-hash",
|
||||
SelectedProfileID: "profile-result",
|
||||
SelectedBackendID: "backend-result",
|
||||
ModelName: "model-result",
|
||||
Endpoint: "https://result.example/v1",
|
||||
EffectiveModelParams: jsonContractExecutionTarget(),
|
||||
InputHashes: map[string]string{"input": "input-hash"},
|
||||
Usage: promptkit.TokenUsage{
|
||||
PromptTokens: 101,
|
||||
CompletionTokens: 202,
|
||||
TotalTokens: 303,
|
||||
CachedTokens: 44,
|
||||
CacheWriteTokens: 55,
|
||||
},
|
||||
StartTime: start,
|
||||
EndTime: start.Add(1750 * time.Millisecond),
|
||||
Duration: 1750 * time.Millisecond,
|
||||
}
|
||||
}
|
||||
|
||||
func jsonContractExecutionTarget() promptkit.ExecutionTarget {
|
||||
return promptkit.ExecutionTarget{
|
||||
BackendID: "backend-target",
|
||||
Endpoint: "https://target.example/v1",
|
||||
Model: "model-target",
|
||||
Temperature: 0.75,
|
||||
MaxTokens: 321,
|
||||
TopP: 0.875,
|
||||
TimeoutSeconds: 43,
|
||||
ServiceTier: "priority",
|
||||
ReasoningEffort: "high",
|
||||
APIKeyEnv: "PROMPTKIT_JSON_CONTRACT_KEY",
|
||||
ExtraParams: map[string]any{
|
||||
"enabled": true,
|
||||
"weight": float64(1.25),
|
||||
"nested": map[string]any{"name": "value"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func durationPayload(milliseconds int64) []byte {
|
||||
return []byte(`{"duration_ms":` + strconv.FormatInt(milliseconds, 10) + `}`)
|
||||
}
|
||||
|
||||
func decodeJSONObject(t *testing.T, payload []byte) map[string]json.RawMessage {
|
||||
t.Helper()
|
||||
var object map[string]json.RawMessage
|
||||
if err := json.Unmarshal(payload, &object); err != nil {
|
||||
t.Fatalf("decode JSON object: %v", err)
|
||||
}
|
||||
return object
|
||||
}
|
||||
|
||||
func assertJSONFields(t *testing.T, object map[string]json.RawMessage, fields ...string) {
|
||||
t.Helper()
|
||||
want := make(map[string]struct{}, len(fields))
|
||||
for _, field := range fields {
|
||||
want[field] = struct{}{}
|
||||
}
|
||||
for field := range object {
|
||||
if _, ok := want[field]; !ok {
|
||||
t.Errorf("unexpected JSON field %q", field)
|
||||
}
|
||||
}
|
||||
for field := range want {
|
||||
if _, ok := object[field]; !ok {
|
||||
t.Errorf("missing JSON field %q", field)
|
||||
}
|
||||
}
|
||||
}
|
||||
96
llm_adapter_internal_test.go
Normal file
96
llm_adapter_internal_test.go
Normal file
@@ -0,0 +1,96 @@
|
||||
package promptkit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
||||
)
|
||||
|
||||
func TestPublicLLMClientAdapterGivesClientOwnedNestedValues(t *testing.T) {
|
||||
source := adapterOwnershipRequest()
|
||||
want := adapterOwnershipRequest()
|
||||
client := &retainingMutatingLLMClient{}
|
||||
adapter := publicLLMClientAdapter{client: client}
|
||||
|
||||
if _, err := adapter.Generate(context.Background(), source); err != nil {
|
||||
t.Fatalf("generate: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(source, want) {
|
||||
t.Fatalf("client mutation changed prepared source:\ngot %#v\nwant %#v", source, want)
|
||||
}
|
||||
|
||||
var mutations sync.WaitGroup
|
||||
mutations.Add(1)
|
||||
go func() {
|
||||
defer mutations.Done()
|
||||
for i := 0; i < 10_000; i++ {
|
||||
mutateGenerateRequest(&client.retained, strconv.Itoa(i))
|
||||
}
|
||||
}()
|
||||
for i := 0; i < 10_000; i++ {
|
||||
laterRequest := fromDomainGenerateRequest(source)
|
||||
if laterRequest.Prompt.Messages[0].Content != "source-message" ||
|
||||
laterRequest.Prompt.Messages[0].CacheControl.TTL != "source-ttl" ||
|
||||
laterRequest.Target.ExtraParams["nested"].([]any)[0] != "source-extra" ||
|
||||
laterRequest.StructuredOutput.JSONSchema.Schema.(map[string]any)["enum"].([]any)[0] != "source-schema" {
|
||||
t.Fatal("retained client mutation reached a later execution request")
|
||||
}
|
||||
}
|
||||
mutations.Wait()
|
||||
|
||||
if !reflect.DeepEqual(source, want) {
|
||||
t.Fatalf("retained client mutation changed prepared source:\ngot %#v\nwant %#v", source, want)
|
||||
}
|
||||
}
|
||||
|
||||
type retainingMutatingLLMClient struct {
|
||||
retained GenerateRequest
|
||||
}
|
||||
|
||||
func (c *retainingMutatingLLMClient) Generate(_ context.Context, request GenerateRequest) (*GenerateResponse, error) {
|
||||
c.retained = request
|
||||
mutateGenerateRequest(&c.retained, "client-mutation")
|
||||
return &GenerateResponse{Content: "generated"}, nil
|
||||
}
|
||||
|
||||
func adapterOwnershipRequest() domain.GenerateRequest {
|
||||
return domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{
|
||||
SessionID: "source-session",
|
||||
Messages: []domain.RenderedMessage{{
|
||||
Role: "user",
|
||||
Content: "source-message",
|
||||
CacheControl: &domain.CacheControl{
|
||||
Type: domain.CacheControlEphemeral,
|
||||
TTL: "source-ttl",
|
||||
},
|
||||
}},
|
||||
},
|
||||
Target: domain.ExecutionTarget{
|
||||
Model: "source-model",
|
||||
APIKey: "source-api-key",
|
||||
ExtraParams: map[string]any{
|
||||
"nested": []any{"source-extra"},
|
||||
},
|
||||
},
|
||||
StructuredOutput: &domain.StructuredOutputSpec{
|
||||
Type: domain.StructuredOutputJSONSchema,
|
||||
JSONSchema: &domain.StructuredOutputJSONSpec{
|
||||
Name: "source-schema-name",
|
||||
Strict: true,
|
||||
Schema: map[string]any{"enum": []any{"source-schema"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func mutateGenerateRequest(request *GenerateRequest, value string) {
|
||||
request.Prompt.Messages[0].Content = value
|
||||
request.Prompt.Messages[0].CacheControl.TTL = value
|
||||
request.Target.ExtraParams["nested"].([]any)[0] = value
|
||||
request.StructuredOutput.JSONSchema.Schema.(map[string]any)["enum"].([]any)[0] = value
|
||||
}
|
||||
48
output_contract_contract_test.go
Normal file
48
output_contract_contract_test.go
Normal file
@@ -0,0 +1,48 @@
|
||||
package promptkit_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit"
|
||||
)
|
||||
|
||||
func TestPreparationRejectsInvalidOutputContractWithPublicError(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "profile", "message"), "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "profile",
|
||||
Endpoint: "http://example.test/v1",
|
||||
Model: "model",
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
|
||||
req := promptkit.RunRequest{
|
||||
PromptID: "prompt",
|
||||
Validation: &promptkit.OutputContract{
|
||||
Format: promptkit.OutputFormat("binary"),
|
||||
ValidationMode: promptkit.ValidationNone,
|
||||
},
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
55
prepared_execution.go
Normal file
55
prepared_execution.go
Normal file
@@ -0,0 +1,55 @@
|
||||
package promptkit
|
||||
|
||||
import "gitea.maximumdirect.net/eric/promptkit/internal/usecase"
|
||||
|
||||
const preparedExecutionString = "promptkit.PreparedExecution{opaque}"
|
||||
|
||||
// PreparedExecution is an opaque, in-process handle for one completely
|
||||
// prepared execution. A handle is bound to the [Engine] that created it and
|
||||
// permits one [Engine.RunPrepared] invocation.
|
||||
//
|
||||
// PreparedExecution contains no supported serializable state and cannot be
|
||||
// used as a restartable job. Copying the value preserves the same shared
|
||||
// lifecycle; it does not create another execution attempt.
|
||||
type PreparedExecution struct {
|
||||
internal *usecase.PreparedExecution
|
||||
}
|
||||
|
||||
// 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.
|
||||
//
|
||||
// A nil receiver or zero-value PreparedExecution returns a zero [PreparedRun].
|
||||
func (p *PreparedExecution) Details() PreparedRun {
|
||||
if p == nil || p.internal == nil {
|
||||
return PreparedRun{}
|
||||
}
|
||||
details := fromDomainPreparedRun(p.internal.Details())
|
||||
if details == nil {
|
||||
return PreparedRun{}
|
||||
}
|
||||
return *details
|
||||
}
|
||||
|
||||
// Discard invalidates an unclaimed handle and drops Promptkit's references to
|
||||
// its execution-only state. Discard is nil-safe and idempotent. It does not
|
||||
// cancel an execution that has already claimed the handle; use the
|
||||
// [Engine.RunPrepared] context for cancellation.
|
||||
func (p *PreparedExecution) Discard() {
|
||||
if p == nil || p.internal == nil {
|
||||
return
|
||||
}
|
||||
p.internal.Discard()
|
||||
}
|
||||
|
||||
// String returns a constant representation that exposes no retained request,
|
||||
// rendered content, or credential data.
|
||||
func (p PreparedExecution) String() string {
|
||||
return preparedExecutionString
|
||||
}
|
||||
|
||||
// GoString returns a constant Go-syntax representation that exposes no
|
||||
// retained request, rendered content, or credential data.
|
||||
func (p PreparedExecution) GoString() string {
|
||||
return preparedExecutionString
|
||||
}
|
||||
794
prepared_execution_contract_test.go
Normal file
794
prepared_execution_contract_test.go
Normal file
@@ -0,0 +1,794 @@
|
||||
package promptkit_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit"
|
||||
)
|
||||
|
||||
func TestPreparedExecutionFreezesSourcesAndReturnsIndependentDetails(t *testing.T) {
|
||||
promptSource := preparedPromptSource("original")
|
||||
profileSource := preparedProfileSource("original-model")
|
||||
schemaSource := preparedSchemaSource()
|
||||
reader := &mutablePreparedArtifactReader{
|
||||
body: "original artifact",
|
||||
hash: "original-input-hash",
|
||||
}
|
||||
client := &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: `{"value":3}`},
|
||||
}
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(promptSource, "."),
|
||||
promptkit.WithProfileFS(profileSource, "."),
|
||||
promptkit.WithSchemaFS(schemaSource, "."),
|
||||
promptkit.WithArtifactReader(reader),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
|
||||
temperature := 0.25
|
||||
extraParams := map[string]any{
|
||||
"nested": map[string]any{"source": "original"},
|
||||
}
|
||||
request := promptkit.RunRequest{
|
||||
PromptID: "prepared",
|
||||
Inputs: map[string]promptkit.ArtifactRef{
|
||||
"input": promptkit.Inline("original request input"),
|
||||
},
|
||||
Vars: map[string]string{"label": "original variable"},
|
||||
Execution: &promptkit.ExecutionTargetOverride{
|
||||
Temperature: &temperature,
|
||||
ExtraParams: extraParams,
|
||||
},
|
||||
}
|
||||
preparationContext, cancelPreparation := context.WithCancel(context.Background())
|
||||
prepared, err := engine.PrepareExecution(preparationContext, request)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
cancelPreparation()
|
||||
|
||||
request.PromptID = "changed"
|
||||
request.Inputs["input"] = promptkit.Inline("changed request input")
|
||||
request.Vars["label"] = "changed variable"
|
||||
temperature = 1.5
|
||||
extraParams["nested"].(map[string]any)["source"] = "changed"
|
||||
promptSource["prompt.yaml"] = &fstest.MapFile{Data: []byte(`id: changed`)}
|
||||
profileSource["profile.yaml"] = &fstest.MapFile{Data: []byte(`id: changed`)}
|
||||
schemaSource["schema.json"] = &fstest.MapFile{Data: []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"title": "changed root",
|
||||
"type": "string"
|
||||
}`)}
|
||||
schemaSource["value.json"] = &fstest.MapFile{Data: []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"type": "string"
|
||||
}`)}
|
||||
reader.set("changed artifact", "changed-input-hash")
|
||||
|
||||
first := prepared.Details()
|
||||
first.Messages[0].Content = "changed details"
|
||||
first.InputHashes["input"] = "changed-details-hash"
|
||||
first.EffectiveModelParams.ExtraParams["nested"].(map[string]any)["source"] = "changed details"
|
||||
first.StructuredOutput.JSONSchema.Schema.(map[string]any)["title"] = "changed details"
|
||||
|
||||
second := prepared.Details()
|
||||
if second.Messages[0].Content != "Input=original artifact Label=original variable" {
|
||||
t.Fatalf("details message changed: %q", second.Messages[0].Content)
|
||||
}
|
||||
if second.InputHashes["input"] != "original-input-hash" {
|
||||
t.Fatalf("details input hash changed: %q", second.InputHashes["input"])
|
||||
}
|
||||
if second.EffectiveModelParams.Model != "original-model" ||
|
||||
second.EffectiveModelParams.Temperature != 0.25 ||
|
||||
second.EffectiveModelParams.ExtraParams["nested"].(map[string]any)["source"] != "original" {
|
||||
t.Fatalf("details target changed: %+v", second.EffectiveModelParams)
|
||||
}
|
||||
schema := second.StructuredOutput.JSONSchema.Schema.(map[string]any)
|
||||
if schema["title"] != "original root" {
|
||||
t.Fatalf("details schema changed: %#v", schema)
|
||||
}
|
||||
|
||||
result, err := engine.RunPrepared(context.Background(), prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("run prepared after preparation-context cancellation: %v", err)
|
||||
}
|
||||
if result.Validation.Status != promptkit.ValidationPassed || !result.Validation.IsValid {
|
||||
t.Fatalf("frozen schema did not validate original output: %+v", result.Validation)
|
||||
}
|
||||
if reader.callCount() != 1 {
|
||||
t.Fatalf("execution reopened artifact source: calls=%d", reader.callCount())
|
||||
}
|
||||
|
||||
requests := client.snapshot()
|
||||
if len(requests) != 1 {
|
||||
t.Fatalf("generation calls=%d, want 1", len(requests))
|
||||
}
|
||||
generated := requests[0]
|
||||
if generated.Prompt.Messages[0].Content != second.Messages[0].Content ||
|
||||
generated.Target.Model != second.EffectiveModelParams.Model ||
|
||||
!reflect.DeepEqual(generated.Target.ExtraParams, second.EffectiveModelParams.ExtraParams) ||
|
||||
!reflect.DeepEqual(generated.StructuredOutput, second.StructuredOutput) {
|
||||
t.Fatalf("generation did not use frozen details:\nrequest=%+v\ndetails=%+v", generated, second)
|
||||
}
|
||||
if result.PromptID != second.PromptID ||
|
||||
result.PromptVersion != second.PromptVersion ||
|
||||
result.PromptHash != second.PromptHash ||
|
||||
result.SessionID != second.SessionID ||
|
||||
result.RenderedPromptHash != second.RenderedPromptHash ||
|
||||
result.SelectedProfileID != second.SelectedProfileID ||
|
||||
result.SelectedBackendID != second.SelectedBackendID ||
|
||||
!reflect.DeepEqual(result.EffectiveModelParams, second.EffectiveModelParams) ||
|
||||
!reflect.DeepEqual(result.InputHashes, second.InputHashes) {
|
||||
t.Fatalf("result provenance does not match details:\nresult=%+v\ndetails=%+v", result, second)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedExecutionLifecycleAndEngineBinding(t *testing.T) {
|
||||
ownerClient := &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: "ok"},
|
||||
}
|
||||
owner := newPreparedContractEngine(t, ownerClient, "owner content")
|
||||
foreign := newPreparedContractEngine(t, &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: "unexpected"},
|
||||
}, "foreign content")
|
||||
|
||||
prepared, err := owner.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prepared"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
copied := *prepared
|
||||
|
||||
var nilEngine *promptkit.Engine
|
||||
if result, err := nilEngine.RunPrepared(context.Background(), prepared); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidConfig) {
|
||||
t.Fatalf("nil engine result=(%+v, %v), want ErrInvalidConfig", result, err)
|
||||
}
|
||||
if result, err := foreign.RunPrepared(context.Background(), prepared); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("foreign engine result=(%+v, %v), want ErrInvalidRequest", result, err)
|
||||
}
|
||||
if result, err := owner.RunPrepared(context.Background(), nil); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("nil handle result=(%+v, %v), want ErrInvalidRequest", result, err)
|
||||
}
|
||||
if result, err := owner.RunPrepared(context.Background(), &promptkit.PreparedExecution{}); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("zero handle result=(%+v, %v), want ErrInvalidRequest", result, err)
|
||||
}
|
||||
|
||||
result, err := owner.RunPrepared(context.Background(), &copied)
|
||||
if err != nil || result == nil {
|
||||
t.Fatalf("owner run prepared=(%+v, %v), want success", result, err)
|
||||
}
|
||||
for name, handle := range map[string]*promptkit.PreparedExecution{
|
||||
"original": prepared,
|
||||
"copy": &copied,
|
||||
} {
|
||||
if result, err := owner.RunPrepared(context.Background(), handle); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("%s reused handle result=(%+v, %v), want ErrInvalidRequest", name, result, err)
|
||||
}
|
||||
if handle.Details().PromptID != "prepared" {
|
||||
t.Fatalf("%s details unavailable after execution", name)
|
||||
}
|
||||
}
|
||||
if len(ownerClient.snapshot()) != 1 {
|
||||
t.Fatalf("owner generation calls=%d, want 1", len(ownerClient.snapshot()))
|
||||
}
|
||||
|
||||
collaboratorFailure := errors.New("prepared collaborator failure")
|
||||
failingClient := &preparedRecordingClient{err: collaboratorFailure}
|
||||
failingEngine := newPreparedContractEngine(t, failingClient, "failure content")
|
||||
failing, err := failingEngine.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prepared"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare failing execution: %v", err)
|
||||
}
|
||||
if result, err := failingEngine.RunPrepared(context.Background(), failing); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrLLMGenerate) ||
|
||||
!errors.Is(err, collaboratorFailure) {
|
||||
t.Fatalf("generation failure result=(%+v, %v), want public and collaborator identities", result, err)
|
||||
}
|
||||
if result, err := failingEngine.RunPrepared(context.Background(), failing); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("failed execution was reusable: result=(%+v, %v)", result, err)
|
||||
}
|
||||
|
||||
cancellationRelease := make(chan struct{})
|
||||
cancellationStarted := make(chan struct{}, 1)
|
||||
cancelingEngine := newPreparedContractEngine(t, &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: "unexpected"},
|
||||
started: cancellationStarted,
|
||||
release: cancellationRelease,
|
||||
}, "cancellation content")
|
||||
canceling, err := cancelingEngine.PrepareExecution(
|
||||
context.Background(),
|
||||
promptkit.RunRequest{PromptID: "prepared"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare canceled execution: %v", err)
|
||||
}
|
||||
executionContext, cancelExecution := context.WithCancel(context.Background())
|
||||
type canceledOutcome struct {
|
||||
result *promptkit.RunResult
|
||||
err error
|
||||
}
|
||||
canceledResult := make(chan canceledOutcome, 1)
|
||||
go func() {
|
||||
result, runErr := cancelingEngine.RunPrepared(executionContext, canceling)
|
||||
canceledResult <- canceledOutcome{result: result, err: runErr}
|
||||
}()
|
||||
select {
|
||||
case <-cancellationStarted:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for cancelable generation")
|
||||
}
|
||||
cancelExecution()
|
||||
select {
|
||||
case outcome := <-canceledResult:
|
||||
if outcome.result != nil ||
|
||||
!errors.Is(outcome.err, promptkit.ErrLLMGenerate) ||
|
||||
!errors.Is(outcome.err, context.Canceled) {
|
||||
t.Fatalf(
|
||||
"canceled execution=(%+v, %v), want generation and context identities",
|
||||
outcome.result,
|
||||
outcome.err,
|
||||
)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for canceled execution")
|
||||
}
|
||||
if result, err := cancelingEngine.RunPrepared(context.Background(), canceling); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("canceled execution was reusable: result=(%+v, %v)", result, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedExecutionConcurrentClaimAllowsOneGeneration(t *testing.T) {
|
||||
release := make(chan struct{})
|
||||
client := &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: "ok"},
|
||||
started: make(chan struct{}, 1),
|
||||
release: release,
|
||||
}
|
||||
engine := newPreparedContractEngine(t, client, "concurrent content")
|
||||
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prepared"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
|
||||
type outcome struct {
|
||||
result *promptkit.RunResult
|
||||
err error
|
||||
}
|
||||
outcomes := make(chan outcome, 2)
|
||||
for i := 0; i < 2; i++ {
|
||||
go func() {
|
||||
result, runErr := engine.RunPrepared(context.Background(), prepared)
|
||||
outcomes <- outcome{result: result, err: runErr}
|
||||
}()
|
||||
}
|
||||
|
||||
select {
|
||||
case <-client.started:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for generation")
|
||||
}
|
||||
select {
|
||||
case loser := <-outcomes:
|
||||
if loser.result != nil || !errors.Is(loser.err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("concurrent loser=(%+v, %v), want ErrInvalidRequest", loser.result, loser.err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for rejected concurrent claim")
|
||||
}
|
||||
|
||||
close(release)
|
||||
select {
|
||||
case winner := <-outcomes:
|
||||
if winner.err != nil || winner.result == nil {
|
||||
t.Fatalf("concurrent winner=(%+v, %v), want success", winner.result, winner.err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for successful concurrent claim")
|
||||
}
|
||||
if len(client.snapshot()) != 1 {
|
||||
t.Fatalf("generation calls=%d, want 1", len(client.snapshot()))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedExecutionRunAndDiscardRaceHasOneWinner(t *testing.T) {
|
||||
const attempts = 32
|
||||
|
||||
for i := 0; i < attempts; i++ {
|
||||
client := &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: "ok"},
|
||||
}
|
||||
engine := newPreparedContractEngine(t, client, "race content")
|
||||
prepared, err := engine.PrepareExecution(
|
||||
context.Background(),
|
||||
promptkit.RunRequest{PromptID: "prepared"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("attempt %d prepare execution: %v", i, err)
|
||||
}
|
||||
|
||||
start := make(chan struct{})
|
||||
type outcome struct {
|
||||
result *promptkit.RunResult
|
||||
err error
|
||||
}
|
||||
runOutcome := make(chan outcome, 1)
|
||||
discardDone := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
<-start
|
||||
result, runErr := engine.RunPrepared(context.Background(), prepared)
|
||||
runOutcome <- outcome{result: result, err: runErr}
|
||||
}()
|
||||
go func() {
|
||||
<-start
|
||||
prepared.Discard()
|
||||
close(discardDone)
|
||||
}()
|
||||
|
||||
close(start)
|
||||
runResult := <-runOutcome
|
||||
<-discardDone
|
||||
calls := len(client.snapshot())
|
||||
switch {
|
||||
case runResult.err == nil:
|
||||
if runResult.result == nil || calls != 1 {
|
||||
t.Fatalf(
|
||||
"attempt %d run won with outcome=(%+v, %v), generation calls=%d",
|
||||
i,
|
||||
runResult.result,
|
||||
runResult.err,
|
||||
calls,
|
||||
)
|
||||
}
|
||||
case errors.Is(runResult.err, promptkit.ErrInvalidRequest):
|
||||
if runResult.result != nil || calls != 0 {
|
||||
t.Fatalf(
|
||||
"attempt %d discard won with outcome=(%+v, %v), generation calls=%d",
|
||||
i,
|
||||
runResult.result,
|
||||
runResult.err,
|
||||
calls,
|
||||
)
|
||||
}
|
||||
default:
|
||||
t.Fatalf("attempt %d unexpected run outcome=(%+v, %v)", i, runResult.result, runResult.err)
|
||||
}
|
||||
if prepared.Details().PromptID != "prepared" {
|
||||
t.Fatalf("attempt %d details unavailable after race", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedExecutionDiscardAndFormattingDoNotExposePrivateState(t *testing.T) {
|
||||
const (
|
||||
directCredential = "pk-test-direct-credential-41f7"
|
||||
renderedContent = "rendered-content-sentinel-98d2"
|
||||
)
|
||||
client := &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: "generated output"},
|
||||
}
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prepared", "profile", renderedContent), "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "profile",
|
||||
Endpoint: "http://example.test/v1",
|
||||
Model: "model",
|
||||
APIKeyRequired: true,
|
||||
}),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{
|
||||
PromptID: "prepared",
|
||||
APIKey: directCredential,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution: %v", err)
|
||||
}
|
||||
|
||||
copied := *prepared
|
||||
zeroValue := promptkit.PreparedExecution{}
|
||||
var nilHandle *promptkit.PreparedExecution
|
||||
for name, value := range map[string]any{
|
||||
"original pointer": prepared,
|
||||
"copied value": copied,
|
||||
"zero value": zeroValue,
|
||||
"zero pointer": &zeroValue,
|
||||
} {
|
||||
for format, formatted := range map[string]string{
|
||||
"String": fmt.Sprintf("%s", value),
|
||||
"GoString": fmt.Sprintf("%#v", value),
|
||||
"v": fmt.Sprintf("%v", value),
|
||||
"+v": fmt.Sprintf("%+v", value),
|
||||
} {
|
||||
if formatted != "promptkit.PreparedExecution{opaque}" {
|
||||
t.Fatalf("%s %s formatting = %q, want opaque representation", name, format, formatted)
|
||||
}
|
||||
assertPreparedPrivateValuesAbsent(t, formatted, directCredential, renderedContent)
|
||||
}
|
||||
}
|
||||
for format, formatted := range map[string]string{
|
||||
"String": fmt.Sprintf("%s", nilHandle),
|
||||
"GoString": fmt.Sprintf("%#v", nilHandle),
|
||||
"v": fmt.Sprintf("%v", nilHandle),
|
||||
"+v": fmt.Sprintf("%+v", nilHandle),
|
||||
} {
|
||||
if formatted != "<nil>" {
|
||||
t.Fatalf("nil pointer %s formatting = %q, want <nil>", format, formatted)
|
||||
}
|
||||
assertPreparedPrivateValuesAbsent(t, formatted, directCredential, renderedContent)
|
||||
}
|
||||
payload, err := json.Marshal(prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal opaque handle: %v", err)
|
||||
}
|
||||
assertPreparedPrivateValuesAbsent(t, string(payload), directCredential, renderedContent)
|
||||
|
||||
detailsBefore := prepared.Details()
|
||||
detailsJSON, err := json.Marshal(detailsBefore)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal prepared details: %v", err)
|
||||
}
|
||||
assertPreparedPrivateValuesAbsent(t, string(detailsJSON), directCredential)
|
||||
|
||||
executionResult, err := engine.RunPrepared(context.Background(), &copied)
|
||||
if err != nil {
|
||||
t.Fatalf("run copied execution after formatting: %v", err)
|
||||
}
|
||||
requests := client.snapshot()
|
||||
if len(requests) != 1 || requests[0].APIKey != directCredential {
|
||||
t.Fatalf("direct credential did not reach only the client credential field: %#v", requests)
|
||||
}
|
||||
requestJSON, err := json.Marshal(requests[0])
|
||||
if err != nil {
|
||||
t.Fatalf("marshal captured generate request: %v", err)
|
||||
}
|
||||
for _, value := range []string{
|
||||
fmt.Sprint(requests[0]),
|
||||
fmt.Sprintf("%+v", requests[0]),
|
||||
fmt.Sprintf("%#v", requests[0]),
|
||||
string(requestJSON),
|
||||
fmt.Sprint(executionResult),
|
||||
} {
|
||||
assertPreparedPrivateValuesAbsent(t, value, directCredential)
|
||||
}
|
||||
resultJSON, err := json.Marshal(executionResult)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal execution result: %v", err)
|
||||
}
|
||||
assertPreparedPrivateValuesAbsent(t, string(resultJSON), directCredential)
|
||||
|
||||
nilHandle.Discard()
|
||||
if !reflect.DeepEqual(nilHandle.Details(), promptkit.PreparedRun{}) {
|
||||
t.Fatalf("nil handle details=%+v, want zero value", nilHandle.Details())
|
||||
}
|
||||
zeroHandle := &zeroValue
|
||||
zeroHandle.Discard()
|
||||
if !reflect.DeepEqual(zeroHandle.Details(), promptkit.PreparedRun{}) {
|
||||
t.Fatalf("zero handle details=%+v, want zero value", zeroHandle.Details())
|
||||
}
|
||||
|
||||
discarded, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{
|
||||
PromptID: "prepared",
|
||||
APIKey: directCredential,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution for discard: %v", err)
|
||||
}
|
||||
discardedDetails := discarded.Details()
|
||||
discarded.Discard()
|
||||
discarded.Discard()
|
||||
result, lifecycleErr := engine.RunPrepared(context.Background(), discarded)
|
||||
if result != nil || !errors.Is(lifecycleErr, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("discarded execution result=(%+v, %v), want ErrInvalidRequest", result, lifecycleErr)
|
||||
}
|
||||
assertPreparedPrivateValuesAbsent(t, lifecycleErr.Error(), directCredential, renderedContent)
|
||||
if !reflect.DeepEqual(discarded.Details(), discardedDetails) {
|
||||
t.Fatal("details changed after discard")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedExecutionCredentialCapacityAndTimingBoundaries(t *testing.T) {
|
||||
t.Run("credential is rechecked before generation", func(t *testing.T) {
|
||||
const (
|
||||
environmentName = "PROMPTKIT_PREPARED_CONTRACT_KEY"
|
||||
environmentKey = "environment-credential-sentinel"
|
||||
)
|
||||
t.Setenv(environmentName, environmentKey)
|
||||
client := &preparedRecordingClient{
|
||||
response: &promptkit.GenerateResponse{Content: "unexpected"},
|
||||
}
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prepared", "profile", "content"), "."),
|
||||
promptkit.WithProfileFS(preparedCredentialProfileSource(environmentName), "."),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct credential engine: %v", err)
|
||||
}
|
||||
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prepared"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare credential execution: %v", err)
|
||||
}
|
||||
if err := os.Unsetenv(environmentName); err != nil {
|
||||
t.Fatalf("unset credential environment: %v", err)
|
||||
}
|
||||
|
||||
result, err := engine.RunPrepared(context.Background(), prepared)
|
||||
if result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) ||
|
||||
!errors.Is(err, promptkit.ErrAPIKeyEnvMissing) {
|
||||
t.Fatalf("credential execution=(%+v, %v), want credential identities", result, err)
|
||||
}
|
||||
if len(client.snapshot()) != 0 {
|
||||
t.Fatalf("credential failure reached generation: %d calls", len(client.snapshot()))
|
||||
}
|
||||
assertPreparedPrivateValuesAbsent(t, err.Error(), environmentKey)
|
||||
if result, err := engine.RunPrepared(context.Background(), prepared); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("credential failure did not consume handle: result=(%+v, %v)", result, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("preparation does not admit and execution timing starts after retention", func(t *testing.T) {
|
||||
release := make(chan struct{})
|
||||
client := newCapacityGateClient(release, 4)
|
||||
engine := newBackendCapacityEngine(t, client, 1, capacityInt(0), nil)
|
||||
|
||||
activeRun := make(chan capacityRunResult, 1)
|
||||
go runCapacityRequest(
|
||||
engine,
|
||||
context.Background(),
|
||||
promptkit.RunRequest{PromptID: "prompt"},
|
||||
activeRun,
|
||||
)
|
||||
awaitCapacityRequest(t, client.started)
|
||||
|
||||
prepared, err := engine.PrepareExecution(
|
||||
context.Background(),
|
||||
promptkit.RunRequest{PromptID: "prompt"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare while capacity is full: %v", err)
|
||||
}
|
||||
if _, _, calls := client.snapshot(); calls != 1 {
|
||||
t.Fatalf("preparation invoked generation: calls=%d", calls)
|
||||
}
|
||||
if result, err := engine.RunPrepared(context.Background(), prepared); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrCapacityExceeded) {
|
||||
t.Fatalf("capacity execution=(%+v, %v), want ErrCapacityExceeded", result, err)
|
||||
} else {
|
||||
var capacityErr *promptkit.CapacityError
|
||||
if !errors.As(err, &capacityErr) || capacityErr == nil || capacityErr.BackendID != "limited" {
|
||||
t.Fatalf("capacity execution=%v, want limited CapacityError", err)
|
||||
}
|
||||
}
|
||||
if result, err := engine.RunPrepared(context.Background(), prepared); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("capacity rejection did not consume handle: result=(%+v, %v)", result, err)
|
||||
}
|
||||
if prepared.Details().PromptID != "prompt" {
|
||||
t.Fatal("details unavailable after capacity rejection")
|
||||
}
|
||||
|
||||
close(release)
|
||||
activeOutcome := awaitCapacityRun(t, activeRun)
|
||||
if activeOutcome.err != nil || activeOutcome.result == nil {
|
||||
t.Fatalf("active run outcome=(%+v, %v), want success", activeOutcome.result, activeOutcome.err)
|
||||
}
|
||||
|
||||
timed, err := engine.PrepareExecution(
|
||||
context.Background(),
|
||||
promptkit.RunRequest{PromptID: "prompt"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare timed execution: %v", err)
|
||||
}
|
||||
details := timed.Details()
|
||||
time.Sleep(25 * time.Millisecond)
|
||||
executionFloor := time.Now().UTC()
|
||||
result, err := engine.RunPrepared(context.Background(), timed)
|
||||
if err != nil {
|
||||
t.Fatalf("run timed execution: %v", err)
|
||||
}
|
||||
if result.StartTime.Before(executionFloor) ||
|
||||
!result.StartTime.After(details.EndTime) ||
|
||||
result.EndTime.Before(result.StartTime) ||
|
||||
result.Duration != result.EndTime.Sub(result.StartTime) {
|
||||
t.Fatalf(
|
||||
"execution timing includes preparation or retention: details_end=%s floor=%s result=%+v",
|
||||
details.EndTime,
|
||||
executionFloor,
|
||||
result,
|
||||
)
|
||||
}
|
||||
if _, _, calls := client.snapshot(); calls != 2 {
|
||||
t.Fatalf("generation calls=%d, want active and timed executions only", calls)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
type mutablePreparedArtifactReader struct {
|
||||
mu sync.Mutex
|
||||
body string
|
||||
hash string
|
||||
calls int
|
||||
}
|
||||
|
||||
func (r *mutablePreparedArtifactReader) Read(
|
||||
_ context.Context,
|
||||
_ promptkit.ArtifactRef,
|
||||
) (*promptkit.Artifact, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.calls++
|
||||
return &promptkit.Artifact{
|
||||
Body: []byte(r.body),
|
||||
Hash: r.hash,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *mutablePreparedArtifactReader) set(body, hash string) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.body = body
|
||||
r.hash = hash
|
||||
}
|
||||
|
||||
func (r *mutablePreparedArtifactReader) callCount() int {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return r.calls
|
||||
}
|
||||
|
||||
type preparedRecordingClient struct {
|
||||
mu sync.Mutex
|
||||
response *promptkit.GenerateResponse
|
||||
err error
|
||||
requests []promptkit.GenerateRequest
|
||||
started chan struct{}
|
||||
release <-chan struct{}
|
||||
}
|
||||
|
||||
func (c *preparedRecordingClient) Generate(
|
||||
ctx context.Context,
|
||||
request promptkit.GenerateRequest,
|
||||
) (*promptkit.GenerateResponse, error) {
|
||||
c.mu.Lock()
|
||||
c.requests = append(c.requests, request)
|
||||
c.mu.Unlock()
|
||||
|
||||
if c.started != nil {
|
||||
c.started <- struct{}{}
|
||||
}
|
||||
if c.release != nil {
|
||||
select {
|
||||
case <-c.release:
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
if c.err != nil {
|
||||
return nil, c.err
|
||||
}
|
||||
return c.response, nil
|
||||
}
|
||||
|
||||
func (c *preparedRecordingClient) snapshot() []promptkit.GenerateRequest {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return append([]promptkit.GenerateRequest(nil), c.requests...)
|
||||
}
|
||||
|
||||
func newPreparedContractEngine(
|
||||
t *testing.T,
|
||||
client promptkit.LLMClient,
|
||||
message string,
|
||||
) *promptkit.Engine {
|
||||
t.Helper()
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prepared", "profile", message), "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "profile",
|
||||
Endpoint: "http://example.test/v1",
|
||||
Model: "model",
|
||||
}),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct prepared execution engine: %v", err)
|
||||
}
|
||||
return engine
|
||||
}
|
||||
|
||||
func preparedPromptSource(label string) fstest.MapFS {
|
||||
return fstest.MapFS{
|
||||
"prompt.yaml": &fstest.MapFile{Data: []byte(`id: prepared
|
||||
version: "1"
|
||||
default_profile: profile
|
||||
inputs:
|
||||
- name: input
|
||||
required: true
|
||||
messages:
|
||||
- role: user
|
||||
content: 'Input={{input "input"}} Label={{.label}}'
|
||||
description: ` + label + `
|
||||
output:
|
||||
format: json
|
||||
validation_mode: json_schema
|
||||
schema_path: schema.json
|
||||
`)},
|
||||
}
|
||||
}
|
||||
|
||||
func preparedProfileSource(model string) fstest.MapFS {
|
||||
return fstest.MapFS{
|
||||
"profile.yaml": &fstest.MapFile{Data: []byte(`id: profile
|
||||
endpoint: http://example.test/v1
|
||||
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(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"title": "original root",
|
||||
"type": "object",
|
||||
"required": ["value"],
|
||||
"properties": {
|
||||
"value": {"$ref": "value.json"}
|
||||
}
|
||||
}`)},
|
||||
"value.json": &fstest.MapFile{Data: []byte(`{
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"type": "integer",
|
||||
"minimum": 2
|
||||
}`)},
|
||||
}
|
||||
}
|
||||
|
||||
func assertPreparedPrivateValuesAbsent(t *testing.T, value string, privateValues ...string) {
|
||||
t.Helper()
|
||||
for _, privateValue := range privateValues {
|
||||
if strings.Contains(value, privateValue) {
|
||||
t.Fatalf("value exposed private data %q: %s", privateValue, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
33
profiles.go
33
profiles.go
@@ -100,33 +100,34 @@ func toDomainProfile(publicProfile Profile) (domain.ExecutionProfile, error) {
|
||||
APIKeyRequired: publicProfile.APIKeyRequired,
|
||||
ExtraParams: extraParams,
|
||||
}
|
||||
if err := validatePublicProfile(prof); err != nil {
|
||||
if err := normalizeAndValidatePublicProfile(&prof); err != nil {
|
||||
return domain.ExecutionProfile{}, err
|
||||
}
|
||||
return prof, nil
|
||||
}
|
||||
|
||||
func validatePublicProfile(prof domain.ExecutionProfile) error {
|
||||
func normalizeAndValidatePublicProfile(prof *domain.ExecutionProfile) error {
|
||||
if strings.TrimSpace(prof.ID) == "" {
|
||||
return errors.New("id is required")
|
||||
}
|
||||
if strings.TrimSpace(prof.BackendID) == "" && strings.TrimSpace(prof.Endpoint) == "" {
|
||||
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")
|
||||
}
|
||||
if prof.Temperature < 0 || prof.Temperature > 2 {
|
||||
return errors.New("temperature must be between 0 and 2")
|
||||
}
|
||||
if prof.MaxTokens < 0 {
|
||||
return errors.New("max_tokens must be greater than or equal to 0")
|
||||
}
|
||||
if prof.TopP < 0 || prof.TopP > 1 {
|
||||
return errors.New("top_p must be between 0 and 1")
|
||||
}
|
||||
if prof.TimeoutSeconds < 0 {
|
||||
return errors.New("timeout_seconds must be greater than or equal to 0")
|
||||
}
|
||||
return nil
|
||||
return domain.ValidateExecutionTargetSettings(domain.ExecutionTarget{
|
||||
Temperature: prof.Temperature,
|
||||
MaxTokens: prof.MaxTokens,
|
||||
TopP: prof.TopP,
|
||||
TimeoutSeconds: prof.TimeoutSeconds,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -5,25 +5,423 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/promptkit"
|
||||
)
|
||||
|
||||
func TestPreparedRunJSONOmitsZeroTimingValues(t *testing.T) {
|
||||
payload, err := json.Marshal(promptkit.PreparedRun{})
|
||||
type inspectionCountingFS struct {
|
||||
opens atomic.Int64
|
||||
}
|
||||
|
||||
func (f *inspectionCountingFS) Open(string) (fs.File, error) {
|
||||
f.opens.Add(1)
|
||||
return nil, fs.ErrNotExist
|
||||
}
|
||||
|
||||
func TestInspectProfileResolvesCredentialStatesWithoutPromptOrGeneration(t *testing.T) {
|
||||
const environmentName = "PROMPTKIT_INSPECTION_ABSENT_KEY"
|
||||
t.Setenv(environmentName, "")
|
||||
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "unexpected"}}
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(fstest.MapFS{}, "."),
|
||||
promptkit.WithProfileFS(fstest.MapFS{
|
||||
"environment.yaml": &fstest.MapFile{Data: []byte(`id: environment
|
||||
endpoint: http://environment.example/v1
|
||||
model: environment-model
|
||||
api_key_env: PROMPTKIT_INSPECTION_ABSENT_KEY
|
||||
`)},
|
||||
}, "."),
|
||||
promptkit.WithProfiles(
|
||||
promptkit.Profile{ID: "direct", Endpoint: "http://direct.example/v1", Model: "direct-model", APIKeyRequired: true},
|
||||
promptkit.Profile{ID: "none", Endpoint: "http://none.example/v1", Model: "none-model"},
|
||||
),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal prepared run: %v", err)
|
||||
t.Fatalf("construct inspection engine: %v", err)
|
||||
}
|
||||
for _, field := range []string{"start_time", "end_time", "duration_ms"} {
|
||||
if strings.Contains(string(payload), `"`+field+`"`) {
|
||||
t.Fatalf("expected zero %s to be omitted, got %s", field, payload)
|
||||
|
||||
for _, tc := range []struct {
|
||||
profileID string
|
||||
wantEnv string
|
||||
wantDirectKey bool
|
||||
wantEndpoint string
|
||||
}{
|
||||
{profileID: " environment ", wantEnv: environmentName, wantEndpoint: "http://environment.example/v1"},
|
||||
{profileID: "direct", wantDirectKey: true, wantEndpoint: "http://direct.example/v1"},
|
||||
{profileID: "none", wantEndpoint: "http://none.example/v1"},
|
||||
} {
|
||||
t.Run(tc.profileID, func(t *testing.T) {
|
||||
inspection, err := engine.InspectProfile(context.Background(), tc.profileID)
|
||||
if err != nil {
|
||||
t.Fatalf("inspect profile: %v", err)
|
||||
}
|
||||
if inspection.ProfileID != strings.TrimSpace(tc.profileID) ||
|
||||
inspection.EffectiveModelParams.Endpoint != tc.wantEndpoint ||
|
||||
inspection.EffectiveModelParams.BackendID != "" ||
|
||||
inspection.EffectiveModelParams.APIKeyEnv != tc.wantEnv ||
|
||||
inspection.APIKeyRequired != tc.wantDirectKey {
|
||||
t.Fatalf("unexpected inspection: %#v", inspection)
|
||||
}
|
||||
})
|
||||
}
|
||||
if len(client.requests) != 0 {
|
||||
t.Fatalf("inspection invoked the model client %d times", len(client.requests))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileProfileNormalizedIDMatchesInspectionAndPreparation(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "normalized-profile", "message"), "."),
|
||||
promptkit.WithProfileFS(fstest.MapFS{
|
||||
"profile.yaml": &fstest.MapFile{Data: []byte(`
|
||||
id: " normalized-profile "
|
||||
endpoint: http://profile.example/v1
|
||||
model: normalized-model
|
||||
`)},
|
||||
}, "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
|
||||
inspection, err := engine.InspectProfile(context.Background(), " normalized-profile ")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect normalized profile: %v", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{
|
||||
PromptID: "prompt",
|
||||
ProfileID: " normalized-profile ",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare with normalized profile: %v", err)
|
||||
}
|
||||
|
||||
if inspection.ProfileID != "normalized-profile" ||
|
||||
prepared.SelectedProfileID != inspection.ProfileID ||
|
||||
inspection.EffectiveModelParams.Model != "normalized-model" ||
|
||||
prepared.EffectiveModelParams.Model != inspection.EffectiveModelParams.Model {
|
||||
t.Fatalf("inspection=%#v prepared=%#v", inspection, prepared)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectProfilePreservesPublicErrorIdentities(t *testing.T) {
|
||||
newEngine := func(t *testing.T, options ...promptkit.Option) *promptkit.Engine {
|
||||
t.Helper()
|
||||
engine, err := promptkit.NewEngine(
|
||||
promptkit.Config{},
|
||||
append([]promptkit.Option{promptkit.WithPromptFS(fstest.MapFS{}, ".")}, options...)...,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct inspection engine: %v", err)
|
||||
}
|
||||
return engine
|
||||
}
|
||||
|
||||
var nilEngine *promptkit.Engine
|
||||
if result, err := nilEngine.InspectProfile(context.Background(), "profile"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidConfig) {
|
||||
t.Fatalf("nil engine result=(%#v, %v), want ErrInvalidConfig", result, err)
|
||||
}
|
||||
|
||||
valid := newEngine(t, promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "profile", Endpoint: "http://profile.example/v1", Model: "model",
|
||||
}))
|
||||
if result, err := valid.InspectProfile(context.Background(), " \t "); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("blank profile result=(%#v, %v), want ErrInvalidRequest", result, err)
|
||||
}
|
||||
if result, err := valid.InspectProfile(context.Background(), "missing"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrProfileNotFound) || errors.Is(err, promptkit.ErrProfileLoad) {
|
||||
t.Fatalf("missing profile result=(%#v, %v), want only ErrProfileNotFound", result, err)
|
||||
}
|
||||
|
||||
malformed := newEngine(t, promptkit.WithProfileFS(fstest.MapFS{
|
||||
"broken.yaml": &fstest.MapFile{Data: []byte("id: broken\nendpoint: http://broken.example/v1\nmodel: model\nextra_params:\n invalid: .nan\n")},
|
||||
}, "."))
|
||||
if result, err := malformed.InspectProfile(context.Background(), "broken"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrProfileLoad) {
|
||||
t.Fatalf("malformed profile result=(%#v, %v), want ErrProfileLoad", result, err)
|
||||
}
|
||||
|
||||
unknownBackend := newEngine(t, promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "unknown-backend", BackendID: "unknown", Model: "model",
|
||||
}))
|
||||
if result, err := unknownBackend.InspectProfile(context.Background(), "unknown-backend"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrProfileLoad) {
|
||||
t.Fatalf("unknown backend result=(%#v, %v), want ErrProfileLoad", result, err)
|
||||
}
|
||||
|
||||
countingFS := &inspectionCountingFS{}
|
||||
canceled := newEngine(t, promptkit.WithProfileFS(countingFS, "."))
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if result, err := canceled.InspectProfile(ctx, "profile"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrProfileLoad) || !errors.Is(err, context.Canceled) || countingFS.opens.Load() != 0 {
|
||||
t.Fatalf("canceled inspection result=(%#v, %v), opens=%d", result, err, countingFS.opens.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectProfileReturnsIndependentTargetMatchingPreparation(t *testing.T) {
|
||||
extraParams := map[string]any{
|
||||
"nested": map[string]any{"value": "original"},
|
||||
}
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "profile", "message"), "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "profile", Endpoint: "http://profile.example/v1", Model: "model", ExtraParams: extraParams,
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct inspection engine: %v", err)
|
||||
}
|
||||
|
||||
first, err := engine.InspectProfile(context.Background(), "profile")
|
||||
if err != nil {
|
||||
t.Fatalf("first inspection: %v", err)
|
||||
}
|
||||
first.EffectiveModelParams.ExtraParams["nested"].(map[string]any)["value"] = "changed"
|
||||
first.EffectiveModelParams.ExtraParams["later"] = true
|
||||
|
||||
second, err := engine.InspectProfile(context.Background(), "profile")
|
||||
if err != nil {
|
||||
t.Fatalf("second inspection: %v", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare after inspection mutation: %v", err)
|
||||
}
|
||||
for _, target := range []promptkit.ExecutionTarget{second.EffectiveModelParams, prepared.EffectiveModelParams} {
|
||||
if target.ExtraParams["nested"].(map[string]any)["value"] != "original" || target.ExtraParams["later"] != nil {
|
||||
t.Fatalf("inspection mutation reached engine-owned target: %#v", target.ExtraParams)
|
||||
}
|
||||
}
|
||||
if !reflect.DeepEqual(second.EffectiveModelParams, prepared.EffectiveModelParams) {
|
||||
t.Fatalf("inspection target=%#v, preparation target=%#v", second.EffectiveModelParams, prepared.EffectiveModelParams)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectPromptReturnsDeclaredMetadataWithoutExecutionWork(t *testing.T) {
|
||||
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "unexpected"}}
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(fstest.MapFS{
|
||||
"report-v1.yaml": &fstest.MapFile{Data: []byte(`id: report
|
||||
version: "1.0.0"
|
||||
messages:
|
||||
- role: user
|
||||
content: old report
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
"report-v2.yaml": &fstest.MapFile{Data: []byte(`id: report
|
||||
version: "2.0.0"
|
||||
default_profile: missing-profile
|
||||
inputs:
|
||||
- name: location
|
||||
required: true
|
||||
content_type: text/plain
|
||||
description: Forecast location.
|
||||
- name: units
|
||||
content_type: text/plain
|
||||
description: Unit preference.
|
||||
messages:
|
||||
- role: user
|
||||
content_file: messages/report.md
|
||||
output:
|
||||
format: json
|
||||
validation_mode: json_schema
|
||||
schema_path: schemas/report.json
|
||||
`)},
|
||||
"messages/report.md": &fstest.MapFile{Data: []byte("rendered report body is not returned")},
|
||||
}, "."),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct prompt-inspection engine: %v", err)
|
||||
}
|
||||
|
||||
inspection, err := engine.InspectPrompt(context.Background(), "report", "2.0.0")
|
||||
if err != nil {
|
||||
t.Fatalf("inspect prompt: %v", err)
|
||||
}
|
||||
if inspection.PromptID != "report" ||
|
||||
inspection.PromptVersion != "2.0.0" ||
|
||||
inspection.PromptHash == "" ||
|
||||
inspection.DefaultProfileID != "missing-profile" ||
|
||||
inspection.OutputContract != (promptkit.OutputContract{
|
||||
Format: promptkit.FormatJSON,
|
||||
ValidationMode: promptkit.ValidationJSONSchema,
|
||||
SchemaPath: "schemas/report.json",
|
||||
}) {
|
||||
t.Fatalf("unexpected inspection metadata: %#v", inspection)
|
||||
}
|
||||
wantInputs := []promptkit.PromptInputDefinition{
|
||||
{Name: "location", Required: true, ContentType: "text/plain", Description: "Forecast location."},
|
||||
{Name: "units", ContentType: "text/plain", Description: "Unit preference."},
|
||||
}
|
||||
if !reflect.DeepEqual(inspection.Inputs, wantInputs) {
|
||||
t.Fatalf("inspection inputs=%#v, want %#v", inspection.Inputs, wantInputs)
|
||||
}
|
||||
if len(client.requests) != 0 {
|
||||
t.Fatalf("inspection invoked the model client %d times", len(client.requests))
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectPromptPreservesPublicErrorIdentities(t *testing.T) {
|
||||
newEngine := func(t *testing.T, source fstest.MapFS) *promptkit.Engine {
|
||||
t.Helper()
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{}, promptkit.WithPromptFS(source, "."))
|
||||
if err != nil {
|
||||
t.Fatalf("construct prompt-inspection engine: %v", err)
|
||||
}
|
||||
return engine
|
||||
}
|
||||
validSource := fstest.MapFS{
|
||||
"prompt.yaml": &fstest.MapFile{Data: []byte(`id: prompt
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: body
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
}
|
||||
|
||||
var nilEngine *promptkit.Engine
|
||||
if result, err := nilEngine.InspectPrompt(context.Background(), "prompt", "1"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidConfig) {
|
||||
t.Fatalf("nil engine result=(%#v, %v), want ErrInvalidConfig", result, err)
|
||||
}
|
||||
|
||||
valid := newEngine(t, validSource)
|
||||
if result, err := valid.InspectPrompt(context.Background(), " \t ", "1"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrInvalidRequest) {
|
||||
t.Fatalf("blank prompt result=(%#v, %v), want ErrInvalidRequest", result, err)
|
||||
}
|
||||
if result, err := valid.InspectPrompt(context.Background(), "missing", "1"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrPromptNotFound) || errors.Is(err, promptkit.ErrPromptLoad) {
|
||||
t.Fatalf("missing prompt result=(%#v, %v), want only ErrPromptNotFound", result, err)
|
||||
}
|
||||
if result, err := valid.InspectPrompt(context.Background(), "prompt", "missing"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrPromptNotFound) || errors.Is(err, promptkit.ErrPromptLoad) {
|
||||
t.Fatalf("missing version result=(%#v, %v), want only ErrPromptNotFound", result, err)
|
||||
}
|
||||
|
||||
ambiguous := newEngine(t, fstest.MapFS{
|
||||
"one.yaml": &fstest.MapFile{Data: []byte(`id: prompt
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content: first
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
"two.yaml": &fstest.MapFile{Data: []byte(`id: prompt
|
||||
version: "2"
|
||||
messages:
|
||||
- role: user
|
||||
content: second
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
})
|
||||
if result, err := ambiguous.InspectPrompt(context.Background(), "prompt", ""); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrPromptLoad) {
|
||||
t.Fatalf("ambiguous prompt result=(%#v, %v), want ErrPromptLoad", result, err)
|
||||
}
|
||||
|
||||
for name, source := range map[string]fstest.MapFS{
|
||||
"malformed definition": {
|
||||
"broken.yaml": &fstest.MapFile{Data: []byte("id: broken\nversion: \"1\"\nunknown: value\n")},
|
||||
},
|
||||
"missing content file": {
|
||||
"broken.yaml": &fstest.MapFile{Data: []byte(`id: broken
|
||||
version: "1"
|
||||
messages:
|
||||
- role: user
|
||||
content_file: missing.md
|
||||
output:
|
||||
format: text
|
||||
validation_mode: none
|
||||
`)},
|
||||
},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if result, err := newEngine(t, source).InspectPrompt(context.Background(), "broken", "1"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrPromptLoad) {
|
||||
t.Fatalf("broken prompt result=(%#v, %v), want ErrPromptLoad", result, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
countingFS := &inspectionCountingFS{}
|
||||
canceled, err := promptkit.NewEngine(promptkit.Config{}, promptkit.WithPromptFS(countingFS, "."))
|
||||
if err != nil {
|
||||
t.Fatalf("construct canceled prompt-inspection engine: %v", err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if result, err := canceled.InspectPrompt(ctx, "prompt", "1"); result != nil ||
|
||||
!errors.Is(err, promptkit.ErrPromptLoad) || !errors.Is(err, context.Canceled) || countingFS.opens.Load() != 0 {
|
||||
t.Fatalf("canceled inspection result=(%#v, %v), opens=%d", result, err, countingFS.opens.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectPromptReturnsIndependentMetadataMatchingPreparation(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(fstest.MapFS{
|
||||
"prompt.yaml": &fstest.MapFile{Data: []byte(`id: prompt
|
||||
version: "1"
|
||||
default_profile: profile
|
||||
inputs:
|
||||
- name: subject
|
||||
content_type: text/plain
|
||||
description: Summary subject.
|
||||
messages:
|
||||
- role: user
|
||||
content: summarize
|
||||
output:
|
||||
format: markdown
|
||||
validation_mode: basic
|
||||
`)},
|
||||
}, "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "profile", Endpoint: "http://profile.example/v1", Model: "model",
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct prompt-inspection engine: %v", err)
|
||||
}
|
||||
|
||||
first, err := engine.InspectPrompt(context.Background(), "prompt", "1")
|
||||
if err != nil {
|
||||
t.Fatalf("first inspection: %v", err)
|
||||
}
|
||||
first.Inputs[0].Name = "changed"
|
||||
first.OutputContract.SchemaPath = "changed.json"
|
||||
|
||||
second, err := engine.InspectPrompt(context.Background(), "prompt", "1")
|
||||
if err != nil {
|
||||
t.Fatalf("second inspection: %v", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt", PromptVersion: "1"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare after inspection mutation: %v", err)
|
||||
}
|
||||
if second.Inputs[0].Name != "subject" || second.OutputContract.SchemaPath != "" ||
|
||||
prepared.OutputContract.SchemaPath != "" || second.PromptHash != prepared.PromptHash {
|
||||
t.Fatalf("inspection mutation reached engine-owned prompt metadata: inspection=%#v prepared=%#v", second, prepared)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -138,6 +536,49 @@ func TestUnknownProfileBackendHasProfileLoadIdentity(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalBackendConstructsAndRegistersConventionalBackend(t *testing.T) {
|
||||
const (
|
||||
localBackendID = "local"
|
||||
limit = 2
|
||||
)
|
||||
if promptkit.BackendLocal != localBackendID {
|
||||
t.Fatalf("BackendLocal=%q, want %q", promptkit.BackendLocal, localBackendID)
|
||||
}
|
||||
|
||||
endpoint := "http://local.example/v1"
|
||||
backend := promptkit.LocalBackend(endpoint, limit)
|
||||
want := promptkit.Backend{
|
||||
ID: localBackendID,
|
||||
Endpoint: endpoint,
|
||||
ConcurrencyLimit: limit,
|
||||
}
|
||||
if !reflect.DeepEqual(backend, want) {
|
||||
t.Fatalf("LocalBackend()=%+v, want %+v", backend, want)
|
||||
}
|
||||
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "local-profile", "message"), "."),
|
||||
promptkit.WithBackend(backend),
|
||||
promptkit.WithProfiles(promptkit.Profile{
|
||||
ID: "local-profile",
|
||||
BackendID: localBackendID,
|
||||
Model: "model",
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine with local backend: %v", err)
|
||||
}
|
||||
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare with local backend: %v", err)
|
||||
}
|
||||
if prepared.SelectedBackendID != localBackendID ||
|
||||
prepared.EffectiveModelParams.Endpoint != endpoint {
|
||||
t.Fatalf("unexpected local backend preparation: %+v", prepared)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomBackendFlowsThroughProfilesOverridesAndInjectedClient(t *testing.T) {
|
||||
t.Setenv("CUSTOM_LLM_KEY", "test-key")
|
||||
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "ok"}}
|
||||
@@ -338,6 +779,7 @@ func TestBackendRegistrationRejectsInvalidAndDuplicateDefinitions(t *testing.T)
|
||||
{name: "reserved extra parameter", backends: []promptkit.Backend{{ID: "custom", Endpoint: "http://example.test/v1", ExtraParams: map[string]any{"model": "override"}}}},
|
||||
{name: "cyclic extra parameter", backends: []promptkit.Backend{{ID: "custom", Endpoint: "http://example.test/v1", ExtraParams: cycle}}},
|
||||
{name: "malformed JSON number", backends: []promptkit.Backend{{ID: "custom", Endpoint: "http://example.test/v1", ExtraParams: map[string]any{"value": json.Number("01")}}}},
|
||||
{name: "excessively deep extra parameter", backends: []promptkit.Backend{{ID: "custom", Endpoint: "http://example.test/v1", ExtraParams: map[string]any{"value": excessivelyDeepJSONValue()}}}},
|
||||
{name: "duplicate consumer id", backends: []promptkit.Backend{
|
||||
{ID: " custom ", Endpoint: "http://one.example/v1"},
|
||||
{ID: "custom", Endpoint: "http://two.example/v1"},
|
||||
@@ -399,92 +841,6 @@ func TestBackendExtraParamsAreDeeplyCopiedAtConstructionAndLookup(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreparedRunJSONTimingRoundTrips(t *testing.T) {
|
||||
start := time.Date(2026, time.July, 29, 12, 0, 0, 0, time.UTC)
|
||||
prepared := promptkit.PreparedRun{
|
||||
PromptID: "prompt",
|
||||
StartTime: start,
|
||||
EndTime: start.Add(1250 * time.Millisecond),
|
||||
DurationMS: 1250,
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(prepared)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal prepared run: %v", err)
|
||||
}
|
||||
var decoded promptkit.PreparedRun
|
||||
if err := json.Unmarshal(payload, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal prepared run: %v", err)
|
||||
}
|
||||
if decoded.DurationMS != prepared.DurationMS ||
|
||||
!decoded.StartTime.Equal(prepared.StartTime) ||
|
||||
!decoded.EndTime.Equal(prepared.EndTime) {
|
||||
t.Fatalf("timing values did not round trip: got %#v, want %#v", decoded, prepared)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunResultJSONUsesMillisecondsAndRoundTrips(t *testing.T) {
|
||||
start := time.Date(2026, time.July, 29, 12, 0, 0, 0, time.UTC)
|
||||
result := promptkit.RunResult{
|
||||
RunID: "opaque-run-id",
|
||||
Artifact: promptkit.Artifact{Name: "output", ContentType: "text/plain", Body: []byte("ok")},
|
||||
SessionID: "session-123",
|
||||
StartTime: start,
|
||||
EndTime: start.Add(1500 * time.Millisecond),
|
||||
Duration: 1500 * time.Millisecond,
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(result)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal run result: %v", err)
|
||||
}
|
||||
var object map[string]any
|
||||
if err := json.Unmarshal(payload, &object); err != nil {
|
||||
t.Fatalf("decode run result JSON: %v", err)
|
||||
}
|
||||
if got := object["duration_ms"]; got != float64(1500) {
|
||||
t.Fatalf("expected duration_ms=1500, got %#v in %s", got, payload)
|
||||
}
|
||||
if _, exists := object["duration"]; exists {
|
||||
t.Fatalf("unexpected nanosecond duration field in %s", payload)
|
||||
}
|
||||
if got := object["session_id"]; got != result.SessionID {
|
||||
t.Fatalf("expected session_id=%q, got %#v in %s", result.SessionID, got, payload)
|
||||
}
|
||||
artifact, ok := object["artifact"].(map[string]any)
|
||||
if !ok || artifact["content_type"] != "text/plain" {
|
||||
t.Fatalf("expected stable artifact JSON fields, got %#v", object["artifact"])
|
||||
}
|
||||
|
||||
var decoded promptkit.RunResult
|
||||
if err := json.Unmarshal(payload, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal run result: %v", err)
|
||||
}
|
||||
if decoded.SessionID != result.SessionID ||
|
||||
decoded.Duration != result.Duration ||
|
||||
!decoded.StartTime.Equal(result.StartTime) ||
|
||||
!decoded.EndTime.Equal(result.EndTime) {
|
||||
t.Fatalf("timing values did not round trip: got %#v, want %#v", decoded, result)
|
||||
}
|
||||
|
||||
payload, err = json.Marshal(promptkit.RunResult{})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal zero run result: %v", err)
|
||||
}
|
||||
for _, field := range []string{"session_id", "start_time", "end_time", "duration_ms"} {
|
||||
if strings.Contains(string(payload), `"`+field+`"`) {
|
||||
t.Fatalf("expected zero %s to be omitted, got %s", field, payload)
|
||||
}
|
||||
}
|
||||
var decodedEmpty promptkit.RunResult
|
||||
if err := json.Unmarshal(payload, &decodedEmpty); err != nil {
|
||||
t.Fatalf("unmarshal run result without session_id: %v", err)
|
||||
}
|
||||
if decodedEmpty.SessionID != "" {
|
||||
t.Fatalf("expected absent session_id to decode empty, got %q", decodedEmpty.SessionID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEngineValidationIsSinglePass(t *testing.T) {
|
||||
client := &fakeLLMClient{
|
||||
response: &promptkit.GenerateResponse{Content: "not-json"},
|
||||
@@ -554,6 +910,24 @@ func TestRepeatedOptionsUseLastValueInEachCategory(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("fallback profile source", func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "profile", "message"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS("profile", "first-model"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS("profile", "second-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare from last fallback profile source: %v", err)
|
||||
}
|
||||
if prepared.EffectiveModelParams.Model != "second-model" {
|
||||
t.Fatalf("expected last fallback profile source, got %q", prepared.EffectiveModelParams.Model)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("in-memory profiles", func(t *testing.T) {
|
||||
first := profile
|
||||
first.Model = "first-model"
|
||||
@@ -638,6 +1012,227 @@ func TestRepeatedOptionsUseLastValueInEachCategory(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestFallbackProfileSourcePrecedence(t *testing.T) {
|
||||
const profileID = "application-profile"
|
||||
|
||||
prepareModel := func(t *testing.T, engine *promptkit.Engine, promptID string) string {
|
||||
t.Helper()
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: promptID})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare: %v", err)
|
||||
}
|
||||
return prepared.EffectiveModelParams.Model
|
||||
}
|
||||
|
||||
t.Run("in-memory profiles override ordinary and fallback profiles", func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", profileID, "message"), "."),
|
||||
promptkit.WithProfiles(promptkit.Profile{ID: profileID, Endpoint: "http://example.test/v1", Model: "memory-model"}),
|
||||
promptkit.WithProfileFS(contractProfileFS(profileID, "ordinary-model"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS(profileID, "fallback-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if model := prepareModel(t, engine, "prompt"); model != "memory-model" {
|
||||
t.Fatalf("expected in-memory profile, got %q", model)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ordinary filesystem source overrides fallback profile", func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", profileID, "message"), "."),
|
||||
promptkit.WithProfileFS(contractProfileFS(profileID, "ordinary-model"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS(profileID, "fallback-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if model := prepareModel(t, engine, "prompt"); model != "ordinary-model" {
|
||||
t.Fatalf("expected ordinary profile, got %q", model)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ordinary option replaces configured directory", func(t *testing.T) {
|
||||
profileDir := t.TempDir()
|
||||
writePublicProfileFile(t, profileDir, profileID, "http://example.test/v1", "directory-model")
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{ProfileDir: profileDir},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", profileID, "message"), "."),
|
||||
promptkit.WithProfileFS(contractProfileFS(profileID, "option-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if model := prepareModel(t, engine, "prompt"); model != "option-model" {
|
||||
t.Fatalf("expected ordinary option profile, got %q", model)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("configured directory overrides fallback profile", func(t *testing.T) {
|
||||
profileDir := t.TempDir()
|
||||
writePublicProfileFile(t, profileDir, profileID, "http://example.test/v1", "directory-model")
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{ProfileDir: profileDir},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", profileID, "message"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS(profileID, "fallback-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if model := prepareModel(t, engine, "prompt"); model != "directory-model" {
|
||||
t.Fatalf("expected configured directory profile, got %q", model)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("fallback profile overrides built-in profile", func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "mistral-small-3", "message"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS("mistral-small-3", "fallback-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if model := prepareModel(t, engine, "prompt"); model != "fallback-model" {
|
||||
t.Fatalf("expected fallback profile, got %q", model)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing fallback profile uses built-in profile", func(t *testing.T) {
|
||||
t.Setenv("OPENROUTER_API_KEY", "test-key")
|
||||
baseline, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "mistral-small-3", "message"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct baseline engine: %v", err)
|
||||
}
|
||||
want := prepareModel(t, baseline, "prompt")
|
||||
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "mistral-small-3", "message"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS(profileID, "fallback-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if model := prepareModel(t, engine, "prompt"); model != want {
|
||||
t.Fatalf("expected built-in profile model %q, got %q", want, model)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestFallbackProfileSourcePreservesLazyLoadingAndErrors(t *testing.T) {
|
||||
const profileID = "application-profile"
|
||||
|
||||
t.Run("construction defers malformed fallback profiles", func(t *testing.T) {
|
||||
_, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(fstest.MapFS{}, "."),
|
||||
promptkit.WithFallbackProfileFS(fstest.MapFS{
|
||||
"broken.yaml": &fstest.MapFile{Data: []byte("id: broken\nunknown: value\n")},
|
||||
}, "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine with malformed fallback profile: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unrelated malformed fallback profile does not block matching definition", func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", profileID, "message"), "."),
|
||||
promptkit.WithFallbackProfileFS(fstest.MapFS{
|
||||
"broken.yaml": &fstest.MapFile{Data: []byte("id: unrelated\nunknown: value\n")},
|
||||
"valid.yaml": &fstest.MapFile{Data: []byte("id: application-profile\nendpoint: http://example.test/v1\nmodel: fallback-model\n")},
|
||||
}, "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare from valid fallback profile: %v", err)
|
||||
}
|
||||
if prepared.EffectiveModelParams.Model != "fallback-model" {
|
||||
t.Fatalf("unexpected fallback profile model: %q", prepared.EffectiveModelParams.Model)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("matching malformed fallback profile does not reach built-in profile", func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "mistral-small-3", "message"), "."),
|
||||
promptkit.WithFallbackProfileFS(fstest.MapFS{
|
||||
"mistral-small-3.yaml": &fstest.MapFile{Data: []byte("id: mistral-small-3\nendpoint: http://example.test/v1\nmodel: fallback-model\nunknown: value\n")},
|
||||
}, "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if _, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"}); !errors.Is(err, promptkit.ErrProfileLoad) {
|
||||
t.Fatalf("expected ErrProfileLoad, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("matching malformed ordinary profile does not reach fallback profile", func(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", profileID, "message"), "."),
|
||||
promptkit.WithProfileFS(fstest.MapFS{
|
||||
"application-profile.yaml": &fstest.MapFile{Data: []byte("id: application-profile\nendpoint: http://example.test/v1\nmodel: ordinary-model\nunknown: value\n")},
|
||||
}, "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS(profileID, "fallback-model"), "."),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
if _, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"}); !errors.Is(err, promptkit.ErrProfileLoad) {
|
||||
t.Fatalf("expected ErrProfileLoad, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestFallbackProfileSourceWorksAcrossWorkflows(t *testing.T) {
|
||||
const profileID = "application-profile"
|
||||
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "ok"}}
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", profileID, "message"), "."),
|
||||
promptkit.WithFallbackProfileFS(contractProfileFS(profileID, "fallback-model"), "."),
|
||||
promptkit.WithLLMClient(client),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
|
||||
inspection, err := engine.InspectProfile(context.Background(), profileID)
|
||||
if err != nil {
|
||||
t.Fatalf("inspect fallback profile: %v", err)
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare fallback profile: %v", err)
|
||||
}
|
||||
preparedExecution, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare execution with fallback profile: %v", err)
|
||||
}
|
||||
preparedDetails := preparedExecution.Details()
|
||||
preparedResult, err := engine.RunPrepared(context.Background(), preparedExecution)
|
||||
if err != nil {
|
||||
t.Fatalf("run prepared fallback profile: %v", err)
|
||||
}
|
||||
runResult, err := engine.Run(context.Background(), promptkit.RunRequest{PromptID: "prompt"})
|
||||
if err != nil {
|
||||
t.Fatalf("run fallback profile: %v", err)
|
||||
}
|
||||
|
||||
for name, model := range map[string]string{
|
||||
"inspection": inspection.EffectiveModelParams.Model,
|
||||
"preparation": prepared.EffectiveModelParams.Model,
|
||||
"prepared execution": preparedDetails.EffectiveModelParams.Model,
|
||||
"prepared result": preparedResult.EffectiveModelParams.Model,
|
||||
"run result": runResult.EffectiveModelParams.Model,
|
||||
} {
|
||||
if model != "fallback-model" {
|
||||
t.Fatalf("%s model=%q, want fallback-model", name, model)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEngineSupportsConcurrentPrepareAndRun(t *testing.T) {
|
||||
engine, err := promptkit.NewEngine(promptkit.Config{},
|
||||
promptkit.WithPromptFS(contractPromptFS("prompt", "profile", "message"), "."),
|
||||
|
||||
227
types.go
227
types.go
@@ -51,7 +51,8 @@ const (
|
||||
// ValidationPassed means the generated output satisfied its contract.
|
||||
ValidationPassed ValidationStatus = "passed"
|
||||
// ValidationFailed means validation completed and rejected the generated
|
||||
// output. Engine.Run returns this status in a result, not as an error.
|
||||
// output. Engine.Run and Engine.RunPrepared return this status in a result,
|
||||
// not as an error.
|
||||
ValidationFailed ValidationStatus = "failed"
|
||||
// ValidationSkipped means ValidationNone selected no content check.
|
||||
ValidationSkipped ValidationStatus = "skipped"
|
||||
@@ -78,9 +79,11 @@ const (
|
||||
// RunRequest selects one prompt execution. It has no stable JSON
|
||||
// representation.
|
||||
//
|
||||
// Prepare and Run copy the request's maps, pointers, and nested
|
||||
// JSON-compatible values before using them. The caller may mutate the request
|
||||
// after either method returns.
|
||||
// Prepare, PrepareExecution, and Run copy the request's maps, pointers, and
|
||||
// nested JSON-compatible values before using them. The caller may mutate the
|
||||
// request after any method returns. A successful PrepareExecution retains its
|
||||
// own private execution snapshot for RunPrepared. Excessively deep or large
|
||||
// JSON-shaped values are rejected for safety.
|
||||
type RunRequest struct {
|
||||
// PromptID is the required non-empty prompt identifier.
|
||||
PromptID string
|
||||
@@ -98,12 +101,14 @@ type RunRequest struct {
|
||||
// opaque consumer metadata, not a credential, and may be exposed in
|
||||
// prepared values, results, collaborator requests, provider requests, and
|
||||
// provider observability. Callers should use stable, non-sensitive
|
||||
// identifiers. An overlong direct value makes Prepare or Run return an
|
||||
// error matching ErrInvalidRequest.
|
||||
// identifiers. An overlong direct value makes Prepare, PrepareExecution, or
|
||||
// Run return an error matching ErrInvalidRequest.
|
||||
SessionID string
|
||||
// APIKey is a request-scoped direct credential. It takes precedence over
|
||||
// APIKeyEnv, is passed to the selected LLMClient, and is never included in
|
||||
// prepared values, results, hashes, JSON, String, or GoString output.
|
||||
// prepared values, results, hashes, JSON, String, or GoString output. A
|
||||
// successful PrepareExecution retains it only in the opaque handle until
|
||||
// RunPrepared claims the handle or Discard invalidates it.
|
||||
APIKey string `json:"-"`
|
||||
// Inputs maps prompt input names to references. A nil or empty map is valid
|
||||
// only when the selected prompt and its templates require no inputs.
|
||||
@@ -112,16 +117,18 @@ type RunRequest struct {
|
||||
// empty maps are equivalent.
|
||||
Vars map[string]string
|
||||
// Execution optionally overrides individual execution settings. Nil uses
|
||||
// the selected profile over its backend, when any, and framework defaults.
|
||||
// the selected profile over its backend, when any, and the framework
|
||||
// baseline.
|
||||
Execution *ExecutionTargetOverride
|
||||
// Validation optionally replaces the prompt's complete output contract. It
|
||||
// does not merge individual fields. Nil uses the prompt contract.
|
||||
Validation *OutputContract
|
||||
}
|
||||
|
||||
// PreparedRun contains prepared prompt execution state. It does not include
|
||||
// resolved API key values, model output, validation results, or internal target
|
||||
// presence metadata. PreparedRun has a stable JSON representation.
|
||||
// PreparedRun contains prepared prompt execution state returned by
|
||||
// [Engine.Prepare] or [PreparedExecution.Details]. It does not include resolved
|
||||
// API key values, model output, validation results, or internal target presence
|
||||
// metadata. PreparedRun has a stable JSON representation.
|
||||
//
|
||||
// All maps, slices, pointers, and schema values are caller-owned copies. JSON
|
||||
// timestamps use RFC 3339 and zero timing values are omitted. Hash formats are
|
||||
@@ -139,9 +146,10 @@ type PreparedRun struct {
|
||||
// SelectedBackendID equals EffectiveModelParams.BackendID. It is empty for
|
||||
// an endpoint-only profile.
|
||||
SelectedBackendID string `json:"selected_backend_id,omitempty"`
|
||||
// EffectiveModelParams contains framework defaults overlaid by the selected
|
||||
// backend, profile, and then request overrides. It excludes resolved API-key
|
||||
// values.
|
||||
// EffectiveModelParams contains settings resolved from the framework timeout
|
||||
// baseline, selected backend, profile, and then request overrides. Unset
|
||||
// optional provider controls remain zero rather than reporting a provider
|
||||
// default. It excludes resolved API-key values.
|
||||
EffectiveModelParams ExecutionTarget `json:"effective_model_params"`
|
||||
// OutputContract is the complete effective output contract.
|
||||
OutputContract OutputContract `json:"output_contract"`
|
||||
@@ -154,7 +162,8 @@ 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 passes to the LLM client.
|
||||
// Messages are the rendered messages that Run or RunPrepared passes to the
|
||||
// LLM client.
|
||||
Messages []RenderedMessage `json:"messages"`
|
||||
// StartTime is the UTC time at which preparation began.
|
||||
StartTime time.Time `json:"start_time,omitempty"`
|
||||
@@ -214,12 +223,15 @@ type RunResult struct {
|
||||
InputHashes map[string]string `json:"input_hashes,omitempty"`
|
||||
// Usage is the token accounting reported by the LLM client.
|
||||
Usage TokenUsage `json:"usage"`
|
||||
// StartTime is the UTC time immediately before preparation begins.
|
||||
// StartTime is the UTC time immediately before ordinary Run preparation or
|
||||
// after RunPrepared claims its handle.
|
||||
StartTime time.Time `json:"start_time,omitempty"`
|
||||
// EndTime is the UTC time after generation and validation complete.
|
||||
EndTime time.Time `json:"end_time,omitempty"`
|
||||
// Duration covers preparation, generation, and validation. JSON represents
|
||||
// it as integer milliseconds in duration_ms and omits a zero value.
|
||||
// Duration covers preparation, generation, and validation for Run. For
|
||||
// RunPrepared it covers only the execution attempt after claim and excludes
|
||||
// preparation and consumer-held delay. JSON represents it as integer
|
||||
// milliseconds in duration_ms and omits a zero value.
|
||||
Duration time.Duration `json:"-"`
|
||||
}
|
||||
|
||||
@@ -231,8 +243,8 @@ type ArtifactRef struct {
|
||||
// URI is the file path for ArtifactRefFile and optional provenance metadata
|
||||
// for ArtifactRefInline.
|
||||
URI string
|
||||
// Body is the content for ArtifactRefInline and is ignored for
|
||||
// ArtifactRefFile.
|
||||
// Body is the content for ArtifactRefInline, where an empty value is valid,
|
||||
// and is ignored for ArtifactRefFile.
|
||||
Body string
|
||||
}
|
||||
|
||||
@@ -259,10 +271,11 @@ type Artifact struct {
|
||||
// ArtifactReader resolves a prompt input reference into its content.
|
||||
//
|
||||
// Read may be called concurrently. It must honor ctx cancellation to make
|
||||
// Prepare and Run responsive to cancellation. The engine passes a copied ref
|
||||
// and immediately copies the returned Artifact.Body; it does not retain either
|
||||
// value. Readers supply artifact metadata, and the engine assigns an input-map
|
||||
// name only when the returned artifact name is empty.
|
||||
// Prepare, PrepareExecution, and Run responsive to cancellation. The engine
|
||||
// passes a copied ref and immediately copies the returned Artifact.Body; it
|
||||
// does not retain either value. Readers supply artifact metadata, and the
|
||||
// engine assigns an input-map name only when the returned artifact name is
|
||||
// empty.
|
||||
//
|
||||
// An injected reader owns any application-specific path containment,
|
||||
// authorization, content-size, and content-type policy. It must protect
|
||||
@@ -284,17 +297,22 @@ type ExecutionTarget struct {
|
||||
// empty for endpoint-only profiles. It is supplied to injected LLMClient
|
||||
// implementations as part of the effective target.
|
||||
BackendID string `json:"backend_id,omitempty"`
|
||||
// Endpoint is the model-provider base URL.
|
||||
// Endpoint is the normalized absolute HTTP or HTTPS model-provider base URL.
|
||||
// It has a host and no user information, query, or fragment.
|
||||
Endpoint string `json:"endpoint"`
|
||||
// Model is the provider model identifier.
|
||||
Model string `json:"model"`
|
||||
// Temperature is the effective sampling temperature from 0 through 2.
|
||||
// Temperature is the resolved sampling temperature from 0 through 2. Zero
|
||||
// leaves the field unspecified to compatible providers unless the
|
||||
// corresponding ExecutionTargetPresence bit is true.
|
||||
Temperature float64 `json:"temperature"`
|
||||
// MaxTokens is the non-negative effective output-token limit. Zero leaves
|
||||
// the limit unspecified to compatible providers unless it was an explicit
|
||||
// request override.
|
||||
// MaxTokens is the non-negative resolved output-token limit. Zero leaves
|
||||
// the limit unspecified to compatible providers unless the corresponding
|
||||
// ExecutionTargetPresence bit is true.
|
||||
MaxTokens int `json:"max_tokens"`
|
||||
// TopP is the effective nucleus-sampling value from 0 through 1.
|
||||
// TopP is the resolved nucleus-sampling value from 0 through 1. Zero leaves
|
||||
// the field unspecified to compatible providers unless the corresponding
|
||||
// ExecutionTargetPresence bit is true.
|
||||
TopP float64 `json:"top_p"`
|
||||
// TimeoutSeconds is the non-negative per-generation deadline. Zero disables
|
||||
// this deadline without disabling caller cancellation or the transport cap.
|
||||
@@ -310,6 +328,63 @@ type ExecutionTarget struct {
|
||||
ExtraParams map[string]any `json:"extra_params"`
|
||||
}
|
||||
|
||||
// ProfileInspection is the caller-owned result of [Engine.InspectProfile].
|
||||
// It has no stable JSON representation.
|
||||
//
|
||||
// EffectiveModelParams contains a copied effective target. APIKeyRequired is
|
||||
// separate from that target to preserve ExecutionTarget's general execution
|
||||
// and stable JSON contracts.
|
||||
type ProfileInspection struct {
|
||||
// ProfileID is the trimmed, exact profile ID inspected by the engine.
|
||||
ProfileID string
|
||||
// EffectiveModelParams contains settings resolved from the framework timeout
|
||||
// baseline, selected backend, and then profile, without a request override.
|
||||
// Unset optional provider controls remain zero rather than reporting a
|
||||
// provider default. APIKeyEnv is an environment-variable name, never its
|
||||
// credential value.
|
||||
EffectiveModelParams ExecutionTarget
|
||||
// APIKeyRequired reports that a later execution must supply a direct API
|
||||
// key or an explicit request environment override. It is mutually exclusive
|
||||
// with a nonblank EffectiveModelParams.APIKeyEnv.
|
||||
APIKeyRequired bool
|
||||
}
|
||||
|
||||
// PromptInputDefinition describes one declared prompt input.
|
||||
// It has no stable JSON representation.
|
||||
type PromptInputDefinition struct {
|
||||
// Name is the normalized prompt input name.
|
||||
Name string
|
||||
// Required is the prompt definition's declared required flag. When true,
|
||||
// preparation fails if the input is omitted. A false value does not account
|
||||
// for input references in message or session-ID templates.
|
||||
Required bool
|
||||
// ContentType is the declared input media-type metadata.
|
||||
ContentType string
|
||||
// Description is the declared human-readable input description.
|
||||
Description string
|
||||
}
|
||||
|
||||
// PromptInspection is the caller-owned result of [Engine.InspectPrompt].
|
||||
// It has no stable JSON representation.
|
||||
//
|
||||
// Inputs contains copied declared input metadata in definition order.
|
||||
// OutputContract is the normalized contract declared by the prompt definition,
|
||||
// rather than a request-level effective override. PromptHash is opaque.
|
||||
type PromptInspection struct {
|
||||
// PromptID is the normalized ID of the selected prompt definition.
|
||||
PromptID string
|
||||
// PromptVersion is the normalized version of the selected prompt definition.
|
||||
PromptVersion string
|
||||
// PromptHash is the opaque equality value for the selected definition.
|
||||
PromptHash string
|
||||
// DefaultProfileID is declared metadata and is not resolved by inspection.
|
||||
DefaultProfileID string
|
||||
// Inputs contains caller-owned declared input metadata in definition order.
|
||||
Inputs []PromptInputDefinition
|
||||
// OutputContract is the normalized contract declared by the definition.
|
||||
OutputContract OutputContract
|
||||
}
|
||||
|
||||
// ExecutionTargetOverride represents per-request runtime setting overrides and
|
||||
// has no stable JSON representation.
|
||||
//
|
||||
@@ -317,19 +392,28 @@ type ExecutionTarget struct {
|
||||
// fields replace profile values and preserve explicit zero or empty values. A
|
||||
// non-empty ExtraParams map replaces the complete profile or backend map
|
||||
// rather than merging keys. Empty string fields, nil pointers, and a nil or
|
||||
// empty ExtraParams map inherit the selected profile over its backend, when
|
||||
// any, and framework defaults.
|
||||
// empty ExtraParams map inherit lower-precedence values. An optional provider
|
||||
// control that remains zero is unspecified; TimeoutSeconds retains its
|
||||
// framework deadline when no higher-precedence value is present.
|
||||
type ExecutionTargetOverride struct {
|
||||
// Endpoint replaces the profile or backend endpoint when non-empty without
|
||||
// changing the effective BackendID.
|
||||
// changing the effective BackendID. Preparation trims it and requires an
|
||||
// absolute HTTP or HTTPS URL with a host and no user information, query, or
|
||||
// fragment.
|
||||
Endpoint string
|
||||
// Model replaces the profile model when non-empty.
|
||||
Model string
|
||||
// Temperature, when non-nil, must point to a value from 0 through 2.
|
||||
// Temperature, when non-nil, must point to a value from 0 through 2. A
|
||||
// pointed-to zero is explicitly present; nil inherits a lower-precedence
|
||||
// value and otherwise leaves the provider control unspecified.
|
||||
Temperature *float64
|
||||
// MaxTokens, when non-nil, must point to a non-negative value.
|
||||
// MaxTokens, when non-nil, must point to a non-negative value. A pointed-to
|
||||
// zero is explicitly present; nil inherits a lower-precedence value and
|
||||
// otherwise leaves the provider control unspecified.
|
||||
MaxTokens *int
|
||||
// TopP, when non-nil, must point to a value from 0 through 1.
|
||||
// TopP, when non-nil, must point to a value from 0 through 1. A pointed-to
|
||||
// zero is explicitly present; nil inherits a lower-precedence value and
|
||||
// otherwise leaves the provider control unspecified.
|
||||
TopP *float64
|
||||
// TimeoutSeconds, when non-nil, must point to a non-negative value. A
|
||||
// pointed-to zero disables the per-generation deadline.
|
||||
@@ -348,7 +432,8 @@ type ExecutionTargetOverride struct {
|
||||
APIKeyEnv string
|
||||
// ExtraParams, when non-empty, replaces the complete profile or backend map.
|
||||
// Values must be JSON-compatible: nil, booleans, finite numbers, strings,
|
||||
// arrays or slices, and maps with non-empty string keys. Cycles are invalid.
|
||||
// arrays or slices, and maps with non-empty string keys. Cycles and
|
||||
// excessively deep or large values are invalid.
|
||||
ExtraParams map[string]any
|
||||
}
|
||||
|
||||
@@ -360,9 +445,12 @@ type ExecutionTargetOverride struct {
|
||||
// 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. Numeric
|
||||
// zero, blank strings, and an empty ExtraParams map inherit framework defaults;
|
||||
// use ExecutionTargetOverride pointer fields to request explicit numeric zero.
|
||||
// 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.
|
||||
type Profile struct {
|
||||
// ID is the required non-blank profile identifier. WithProfiles trims it.
|
||||
ID string
|
||||
@@ -372,23 +460,24 @@ type Profile struct {
|
||||
BackendID string
|
||||
// Endpoint is the model-provider base URL. It is required only when
|
||||
// BackendID is blank and otherwise overrides the backend endpoint when
|
||||
// non-blank.
|
||||
// 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 string
|
||||
// Temperature is from 0 through 2. Zero inherits the framework default.
|
||||
// Temperature is from 0 through 2. Zero leaves the provider control
|
||||
// unspecified.
|
||||
Temperature float64
|
||||
// MaxTokens is non-negative. Zero inherits the framework default.
|
||||
// MaxTokens is non-negative. Zero leaves the provider control unspecified.
|
||||
MaxTokens int
|
||||
// TopP is from 0 through 1. Zero inherits the framework default rather than
|
||||
// selecting an explicit zero.
|
||||
// TopP is from 0 through 1. Zero leaves the provider control unspecified
|
||||
// rather than selecting an explicit zero.
|
||||
TopP float64
|
||||
// TimeoutSeconds is non-negative. Zero inherits the framework default.
|
||||
// TimeoutSeconds is non-negative. Zero retains the framework deadline.
|
||||
TimeoutSeconds int
|
||||
// ServiceTier is optional; a blank value inherits the framework default.
|
||||
// ServiceTier is optional; a blank value leaves it unspecified.
|
||||
ServiceTier string
|
||||
// ReasoningEffort is optional; a blank value inherits the framework
|
||||
// default.
|
||||
// ReasoningEffort is optional; a blank value leaves it unspecified.
|
||||
ReasoningEffort string
|
||||
// APIKeyRequired clears a backend's inherited API-key environment name and
|
||||
// requires a non-blank RunRequest.APIKey unless the request explicitly
|
||||
@@ -396,7 +485,8 @@ type Profile struct {
|
||||
APIKeyRequired bool
|
||||
// ExtraParams contains provider-specific JSON-compatible values. An empty
|
||||
// map inherits backend request defaults, when any. WithProfiles validates
|
||||
// and deeply copies it during NewEngine.
|
||||
// and deeply copies it during NewEngine. Excessively deep or large values
|
||||
// are rejected for safety.
|
||||
ExtraParams map[string]any
|
||||
}
|
||||
|
||||
@@ -458,8 +548,8 @@ type ExecutionTargetPresence struct {
|
||||
// does not merge fields. The public Engine validates generated output once and
|
||||
// does not install an output repairer.
|
||||
type OutputContract struct {
|
||||
// Format selects generated artifact metadata. An empty effective value
|
||||
// defaults to FormatText.
|
||||
// Format selects generated artifact metadata. An empty value in a non-nil
|
||||
// request replacement defaults to FormatText.
|
||||
Format OutputFormat `json:"format"`
|
||||
// ValidationMode selects the content check. Use one of the declared
|
||||
// ValidationMode constants.
|
||||
@@ -467,8 +557,8 @@ type OutputContract struct {
|
||||
// SchemaPath is required when ValidationMode is ValidationJSONSchema and is
|
||||
// ignored by other modes.
|
||||
SchemaPath string `json:"schema_path"`
|
||||
// RepairAttempts is a requested repair limit. A non-positive value requests
|
||||
// no repairs. The public Engine performs no repairs even when this value is
|
||||
// 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 int `json:"repair_attempts"`
|
||||
}
|
||||
@@ -559,25 +649,26 @@ type StructuredOutputJSONSpec struct {
|
||||
Schema any `json:"schema"`
|
||||
}
|
||||
|
||||
// LLMClient executes rendered prompts for [Engine.Run].
|
||||
// LLMClient executes rendered prompts for [Engine.Run] and
|
||||
// [Engine.RunPrepared].
|
||||
//
|
||||
// Generate is scheduled according to the resolved backend's capacity policy.
|
||||
// It may still be called concurrently for different backend pools or unlimited
|
||||
// backends. Cancellation while waiting for capacity can prevent Generate from
|
||||
// being called. Once invoked, it must honor context cancellation to make Run
|
||||
// responsive to cancellation. The request and all nested maps, slices, and
|
||||
// pointers are client-owned copies and may be mutated or retained without
|
||||
// affecting engine state.
|
||||
// and RunPrepared responsive to cancellation. The request and all nested maps,
|
||||
// slices, and pointers are client-owned copies and may be mutated or retained
|
||||
// without affecting engine state.
|
||||
//
|
||||
// Generate receives rendered messages and may receive a direct API key. A
|
||||
// client must protect those values and any raw output in its logging, storage,
|
||||
// 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 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
|
||||
// Run.
|
||||
// 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.
|
||||
type LLMClient interface {
|
||||
Generate(context.Context, GenerateRequest) (*GenerateResponse, error)
|
||||
}
|
||||
@@ -613,8 +704,12 @@ type GenerateResponse struct {
|
||||
|
||||
// File returns a file-backed artifact reference whose URI is path.
|
||||
//
|
||||
// The default artifact reader opens path as a caller-selected operating-system
|
||||
// path without restricting it to an application root or imposing a size limit.
|
||||
// The default artifact reader accepts path only when it resolves to a regular
|
||||
// operating-system file, checking that condition before and after opening it.
|
||||
// It reads synchronously in bounded chunks and checks context cancellation
|
||||
// before opening, before and after each read, and before returning the
|
||||
// artifact; it cannot interrupt a filesystem operation already in progress.
|
||||
// It does not restrict path to an application root or impose a size limit.
|
||||
// Applications accepting untrusted paths must validate them before calling
|
||||
// Promptkit or use [WithArtifactReader] to enforce application policy.
|
||||
func File(path string) ArtifactRef {
|
||||
@@ -622,13 +717,13 @@ func File(path string) ArtifactRef {
|
||||
}
|
||||
|
||||
// Inline returns an inline artifact reference whose Body is body and whose URI
|
||||
// is empty.
|
||||
// is empty. An empty body is a valid, explicitly supplied input.
|
||||
func Inline(body string) ArtifactRef {
|
||||
return ArtifactRef{Type: ArtifactRefInline, Body: body}
|
||||
}
|
||||
|
||||
// InlineWithURI returns an inline artifact reference with body content and uri
|
||||
// provenance metadata.
|
||||
// provenance metadata. An empty body is a valid, explicitly supplied input.
|
||||
func InlineWithURI(uri string, body string) ArtifactRef {
|
||||
return ArtifactRef{Type: ArtifactRefInline, URI: uri, Body: body}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user