92 Commits

Author SHA1 Message Date
c53250f023 Prepare documentation for the v0.7.0 release 2026-08-25 02:13:43 +00:00
3d99483219 Preserve cancellation after profile lookups 2026-08-25 02:06:47 +00:00
67f788b1e2 Document profile inheritance behavior 2026-08-25 01:46:16 +00:00
764103a2e2 Resolve inherited profiles in engine workflows 2026-08-25 01:41:21 +00:00
a08dd83d1f Add profile inheritance resolver 2026-08-25 01:38:15 +00:00
e8922d8ec5 Add profile inheritance definition support 2026-08-25 01:34:40 +00:00
2d44305a8a Document optional API key environment behavior 2026-08-25 00:44:26 +00:00
a11c80291e Allow unauthenticated optional API key requests 2026-08-25 00:40:49 +00:00
3239567297 Make optional API key environments nonblocking 2026-08-25 00:39:26 +00:00
c239304c2a Document structured generation errors 2026-08-23 19:04:00 +00:00
159b02116f Expose structured generation errors to consumers 2026-08-23 19:01:38 +00:00
e5b7adfb49 Return structured errors for provider status failures 2026-08-23 18:58:15 +00:00
fa2384e696 Bound provider error response body reads 2026-08-23 18:56:15 +00:00
af0bd3f31a Add internal structured provider error parsing 2026-08-23 18:54:42 +00:00
5fff8cd623 Retire the completed Rakestrawhome backend roadmaps 2026-08-23 17:34:14 +00:00
b2a6c47778 Document the Rakestrawhome built-in backend 2026-08-23 17:19:12 +00:00
a1805fe550 Add the Rakestrawhome Gemma profile 2026-08-23 17:12:17 +00:00
93af155254 Add the Rakestrawhome built-in backend 2026-08-23 17:10:34 +00:00
d783b687a5 Plan the built-in Rakestrawhome backend 2026-08-23 17:00:01 +00:00
4ca3be2c14 Finish audit remediation and prepare v0.6.0 2026-08-12 13:01:05 +00:00
227fb35f99 Centralize the maintainer validation workflow 2026-08-12 00:04:48 +00:00
e291b8bfe9 Consolidate provider transport test scaffolding 2026-08-11 23:58:31 +00:00
2b6a7f83c4 Bound and strictly decode provider responses 2026-08-11 23:46:16 +00:00
3a43550f70 Validate and compose provider endpoints 2026-08-11 23:38:45 +00:00
c281f721bc Preserve transport error identities 2026-08-11 23:28:00 +00:00
350b0e76d9 Preserve repair settings and cumulative usage 2026-08-11 23:21:34 +00:00
e43350fd0d Make validation cancellation authoritative 2026-08-11 23:14:01 +00:00
20d3e3b5ee Escape schema resources and reuse compiled plans 2026-08-11 23:05:11 +00:00
a93b799236 Preserve exact JSON validation semantics 2026-08-11 22:55:17 +00:00
e83a3ce179 Honor cancellation while rendering prompts 2026-08-11 22:48:20 +00:00
a04a3bbc5f Correct artifact file and empty input handling 2026-08-11 22:41:37 +00:00
731b66cff5 Avoid decoding unrelated profiles 2026-08-11 22:32:33 +00:00
70e0ea0cf0 Correct profile source validation and identity 2026-08-11 22:28:48 +00:00
d45c474c1e Unify prompt repository source handling 2026-08-11 22:18:50 +00:00
25f1ba0b30 Correct prompt definition selection and decoding 2026-08-11 22:11:39 +00:00
a718762da1 Contain prompt content paths within source roots 2026-08-11 22:03:32 +00:00
58ac3ce298 Correct engine construction edge cases 2026-08-11 21:54:47 +00:00
57f2ce1ce4 Harden public ownership and diagnostic contracts 2026-08-11 21:45:54 +00:00
c8b6d5c490 Consolidate public JSON serialization 2026-08-11 21:40:12 +00:00
abeb50b525 Bound JSON-compatible value copying 2026-08-11 21:32:24 +00:00
1cb07c7d91 Centralize output contract validation 2026-08-11 21:21:33 +00:00
8cfc71c351 Centralize execution setting and session validation 2026-08-11 21:12:56 +00:00
5ccfa4a345 Plan the codebase audit remediation 2026-08-11 20:10:14 +00:00
14e03f19d0 Consolidate and close the codebase audit 2026-08-11 17:21:29 +00:00
7c562a9374 Document cross-cutting architecture audit findings 2026-08-11 17:10:59 +00:00
ef97d85ac9 Document repository-wide test strategy audit 2026-08-11 17:00:59 +00:00
805e48c873 Document capacity scheduling audit results 2026-08-11 16:51:28 +00:00
e9e126dcba Document OpenAI transport audit findings 2026-08-11 16:42:55 +00:00
32e7a3557c Document prepared execution lifecycle audit 2026-08-11 16:31:05 +00:00
1d1b04e2e0 Record ordinary execution and repair audit findings 2026-08-11 16:22:09 +00:00
9748897751 Document inspection and target resolution audit findings 2026-08-11 16:09:50 +00:00
9d020039d5 Record output validation audit findings 2026-08-11 15:59:46 +00:00
5247ce0b73 Record artifact loading and prompt rendering audit findings 2026-08-11 15:41:57 +00:00
c434aa1dae Record profile source audit findings 2026-08-11 15:29:23 +00:00
ac9b3f3d80 Record prompt source audit findings 2026-08-11 15:18:17 +00:00
4f12a89a1b Record backend registry and defaults audit findings 2026-08-11 15:05:15 +00:00
df31e7f58e Record domain and JSON value audit findings 2026-08-11 14:54:56 +00:00
0678d242b9 Record engine operation audit findings 2026-08-11 14:43:20 +00:00
3b4ea21208 Record engine construction audit findings 2026-08-11 14:36:33 +00:00
1430e85147 Record configuration and adapter audit findings 2026-08-11 14:28:07 +00:00
34d7a19da5 Record public value and error audit findings 2026-08-11 14:17:55 +00:00
ebf1602635 Prepare the codebase audit plan 2026-08-11 14:05:30 +00:00
31f2ce3a09 Document Promptkit v0.5.0 2026-08-01 13:18:46 +00:00
fd06e4ca6b Clean up roadmap and fallback profile guidance 2026-08-01 13:16:26 +00:00
e63b8de1e9 Complete application fallback profile implementation 2026-08-01 12:39:01 +00:00
9354d2b373 Add application fallback profile sources 2026-08-01 12:37:01 +00:00
01ca5430bd Move profile composition to the engine facade 2026-08-01 12:31:50 +00:00
ae2179d103 Complete optional parameter omission 2026-08-01 02:43:02 +00:00
a248433d0f Omit unset optional request parameters 2026-08-01 02:41:44 +00:00
bd6cffc9d0 Prepare documentation for Promptkit v0.4.0 2026-07-30 23:48:25 +00:00
e40c4f182b Document structured capacity errors 2026-07-30 23:26:44 +00:00
7428e50c2c Expose structured capacity errors 2026-07-30 23:24:04 +00:00
63c67a4520 Add internal capacity error identity 2026-07-30 23:21:27 +00:00
25a7052a3d Clarify prompt input requirement documentation 2026-07-30 22:10:33 +00:00
fc3255967e Document prompt inspection API 2026-07-30 21:06:33 +00:00
e920168b30 Expose prompt inspection through the engine 2026-07-30 21:03:47 +00:00
272b6a4bc1 Add internal prompt inspection 2026-07-30 21:00:05 +00:00
dde48a31fc Document profile inspection API 2026-07-30 19:57:07 +00:00
242eace4a7 Expose profile inspection through the engine 2026-07-30 19:52:00 +00:00
0bf5f88136 Add internal profile inspection resolution 2026-07-30 19:48:09 +00:00
369ab5392d Add feature roadmap and implementation plan for profile inspection API 2026-07-30 19:42:51 +00:00
2ba0146e5d Tighten prepared execution credential handling 2026-07-30 19:06:17 +00:00
6112c2af0c Document prepared execution workflow 2026-07-30 18:27:03 +00:00
f5e12c00f5 Expose prepared execution handles 2026-07-30 18:20:35 +00:00
49fe402dd2 Add prepared execution lifecycle to the runner 2026-07-30 18:10:28 +00:00
c301eb8d55 Add frozen validation preparation plans 2026-07-30 17:59:45 +00:00
c13e9710d9 Add feature roadmap and implementation plan for downstream consumer wishlist items 2026-07-30 17:54:44 +00:00
87b5ec3d75 Organize downstream feature requests in the future roadmap 2026-07-30 17:13:26 +00:00
cb4028a637 Add feature roadmaps with wishlists from downstream consumers 2026-07-30 16:59:37 +00:00
5a1bff4529 Prepare the local backend convenience release 2026-07-30 04:01:19 +00:00
805a7f965d Document local backend configuration paths 2026-07-30 03:40:52 +00:00
147f5e5ff5 Add local backend convenience constructor 2026-07-30 03:38:23 +00:00
103 changed files with 15393 additions and 4089 deletions

View File

@@ -33,6 +33,21 @@ boundary and constraints that framework work must preserve.
## Release Guidance ## Release Guidance
Consumers upgrading from `v0.6.0` to `v0.7.0` should read the
[v0.7.0 changelog and migration guide](docs/releases/v0.7.0.md).
Earlier adopters can consult the
[v0.6.0 changelog and migration guide](docs/releases/v0.6.0.md).
Consumers upgrading from `v0.4.0` to `v0.5.0` can consult the
[v0.5.0 changelog and migration guide](docs/releases/v0.5.0.md).
Consumers upgrading from `v0.3.0` to `v0.4.0` should read the
[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 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). [v0.2.0 changelog and migration guide](docs/releases/v0.2.0.md).

View File

@@ -9,50 +9,83 @@ import (
// backend. // backend.
const BackendOpenRouter = backend.OpenRouterID const BackendOpenRouter = backend.OpenRouterID
// BackendRakestrawHome is the reserved ID of Promptkit's built-in
// Rakestrawhome backend.
const BackendRakestrawHome = backend.RakestrawHomeID
// BackendLocal is the case-sensitive conventional ID used by [LocalBackend].
// It is not a built-in or reserved backend and must be registered with
// [WithBackend].
const BackendLocal = "local"
// Backend configures one engine-scoped OpenAI-compatible backend. // Backend configures one engine-scoped OpenAI-compatible backend.
// //
// Backend has no stable JSON representation. Use keyed literals so additions // Backend has no stable JSON representation. Use keyed literals so additions
// to this configuration value do not break source compatibility. // to this configuration value do not break source compatibility.
type Backend struct { type Backend struct {
// ID is the stable, case-sensitive registry key. NewEngine trims it and // ID is the stable, case-sensitive registry key. NewEngine trims it and
// requires a non-blank value. BackendOpenRouter is reserved. // requires a non-blank value. Built-in backend IDs are reserved.
ID string ID string
// Endpoint is the OpenAI-compatible base endpoint. NewEngine trims it and // Endpoint is the OpenAI-compatible base endpoint. NewEngine trims it and
// requires an absolute HTTP or HTTPS URL with a host and without user // requires an absolute HTTP or HTTPS URL with a host and without user
// information, a query string, or a fragment. Paths are allowed. // information, a query string, or a fragment. Paths are allowed.
Endpoint string Endpoint string
// APIKeyEnv optionally names the environment variable containing the API // APIKeyEnv optionally names an environment lookup source for an API key.
// key. NewEngine trims it and requires the portable form // NewEngine trims it and requires the portable form [A-Za-z_][A-Za-z0-9_]*.
// [A-Za-z_][A-Za-z0-9_]*. Store only the name, never a credential value. // A direct RunRequest.APIKey takes precedence. When no usable credential is
// available, the built-in client omits Authorization; injected clients own
// their own credential-resolution behavior. Store only the name, never a
// credential value.
APIKeyEnv string APIKeyEnv string
// ExtraParams contains backend-wide request defaults. Values must be // ExtraParams contains backend-wide request defaults. Values must be
// JSON-compatible, finite, acyclic, and keyed by non-empty strings. Keys // JSON-compatible, finite, acyclic, and keyed by non-empty strings. Keys
// must not be model, session_id, messages, temperature, max_tokens, top_p, // must not be model, session_id, messages, temperature, max_tokens, top_p,
// service_tier, reasoning_effort, or response_format. An empty map supplies // 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 ExtraParams map[string]any
// ConcurrencyLimit is the maximum number of simultaneous model-generation // ConcurrencyLimit is the maximum number of simultaneous model-generation
// calls allowed for this backend within one Engine. Zero leaves the backend // calls allowed for this backend within one Engine. Zero leaves the backend
// unlimited. A negative value makes NewEngine fail with ErrInvalidConfig. // unlimited. A negative value makes NewEngine fail with ErrInvalidConfig.
ConcurrencyLimit int ConcurrencyLimit int
// QueueCapacity controls how many additional Run calls may be admitted // QueueCapacity controls how many additional Run or RunPrepared calls may
// beyond ConcurrencyLimit. Nil uses 1024 when ConcurrencyLimit is positive; // be admitted beyond ConcurrencyLimit. Nil uses 1024 when ConcurrencyLimit
// a pointer uses its exact value, including zero. The pointed-to value must // is positive; a pointer uses its exact value, including zero. The pointed-to
// be non-negative, and QueueCapacity must be nil when ConcurrencyLimit is // value must be non-negative, and QueueCapacity must be nil when
// zero. Their sum must fit in an int. WithBackend copies the value and does // ConcurrencyLimit is zero. Their sum must fit in an int. WithBackend copies
// not retain the pointer. // the value and does not retain the pointer.
QueueCapacity *int 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. // WithBackend adds one Backend registration to the constructed Engine.
// //
// Registrations accumulate in option order. Every normalized ID must be unique // Registrations accumulate in option order. Every normalized ID must be unique
// across consumer registrations and built-ins; a duplicate or invalid // across consumer registrations and built-ins; a duplicate or invalid
// definition makes NewEngine fail with ErrInvalidConfig. In particular, // definition makes NewEngine fail with ErrInvalidConfig. Built-in IDs,
// BackendOpenRouter cannot be replaced. The immutable registration is scoped // including [BackendOpenRouter] and [BackendRakestrawHome], cannot be
// to the resulting Engine and cannot be enumerated, replaced, removed, or // replaced. The immutable registration is scoped to the resulting Engine and
// mutated after construction. WithBackend does not install package-global // cannot be enumerated, replaced, removed, or mutated after construction.
// state. // WithBackend does not install package-global state.
func WithBackend(backend Backend) Option { func WithBackend(backend Backend) Option {
queueCapacity := 0 queueCapacity := 0
queueCapacitySet := backend.QueueCapacity != nil queueCapacitySet := backend.QueueCapacity != nil

View File

@@ -60,7 +60,18 @@ func TestEngineRejectsRunBeforeCompletionWhenAdmissionIsFull(t *testing.T) {
awaitCapacitySignal(t, reader.entered, "first artifact read") 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 { if result != nil {
t.Fatalf("capacity rejection returned partial result: %+v", result) 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) { if errors.Is(err, promptkit.ErrInvalidRequest) || errors.Is(err, promptkit.ErrLLMGenerate) {
t.Fatalf("capacity rejection had an unrelated category: %v", err) 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 { if calls := reader.callCount(); calls != 1 {
t.Fatalf("artifact calls=%d, want only the admitted run", calls) t.Fatalf("artifact calls=%d, want only the admitted run", calls)
} }
@@ -174,6 +200,19 @@ func TestCapacityExceededSentinelContract(t *testing.T) {
if promptkit.ErrCapacityExceeded == nil { if promptkit.ErrCapacityExceeded == nil {
t.Fatal("ErrCapacityExceeded is 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{ for _, unrelated := range []error{
promptkit.ErrInvalidConfig, promptkit.ErrInvalidConfig,
promptkit.ErrInvalidRequest, promptkit.ErrInvalidRequest,
@@ -181,7 +220,8 @@ func TestCapacityExceededSentinelContract(t *testing.T) {
promptkit.ErrValidation, promptkit.ErrValidation,
} { } {
if errors.Is(promptkit.ErrCapacityExceeded, unrelated) || 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) t.Fatalf("ErrCapacityExceeded aliases unrelated sentinel %v", unrelated)
} }
} }

40
capacity_error.go Normal file
View 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
}

View File

@@ -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 { func fromDomainExecutionTargetPresence(presence domain.ExecutionTargetPresence) ExecutionTargetPresence {
return ExecutionTargetPresence{ return ExecutionTargetPresence{
Temperature: presence.Temperature, Temperature: presence.Temperature,

47
doc.go
View File

@@ -3,23 +3,29 @@
// //
// Applications construct an [Engine] with [NewEngine], select filesystem or // Applications construct an [Engine] with [NewEngine], select filesystem or
// in-memory sources and optional engine-scoped [Backend] registrations, and // in-memory sources and optional engine-scoped [Backend] registrations, and
// call [Engine.Prepare] or [Engine.Run]. Concrete registries, repositories, // call [Engine.InspectPrompt], [Engine.InspectProfile], [Engine.Prepare],
// validators, and the built-in OpenAI-compatible client remain internal // [Engine.PrepareExecution], [Engine.Run], or [Engine.RunPrepared]. Concrete
// implementation details. // registries, repositories, validators, and the built-in OpenAI-compatible
// client remain internal implementation details.
// //
// # Concurrency and ownership // # Concurrency and ownership
// //
// An Engine supports concurrent Prepare and Run calls. Engine-local backend // An Engine supports concurrent InspectPrompt, InspectProfile, Prepare,
// policies bound admitted Run calls and model generations where configured, // PrepareExecution, Run, and RunPrepared calls. Engine-local backend policies
// while different backend pools and unlimited backends continue independently. // bound admitted Run and RunPrepared calls and model generations where
// An injected [LLMClient] or [ArtifactReader] can therefore still receive // configured, while different backend pools and unlimited backends continue
// concurrent calls and must be safe for that use. // 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 // NewEngine copies in-memory profiles and backend definitions. Prepare,
// copy request maps, slices, pointer values, and JSON-compatible extra // PrepareExecution, and Run copy request maps, slices, pointer values, and
// parameters before using them. Returned values and values passed to extension // JSON-compatible extra parameters before using them. InspectPrompt and
// interfaces are likewise isolated from engine state. Callers own those copies // InspectProfile return copied inspection values. Returned values and values
// and may mutate them after the call that supplied or returned them. // 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. [CapacityError] values are caller-owned and may be mutated
// without affecting engine state or another error. Immutable [GenerationError]
// values are also caller-owned and do not retain shared engine state.
// //
// # Security and sensitive data // # Security and sensitive data
// //
@@ -45,10 +51,17 @@
// [GenerateResponse], [ExecutionTargetPresence], and the string value types // [GenerateResponse], [ExecutionTargetPresence], and the string value types
// used by those values. // used by those values.
// //
// Construction values, including [Config], [Backend], [RunRequest], // Construction, inspection, handle, and error values, including [Config],
// [ArtifactRef], [ExecutionTargetOverride], [Profile], and // [Backend], [RunRequest], [ArtifactRef], [ExecutionTargetOverride], [Profile],
// [OpenAICompatibleProfileConfig], do not have stable JSON representations. // [OpenAICompatibleProfileConfig], [ProfileInspection],
// Direct API keys are nevertheless excluded from JSON for every public value. // [PromptInputDefinition], [PromptInspection], [PreparedExecution],
// [CapacityError], and [GenerationError], do not have stable JSON
// representations. Direct API keys are nevertheless excluded from JSON for
// every public value.
// Provider-derived [GenerationError] accessor values are untrusted and can
// contain sensitive request or schema fragments. Applications must apply their
// own disclosure policy before logging, displaying, or returning them.
// //
// JSON timestamps use time.Time's RFC 3339 encoding and are omitted when zero. // JSON timestamps use time.Time's RFC 3339 encoding and are omitted when zero.
// PreparedRun and RunResult durations are encoded as integer milliseconds in // PreparedRun and RunResult durations are encoded as integer milliseconds in

View File

@@ -40,11 +40,67 @@ validation, and default transport behavior. Source discovery, format
validation, and profile precedence are defined by the validation, and profile precedence are defined by the
[framework format reference](../formats.md). [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 ## Prepare Without Model Execution
[`Engine.Prepare`](../../engine.go) resolves the selected prompt and profile, [`Engine.Prepare`](../../engine.go) resolves the selected prompt and profile,
loads inputs and any structured-output schema, and renders messages without 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 ```go
prepared, err := engine.Prepare(ctx, promptkit.RunRequest{ 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 inline input. Exact request requirements and prepared-result fields belong to
the [`RunRequest` and `PreparedRun` GoDoc](../../types.go). 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 ## Execute And Validate
[`Engine.Run`](../../engine.go) performs the same preparation, invokes the [`Engine.Run`](../../engine.go) performs the same preparation, invokes the
configured model client, classifies the generated artifact, and validates the configured model client, classifies the generated artifact, and validates the
content. A completed content check may return `ValidationFailed` in the result; content in one call. Choose it when the application does not need a preflight
an operational inability to validate returns an error. 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 The maintained
[offline execution example](../../examples/go-library/run/main.go) injects a [offline execution example](../../examples/go-library/run/main.go) injects a
@@ -90,13 +182,90 @@ replace execution settings or the complete output contract.
The [public value GoDoc](../../types.go) defines nil, empty, zero, replacement, The [public value GoDoc](../../types.go) defines nil, empty, zero, replacement,
copy, and credential behavior. The copy, and credential behavior. The
[framework format reference](../formats.md) defines how those request values [framework format reference](../formats.md) defines how those request values
interact with prompt definitions, file-backed profiles, built-ins, schemas, interact with prompt definitions, file-backed and application fallback
and framework defaults. profiles, built-ins, schemas, and framework defaults.
For programmatic profiles, For programmatic profiles,
[`OpenAICompatibleProfile`](../../profiles.go) converts ordinary [`OpenAICompatibleProfile`](../../profiles.go) converts ordinary
OpenAI-compatible settings into a value accepted by `WithProfiles`. OpenAI-compatible settings into a value accepted by `WithProfiles`.
### Alias A Built-In Profile
Give an application-owned profile ID a built-in base when prompts should select
the application ID while inheriting the built-in target. The child can override
only the setting it owns:
```go
promptkit.WithProfiles(promptkit.Profile{
ID: "weather-light",
BaseProfileID: "deepseek-4-flash",
ReasoningEffort: "high",
})
```
Select `weather-light` in a prompt or `RunRequest.ProfileID`; it remains the
reported selected profile. See the [profile inheritance format
reference](../formats.md#profile-inheritance) and the
[`Profile` GoDoc](../../types.go) for exact lookup, merging, and validation
behavior.
### Use The Rakestrawhome Built-In Profile
Set `RAKESTRAWHOME_INFERENCE_API_KEY` in the application environment, then
select `rakestrawhome-gemma-4-31b` as an ordinary profile ID. For example, a
prepared result identifies the selected built-in through
`BackendRakestrawHome`:
```go
prepared, err := engine.Prepare(ctx, promptkit.RunRequest{
PromptID: "meeting.summary",
ProfileID: "rakestrawhome-gemma-4-31b",
Inputs: inputs,
})
if err != nil {
return err
}
if prepared.SelectedBackendID != promptkit.BackendRakestrawHome {
return fmt.Errorf("unexpected backend %q", prepared.SelectedBackendID)
}
```
Do not register `rakestrawhome` manually. When adopting this built-in, remove
an existing `WithBackend` registration with that exact ID; retaining it causes
the intentional duplicate-ID configuration error. Direct request credentials
and runtime endpoint overrides remain supported under their ordinary GoDoc and
format contracts.
### Inspect A Profile Before Prompt Work
Use [`Engine.InspectProfile`](../../engine.go) to validate one configured
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 != "" {
// This is a configured optional environment lookup source.
} 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. A reported `APIKeyEnv` is a configured optional source, while
`APIKeyRequired` is the explicit local requirement. The
[credential format reference](../formats.md#credentials) and the method's
[GoDoc](../../engine.go) own the exact precedence, timing, result, and error
contracts.
### Set A Per-Run Session And Reasoning ### Set A Per-Run Session And Reasoning
Supply a direct session ID when one prompt should be correlated with a Supply a direct session ID when one prompt should be correlated with a
@@ -124,35 +293,86 @@ providers. The
[`RunRequest` and `ExecutionTargetOverride` GoDoc](../../types.go) owns the [`RunRequest` and `ExecutionTargetOverride` GoDoc](../../types.go) owns the
exact normalization, precedence, error, copying, and exposure contract. 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 Choose the smallest configuration that fits how the endpoint will be reused.
profile. This local backend limits model generation to two simultaneous calls;
because `QueueCapacity` is omitted, the engine admits up to 1024 additional #### Use An Endpoint-Only Profile
calls waiting behind them:
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 ```go
engine, err := promptkit.NewEngine(promptkit.Config{ engine, err := promptkit.NewEngine(promptkit.Config{
PromptDir: "prompts", PromptDir: "prompts",
}, },
promptkit.WithBackend(promptkit.Backend{ promptkit.WithProfiles(promptkit.Profile{
ID: "local", ID: "local-summary",
Endpoint: "http://localhost:8000/v1", Endpoint: "http://localhost:8000/v1",
APIKeyEnv: "LOCAL_LLM_API_KEY", Model: "example-model",
ConcurrencyLimit: 2,
}), }),
)
```
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{ promptkit.WithProfiles(promptkit.Profile{
ID: "local-summary", ID: "local-summary",
BackendID: "local", BackendID: promptkit.BackendLocal,
Model: "example-model", 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. Registrations belong to one engine and custom IDs cannot replace built-ins.
The [`Backend` and `WithBackend` GoDoc](../../backends.go) defines validation, The [`Backend`, `LocalBackend`, and `WithBackend` GoDoc](../../backends.go)
copying, uniqueness, exact concurrency-field semantics, and request-default defines exact construction, validation, copying, uniqueness, concurrency, and
behavior. request-default behavior.
Both file-backed and in-memory profiles select a registration through Both file-backed and in-memory profiles select a registration through
`backend` or `Profile.BackendID`. Profile and request endpoint overrides retain `backend` or `Profile.BackendID`. Profile and request endpoint overrides retain
@@ -173,8 +393,8 @@ zero:
```go ```go
noWaiting := 0 noWaiting := 0
backend := promptkit.Backend{ backend := promptkit.Backend{
ID: "local", ID: "local-gpu",
Endpoint: "http://localhost:8000/v1", Endpoint: "http://gpu-host:8000/v1",
ConcurrencyLimit: 2, ConcurrencyLimit: 2,
QueueCapacity: &noWaiting, QueueCapacity: &noWaiting,
} }
@@ -236,6 +456,11 @@ When a limited backend has admitted all active and waiting calls, handle
```go ```go
result, err := engine.Run(ctx, request) result, err := engine.Run(ctx, request)
if errors.Is(err, promptkit.ErrCapacityExceeded) { 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. // Apply application policy: shed work, report overload, or retry later.
} }
``` ```
@@ -243,8 +468,27 @@ if errors.Is(err, promptkit.ErrCapacityExceeded) {
A rejected call returns no partial result and does not invoke the model 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 client. Promptkit does not prescribe retries or map this error to an HTTP
status; those choices remain with the consuming application. The status; those choices remain with the consuming application. The
[`Engine.Run` and error GoDoc](../../engine.go) owns exact error and [`CapacityError` GoDoc](../../capacity_error.go) owns the exact typed-error
cancellation identities. contract, while the [`Engine.Run` and error GoDoc](../../engine.go) owns broad
error and cancellation identities.
For a non-2xx response from the built-in OpenAI-compatible client, inspect the
status and deliberately selected provider diagnostic when useful:
```go
var generationErr *promptkit.GenerationError
if errors.As(err, &generationErr) {
status := generationErr.StatusCode()
message := generationErr.ProviderMessage()
_, _ = status, message // Apply application retry and presentation policy.
}
```
All provider fields are untrusted and can contain sensitive request or schema
fragments. Do not log, display, or return them without an application-specific
disclosure policy. Promptkit does not assign retry or presentation behavior.
The [`GenerationError` GoDoc](../../generation_error.go) owns the exact typed
error contract.
## Application Boundary ## Application Boundary

View File

@@ -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 placeholder documents for packages, APIs, or integrations that do not yet
exist. exist.
## Maintainer-Run Validation ## Maintainer Validation
Promptkit does not currently use hosted CI. Maintainers are responsible for This section is the canonical local validation workflow for Promptkit. Run
running the documented checks before accepting changes. Run the default Go every command from the repository root before accepting a change. The test
validation from the Promptkit repository root: 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 ```sh
go test ./... go test ./...
@@ -57,85 +63,171 @@ go test -race ./...
go vet ./... go vet ./...
go build ./... go build ./...
go run ./examples/go-library/prepare 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 ```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 ### Local Markdown Links
Markdown link and confirm its target exists. Finally, check whitespace:
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 ```sh
git diff --check 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 Inspect `git status --short --untracked-files=all` and the complete diff before
requires link validation and `git diff --check`. Run the Go validation whenever accepting a change. The status may contain only the intended source changes
documentation changes commands, examples, generated output, or another during development. Reject credentials, private keys, environment files,
behavior checked by the module. 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 After committing the accepted change, require a clean candidate:
Use focused checks while iterating, then run the complete validation sequence
before accepting the change. The root package supports:
```sh ```sh
go test . test -z "$(git status --porcelain)"
go vet .
go build .
``` ```
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.

View File

@@ -9,9 +9,10 @@ explains how to select these sources and invoke the engine. The
owns the resulting outbound wire behavior. owns the resulting outbound wire behavior.
Prompt and profile sources recursively discover files ending in `.yaml` or Prompt and profile sources recursively discover files ending in `.yaml` or
`.yml`. YAML decoding is strict: unknown fields are errors for the selected `.yml`. Each prompt-definition and profile file contains exactly one YAML
definition. Definitions are selected by their YAML `id`, not their file name document; comments and trailing whitespace are allowed. YAML decoding is
or directory. strict: unknown fields are errors for the selected definition. Definitions are
selected by their YAML `id`, not their file name or directory.
## Prompt Definitions ## 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 one definition. When it supplies a version, the ID and version pair must be
unique. 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 ### Inputs
Each `inputs` item has these fields: 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`, containing an inline Go template; or
- `content_file`, naming a file whose contents are the Go template. - `content_file`, naming a file whose contents are the Go template.
For directory and `fs.FS` prompt sources, `content_file` resolves relative to `content_file` must be a relative path. It resolves from the directory that
the prompt file and remains within the source root. `WithPromptFile` also contains the prompt file and must remain within the configured prompt source
resolves it relative to that file. 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 Request variables are the template data, so a variable named `audience` is
referenced as `{{.audience}}`. The `{{input "note"}}` helper renders the body referenced as `{{.audience}}`. The `{{input "note"}}` helper renders the body
@@ -131,6 +144,17 @@ A request-level `OutputContract` replaces the complete prompt output contract.
It does not merge individual fields. If its format is empty, Promptkit uses It does not merge individual fields. If its format is empty, Promptkit uses
`text`. `text`.
## Built-In Backends
Every engine provides these reserved OpenAI-compatible backend IDs. Consumers
must not register either ID with `WithBackend`; exact registration and
reservation behavior belongs to the [`Backend` GoDoc](../backends.go).
| ID | Base endpoint | API-key environment variable | Active generation limit | Default queue capacity |
| --- | --- | --- | ---: | ---: |
| `openrouter` | `https://openrouter.ai/api/v1` | `OPENROUTER_API_KEY` | 16 | 1024 |
| `rakestrawhome` | `https://inference.ai.rakestrawhome.com/v1` | `RAKESTRAWHOME_INFERENCE_API_KEY` | 4 | 1024 |
## Profile Definitions ## Profile Definitions
A profile supplies model execution settings: A profile supplies model execution settings:
@@ -149,11 +173,21 @@ extra_params:
provider_option: enabled provider_option: enabled
``` ```
A derived profile can use a named base and override only the settings it owns:
```yaml
id: local-summary-fast
base_profile: local-summary
timeout_seconds: 30
reasoning_effort: low
```
| Field | Required | Meaning | | Field | Required | Meaning |
| --- | --- | --- | | --- | --- | --- |
| `id` | yes | Non-empty profile identifier. IDs must be unique within one source. | | `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. | | `base_profile` | no | One optional parent profile ID. A derived profile may inherit target fields from it. |
| `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. | | `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. | | `model` | yes | Non-empty provider model name. |
| `temperature` | no | Number from 0 through 2. | | `temperature` | no | Number from 0 through 2. |
| `max_tokens` | no | Integer zero or greater. | | `max_tokens` | no | Integer zero or greater. |
@@ -161,16 +195,21 @@ extra_params:
| `timeout_seconds` | no | Per-generation deadline in whole seconds; integer zero or greater. | | `timeout_seconds` | no | Per-generation deadline in whole seconds; integer zero or greater. |
| `service_tier` | no | Provider-specific request tier. | | `service_tier` | no | Provider-specific request tier. |
| `reasoning_effort` | no | Provider-specific reasoning setting. | | `reasoning_effort` | no | Provider-specific reasoning setting. |
| `api_key_env` | no | Name of an environment variable containing the API key. | | `api_key_env` | no | Optional environment-variable lookup source for an API key. |
| `extra_params` | no | JSON-compatible provider-specific outbound fields. | | `extra_params` | no | JSON-compatible provider-specific outbound fields. |
Raw `api_key` is prohibited in profile YAML. Store only an environment Raw `api_key` is prohibited in profile YAML. Store only an environment
variable name in `api_key_env`. variable name in `api_key_env`.
A standalone profile must provide a model and at least one of `backend` or
`endpoint`. A derived profile may omit those target fields because its selected
base chain can provide them. Local parsing still validates a derived profile's
own ID, supplied endpoint, execution-setting bounds, and `extra_params`.
Promptkit does not infer a backend from a model or endpoint. Endpoint-only Promptkit does not infer a backend from a model or endpoint. Endpoint-only
profiles remain supported and have no effective backend ID. profiles remain supported and have no effective backend ID.
The engine always provides the built-in `openrouter` ID. Consumers can add The engine always provides the built-in `openrouter` and `rakestrawhome` IDs.
engine-scoped IDs with Consumers can add engine-scoped IDs with
[`WithBackend`](../backends.go); exact registration validation belongs to its [`WithBackend`](../backends.go); exact registration validation belongs to its
GoDoc. GoDoc.
@@ -178,30 +217,33 @@ GoDoc.
objects with string keys. Keys must be non-empty. With the built-in client, 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 they also cannot collide with the standard fields listed in the
[outbound request contract](integrations/openai-compatible-chat.md#request-body). [outbound request contract](integrations/openai-compatible-chat.md#request-body).
Excessively deep or large JSON-shaped values are rejected for safety.
### Defaults And Overrides ### Defaults And Overrides
Execution settings resolve in this order: Execution settings resolve in this order:
1. framework defaults; 1. the framework timeout baseline;
2. the selected backend, when the profile names one; 2. the selected backend, when the profile names one;
3. the selected profile; and 3. the selected profile; and
4. request `ExecutionTargetOverride` values. 4. request `ExecutionTargetOverride` values.
The framework defaults are: The framework baseline is:
| Setting | Default | | Setting | Default |
| --- | --- | | --- | --- |
| `temperature` | `0` | | `temperature` | Unspecified and omitted from compatible provider requests unless a profile or runtime override selects it. |
| `max_tokens` | `0` | | `max_tokens` | Unspecified and omitted from compatible provider requests unless a profile or runtime override selects it. |
| `top_p` | `1` | | `top_p` | Unspecified and omitted from compatible provider requests unless a profile or runtime override selects it. |
| `timeout_seconds` | `600` | | `timeout_seconds` | `600` |
Numeric zero in a file or in-memory profile means that the profile does not Numeric zero in a file or in-memory profile does not select a numeric value.
replace the framework default. Numeric request overrides use pointers, so an For `temperature`, `max_tokens`, and `top_p`, it leaves the provider control
explicit zero is preserved. In particular, an explicit request unspecified. For `timeout_seconds`, it retains the framework deadline. Numeric
`timeout_seconds` of zero disables the per-generation deadline while leaving request overrides use pointers, so an explicit zero is retained and sent to
the caller context and transport timeout intact. 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 Non-empty profile strings replace backend defaults, and non-empty request
strings replace both. Request reasoning is the exception: a nil strings replace both. Request reasoning is the exception: a nil
@@ -219,54 +261,86 @@ defines how the effective settings are serialized.
### Source And Profile Precedence ### Source And Profile Precedence
An explicit request profile ID takes precedence over the prompt's 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: Profile sources resolve matching IDs in this order:
1. in-memory profiles supplied with `WithProfiles`; 1. in-memory profiles supplied with `WithProfiles`;
2. a profile file, `fs.FS`, or configured profile directory; and 2. the ordinary configured source selected by a profile file, `fs.FS`, or
3. embedded built-in profiles. 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 A profile source supplies a complete definition; definitions and their fields
invalid matching profile is an error and does not fall back. In-memory are not merged across sources. A higher-precedence source falls back only when
`Profile` values follow the same ranges as YAML profiles. They use the requested profile ID is absent. An invalid matching profile is an error and
`APIKeyRequired` for request-scoped credentials instead of `api_key_env`. 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.
When a selected definition names `base_profile`, every profile ID in that
chain is looked up through this same precedence order. A higher-precedence
definition therefore shadows a lower-precedence definition of the same base
ID, including a built-in. References are not source-qualified.
### Profile Inheritance
Promptkit resolves one linear base chain of at most 32 profiles, including the
selected profile. It merges settings from the root base to the selected leaf.
The leaf's `id` remains the selected profile identity. Nonblank string fields
(`backend`, `endpoint`, `model`, `service_tier`, `reasoning_effort`, and
`api_key_env`) and nonzero numeric fields replace inherited values. A nonempty
`extra_params` map replaces the complete inherited map rather than merging
keys, and `APIKeyRequired: true` remains true through the chain. Backend and
endpoint are independent: replacing one does not clear the other.
There is no profile-level clearing syntax. Blank strings, zero numbers, false,
and empty maps remain unspecified and inherit from a base. Use existing
presence-aware request overrides where an execution needs an explicit zero or
empty reasoning setting.
An absent directly selected profile reports the ordinary not-found error. Once
the selected profile exists, a missing base, cycle, overlong chain, or
incomplete resolved target is a profile-load failure. Ordinary operations
resolve chains afresh; prepared execution retains the fully resolved target.
## Built-In Profile Catalog ## Built-In Profile Catalog
Every built-in selects the `openrouter` backend. The engine's built-in backend Every built-in profile selects one maintained built-in backend and inherits
registry supplies `https://openrouter.ai/api/v1` and the environment-variable that backend's connection and credential metadata. Profile files do not repeat
name `OPENROUTER_API_KEY`, so individual profiles contain only model and those values. A configured, application fallback, or in-memory profile with
generation settings. Built-in profile files do not repeat those connection the same profile ID takes precedence.
values. A custom or in-memory profile with the same profile ID takes
precedence.
| Provider | ID | Model | | Provider | ID | Backend | Model |
| --- | --- | --- | | --- | --- | --- | --- |
| aion-labs | `aion-2` | `aion-labs/aion-2.0` | | aion-labs | `aion-2` | `openrouter` | `aion-labs/aion-2.0` |
| anthropic | `claude-fable-latest` | `~anthropic/claude-fable-latest` | | anthropic | `claude-fable-latest` | `openrouter` | `~anthropic/claude-fable-latest` |
| anthropic | `claude-haiku-latest` | `~anthropic/claude-haiku-latest` | | anthropic | `claude-haiku-latest` | `openrouter` | `~anthropic/claude-haiku-latest` |
| anthropic | `claude-opus-latest` | `~anthropic/claude-opus-latest` | | anthropic | `claude-opus-latest` | `openrouter` | `~anthropic/claude-opus-latest` |
| anthropic | `claude-sonnet-latest` | `~anthropic/claude-sonnet-latest` | | anthropic | `claude-sonnet-latest` | `openrouter` | `~anthropic/claude-sonnet-latest` |
| deepseek | `deepseek-3-2` | `deepseek/deepseek-v3.2` | | deepseek | `deepseek-3-2` | `openrouter` | `deepseek/deepseek-v3.2` |
| deepseek | `deepseek-4-flash` | `deepseek/deepseek-v4-flash` | | deepseek | `deepseek-4-flash` | `openrouter` | `deepseek/deepseek-v4-flash` |
| deepseek | `deepseek-4-pro` | `deepseek/deepseek-v4-pro` | | deepseek | `deepseek-4-pro` | `openrouter` | `deepseek/deepseek-v4-pro` |
| google | `gemini-2-flash` | `google/gemini-2.5-flash` | | google | `gemini-2-flash` | `openrouter` | `google/gemini-2.5-flash` |
| google | `gemini-2-flash-lite` | `google/gemini-2.5-flash-lite` | | google | `gemini-2-flash-lite` | `openrouter` | `google/gemini-2.5-flash-lite` |
| google | `gemini-2-pro` | `google/gemini-2.5-pro` | | google | `gemini-2-pro` | `openrouter` | `google/gemini-2.5-pro` |
| google | `gemini-3-flash-lite` | `google/gemini-3.1-flash-lite` | | google | `gemini-3-flash-lite` | `openrouter` | `google/gemini-3.1-flash-lite` |
| google | `gemini-flash-latest` | `~google/gemini-flash-latest` | | google | `gemini-flash-latest` | `openrouter` | `~google/gemini-flash-latest` |
| google | `gemini-pro-latest` | `~google/gemini-pro-latest` | | google | `gemini-pro-latest` | `openrouter` | `~google/gemini-pro-latest` |
| google | `gemma-4-31b` | `google/gemma-4-31b-it:exacto` | | google | `gemma-4-31b` | `openrouter` | `google/gemma-4-31b-it:exacto` |
| minimax | `minimax-m2` | `minimax/minimax-m2.5` | | google | `rakestrawhome-gemma-4-31b` | `rakestrawhome` | `google/gemma-4-31b-it` |
| minimax | `minimax-m3` | `minimax/minimax-m3` | | minimax | `minimax-m2` | `openrouter` | `minimax/minimax-m2.5` |
| mistral | `mistral-large-2512` | `mistralai/mistral-large-2512` | | minimax | `minimax-m3` | `openrouter` | `minimax/minimax-m3` |
| mistral | `mistral-medium-3-5` | `mistralai/mistral-medium-3-5` | | mistral | `mistral-large-2512` | `openrouter` | `mistralai/mistral-large-2512` |
| mistral | `mistral-small-3` | `mistralai/mistral-small-3.2-24b-instruct` | | mistral | `mistral-medium-3-5` | `openrouter` | `mistralai/mistral-medium-3-5` |
| mistral | `mistral-small-4` | `mistralai/mistral-small-2603` | | mistral | `mistral-small-3` | `openrouter` | `mistralai/mistral-small-3.2-24b-instruct` |
| nvidia | `nemotron-3-ultra` | `nvidia/nemotron-3-ultra-550b-a55b` | | mistral | `mistral-small-4` | `openrouter` | `mistralai/mistral-small-2603` |
| openai | `gpt-5-mini` | `openai/gpt-5.4-mini` | | nvidia | `nemotron-3-ultra` | `openrouter` | `nvidia/nemotron-3-ultra-550b-a55b` |
| openai | `gpt-5-nano` | `openai/gpt-5.4-nano` | | openai | `gpt-5-mini` | `openrouter` | `openai/gpt-5.4-mini` |
| openai | `gpt-5-nano` | `openrouter` | `openai/gpt-5.4-nano` |
## Schemas ## Schemas
@@ -285,16 +359,25 @@ schema produces a failed validation result.
Credential values belong at the request or environment boundary, never in Credential values belong at the request or environment boundary, never in
prompt, profile, schema, or example files: prompt, profile, schema, or example files:
- a file profile names an environment variable with `api_key_env`; - a backend or file profile can name an optional environment lookup source
- an in-memory profile may set `APIKeyRequired`; with `APIKeyEnv` or `api_key_env`;
- a request can provide a direct `APIKey` or override `APIKeyEnv`; and - an in-memory profile may set `APIKeyRequired` as an explicit local
requirement;
- a request can provide a direct `APIKey` or override the optional `APIKeyEnv`
source; and
- a direct request key takes precedence over environment lookup. - a direct request key takes precedence over environment lookup.
After a direct request key, the credential-source precedence is request After a direct request key, the credential-source precedence is request
`APIKeyEnv`, profile `api_key_env`, then the backend default. An in-memory `APIKeyEnv`, profile `api_key_env`, then the backend default. An in-memory
profile with `APIKeyRequired` clears an inherited backend environment name and profile with `APIKeyRequired` clears an inherited backend environment name and
requires a direct key unless the request explicitly supplies `APIKeyEnv`. requires a direct key unless the request explicitly supplies `APIKeyEnv`.
Promptkit validates required credential availability during preparation. Named environment sources are optional: when the selected source is absent,
empty, or whitespace-only, the built-in client omits the `Authorization`
header and handles the provider response normally. `APIKeyRequired` is the
only explicit local availability requirement. Promptkit validates required
credential availability during preparation and rechecks it when a prepared
execution runs. Injected clients receive resolved source metadata but define
their own credential-resolution behavior.
Direct keys are excluded from JSON results and redacted by public string Direct keys are excluded from JSON results and redacted by public string
formatters. Environment-variable names may appear in prepared metadata, but formatters. Environment-variable names may appear in prepared metadata, but
their values do not. their values do not.

View File

@@ -14,10 +14,16 @@ that produce these outbound settings.
Generation sends an HTTP `POST` with `Content-Type: application/json`. Generation sends an HTTP `POST` with `Content-Type: application/json`.
Before the client is called, the engine resolves framework, backend, profile, Before the client is called, the engine resolves framework, backend, profile,
and request values into one execution target. A non-empty endpoint from that and request values into one execution target. Endpoint configuration is trimmed
target overrides the client's configured base URL. After trailing slashes are and must be an absolute HTTP or HTTPS URL with a host and without user
removed, `/chat/completions` is appended. Generation fails before sending when information, a query, or a fragment. A non-empty endpoint from the target
neither source supplies an endpoint. 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 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 injected clients. The built-in client does not derive the URL from that ID and
@@ -25,11 +31,13 @@ does not serialize it in the provider request.
## Authentication ## Authentication
A non-empty API key supplied directly on the execution target takes A usable API key supplied directly on the execution target takes precedence.
precedence. Otherwise, when an API-key environment-variable name is supplied, Otherwise, when an API-key environment-variable name is supplied, the client
the client reads that variable and requires a non-empty value. The selected reads and trims that variable. A bearer header is sent only when the resolved
key is sent as `Authorization: Bearer <key>`. No authorization header is sent direct or environment credential is non-empty. When neither source is usable,
when neither mechanism is configured. the client omits `Authorization` and handles the provider response normally.
An explicitly required target with no usable source is rejected before
transport.
The target contains the already resolved environment-variable name: an The target contains the already resolved environment-variable name: an
explicit request override takes precedence over profile metadata, which takes explicit request override takes precedence over profile metadata, which takes
@@ -53,8 +61,9 @@ never also sent as a session header.
The client conditionally includes: The client conditionally includes:
- `temperature`, `max_tokens`, and `top_p` when non-zero or explicitly - `temperature`, `max_tokens`, and `top_p` only when selected by a profile or
present; runtime override, including an explicit runtime zero; they are absent when
unspecified;
- non-empty `service_tier` and effective `reasoning_effort`; an explicitly - non-empty `service_tier` and effective `reasoning_effort`; an explicitly
disabled reasoning setting is empty and therefore omitted; and disabled reasoning setting is empty and therefore omitted; and
- `response_format` for JSON Schema structured output, including its name, - `response_format` for JSON Schema structured output, including its name,
@@ -81,13 +90,45 @@ request fields.
## Response Handling ## Response Handling
Any 2xx response is decoded as an OpenAI-compatible chat response. The client Any 2xx response body is limited to 16 MiB (16,777,216 bytes). A larger
returns the first choice's non-empty message content and maps prompt, declared `Content-Length` is rejected before the body is read, and streamed,
completion, total, cached, and cache-write token counts. 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 The bounded body must contain exactly one OpenAI-compatible JSON response
responses. For a non-2xx status, the error includes the status code but never object followed only by JSON whitespace and EOF. The client returns the first
the provider response body. 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, Promptkit recognizes one JSON document with a top-level
object-valued `error` member. Its optional `message` and `type` fields must be
strings, and `code` may be a string or JSON number. Valid supported fields are
handled independently, numeric codes retain their JSON number text, and
unknown fields are ignored. Missing, invalid, malformed, or multiply framed
envelopes contribute no provider detail.
Non-success bodies have a 65,536-byte limit. A larger declared
`Content-Length` is not read; otherwise the client reads at most one additional
byte to detect streamed or underreported overflow. Empty, unreadable,
oversized, malformed, and unrecognized bodies retain only the received status.
The body is always closed and no oversized stream is drained beyond that probe.
Extracted strings are made valid UTF-8, trimmed, and converted to one line by
collapsing Unicode whitespace, control, and format-character runs. Blank
values are omitted. Codes and types longer than 256 Unicode code points are
omitted; messages longer than 4,096 code points are truncated at a code-point
boundary with an ellipsis inside the limit. Promptkit never exposes raw bodies,
headers, endpoints, credentials, request data, schemas, generated content, or
unsupported provider metadata through this handling.
An outbound `http.Client.Do` failure retains both Promptkit's request-failure
identity and the exact transport error for `errors.Is` and `errors.As` checks.
The rendered error does not include the selected endpoint, request headers,
request content, credentials, or provider body.
## Timeout And Cancellation ## Timeout And Cancellation
@@ -102,5 +143,7 @@ Timeouts are layered:
timeout when the supplied value is not positive. timeout when the supplied value is not positive.
The earliest applicable caller, generation, or transport deadline controls the The earliest applicable caller, generation, or transport deadline controls the
request. Constructing the internal client does not mutate a supplied request. Caller cancellation retains `context.Canceled`; caller, generation,
`http.Client`. 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`.

View File

@@ -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 count, active count, and ordered waiter list. Pool state exists only for the
lifetime of its engine. lifetime of its engine.
## Bounded Run Admission ## Bounded Execution Admission
The runner asks the manager to admit a run after resolving the prompt, profile, For ordinary `Run`, the runner asks the manager to admit after resolving the
selected backend, effective execution target, credentials, and output contract, prompt, profile, selected backend, effective execution target, credentials, and
but before schema loading, artifact loading, or rendering. Admission is output contract, but before schema loading, artifact loading, or rendering.
immediate: a limited pool either reserves a slot or returns the internal `PrepareExecution` performs no admission. `RunPrepared` claims its handle,
`ErrCapacityExceeded` identity. The root facade maps that identity to the rechecks credential availability, and then asks the manager to admit the
public error without treating it as an invalid request or generation failure. 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 The total admitted bound is the active-generation limit plus its configured
waiting capacity. The returned release function is idempotent. The runner waiting capacity. The returned release function is idempotent. The runner
defers it as soon as admission succeeds and holds the lease across remaining defers it as soon as admission succeeds. An ordinary run holds the lease across
preparation, initial generation, validation, every repair attempt, and all remaining preparation, initial generation, validation, every repair attempt,
failure or cancellation exits. A repair is part of its original admission and and all failure or cancellation exits. Prepared execution holds the normal
does not reserve another bounded slot. 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 ## FIFO Generation Permits
`NewClient` wraps the engine's selected internal model client after public `NewClient` wraps the engine's selected internal model client after public
client adaptation or built-in client construction. Initial generation and the 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 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 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 The wrapper passes generation requests, responses, and collaborator errors
through unchanged. It owns scheduling only; the concrete model client remains 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 ## Cancellation And Release
@@ -90,11 +102,17 @@ unlimited admission. The
FIFO transfer, canceled-waiter removal, grant/cancel races, independent pools, FIFO transfer, canceled-waiter removal, grant/cancel races, independent pools,
unlimited calls, passthrough behavior, and panic release. unlimited calls, passthrough behavior, and panic release.
The [runner tests](../../internal/usecase/runner_test.go) own early admission, The [runner tests](../../internal/usecase/runner_test.go) own ordinary early
lease lifetime, failure release, and shared initial/repair scheduling. The 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 [external package capacity tests](../../capacity_contract_test.go) own the
assembled public-engine behavior for configured limits, capacity errors, assembled public-engine behavior for configured limits, capacity errors,
endpoint identity, engine independence, and injected clients. The 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 [root error-boundary tests](../../errors_internal_test.go) own preservation of
the public generation category and context identity when generation is the public generation category and context identity when generation is
canceled. canceled.

View File

@@ -25,16 +25,21 @@ and request precedence. The client uses its endpoint, credential metadata,
generation fields, and extra parameters. `BackendID` remains routing metadata generation fields, and extra parameters. `BackendID` remains routing metadata
for the generation boundary and is not mapped into the provider payload. for the generation boundary and is not mapped into the provider payload.
Construction validates the configured base URL and clones any supplied Construction trims and validates a nonempty configured base URL and clones any
`http.Client` so Promptkit can apply its timeout default without mutating the supplied `http.Client` so Promptkit can apply its timeout default without
caller's client. Generation then: 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; 2. maps the internal request into the OpenAI-compatible chat payload;
3. validates and merges extra parameters; 3. validates and merges extra parameters;
4. resolves authentication; 4. composes `/chat/completions` through parsed URL path operations;
5. performs the outbound request under the applicable deadlines; and 5. resolves authentication;
6. decodes the first response choice and token usage. 6. performs the outbound request under the applicable deadlines; and
7. decodes one strictly framed, size-bounded successful response object and
maps its first choice and token usage, or decodes bounded structured
non-success detail.
`internal/llm` owns the set of reserved OpenAI-compatible request fields used `internal/llm` owns the set of reserved OpenAI-compatible request fields used
when validating extra parameters. Backend registration consumes the same rule when validating extra parameters. Backend registration consumes the same rule
@@ -43,16 +48,63 @@ without making the model client depend on registry configuration.
The implementation has no retry loop, tool-call support, provider catalog, The implementation has no retry loop, tool-call support, provider catalog,
inbound HTTP behavior, or durable session store. 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 a frozen credential
environment-variable name only when the target explicitly requires a
credential. The handle does not retain the environment value; the model client
resolves the value visible when generation begins. For optional sources with no
usable value, the built-in client omits `Authorization` and continues to the
provider. A direct request key remains in private execution state only until
the claimed execution finishes or an unclaimed handle is discarded. Exact
public ownership and redaction semantics belong to the
[`PreparedExecution` GoDoc](../../prepared_execution.go).
## Failure Categories ## Failure Categories
The package preserves distinct error identities for invalid client The package preserves distinct error identities for invalid client
configuration, invalid generation requests, request execution failures, configuration, invalid generation requests, request execution failures,
non-success provider statuses, and malformed successful responses. Provider non-success provider statuses, and malformed successful responses. Provider
response bodies are not included in non-success errors. response bodies are never exposed in raw form through non-success errors.
Caller cancellation and deadline failures during the outbound request are Invalid nonempty configured endpoints are configuration failures. A missing or
reported as request execution failures. The runner classifies these identities invalid final selected endpoint is an invalid generation request and is
without depending on HTTP status mapping. rejected before transport.
Authentication resolves a trimmed direct key before a trimmed configured
environment value. Optional missing, empty, or whitespace-only sources do not
block transport and produce no `Authorization` header. An explicitly required
target with no usable source is rejected before transport with the existing
invalid-request diagnostics.
Successful response bodies have a fixed 16 MiB limit enforced by declared
length and by reading at most one byte beyond the boundary. The decoder accepts
exactly one JSON object plus trailing whitespace and EOF. Size overflow,
truncation, malformed JSON, trailing data, and a second value are malformed
responses with no partial result or provider content in the error. Every body
is closed, and an unbounded oversized stream is not drained.
For a non-success response, `ProviderHTTPError` retains the HTTP status and
only normalized detail from the bounded recognized envelope. It retains
`ErrUnexpectedStatus` through unwrapping. The client owns response closure;
its bounded reader and parser never close or drain a body themselves. The root
facade converts this concrete internal error into the public
[`GenerationError`](../../generation_error.go), while arbitrary injected-client
errors continue through the ordinary generation-error mapping unchanged.
An `http.Client.Do` failure is represented by a redacting multi-cause error:
the package request-failure sentinel and the exact returned transport error are
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 ## Test Ownership
@@ -60,7 +112,15 @@ The
[OpenAI-compatible client tests](../../internal/llm/openai_compatible_client_test.go) [OpenAI-compatible client tests](../../internal/llm/openai_compatible_client_test.go)
own configuration, client cloning, deterministic deadline precedence, own configuration, client cloning, deterministic deadline precedence,
authentication, request and response mapping, malformed data, error identity, authentication, request and response mapping, malformed data, error identity,
cancellation, and response-body suppression. The root transport contract test cancellation, endpoint selection and composition, pre-transport rejection, and
also verifies that resolved backend settings reach this client without bounded single-document successful-response framing, closure, and
serializing backend identity. All use local test servers or test transports; response-body suppression. The focused
the default suite makes no live or paid provider requests. [provider HTTP error tests](../../internal/llm/provider_http_error_test.go)
own envelope parsing, normalization, and bounded-reader cases; their
[transport tests](../../internal/llm/provider_http_error_transport_test.go)
own non-success response closure and integration. Root transport contract tests
own public `GenerationError` conversion, while also verifying that resolved
backend settings reach this client without serializing backend identity and
that ordinary-run cancellation retains its public generation and context
identities. All use local test servers or controlled test transports; the
default suite makes no live or paid provider requests.

View File

@@ -11,23 +11,23 @@ contributor workflow and validation.
| Component | Implemented responsibility | References | | 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 and generation 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/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) | | `examples/go-library/run` | Demonstrates an offline downstream consumer using a prompt file, in-memory profile, inline input, an injected deterministic model client, and `Run`. It is not a public library package. | [Example program](../../examples/go-library/run/main.go) |
| `internal/backend` | Constructs each engine's immutable registry from the built-in OpenRouter definition and consumer additions, validates and defensively copies definitions through the shared JSON-value package, and consumes the LLM-owned OpenAI-compatible reserved request-field rule. | [Backend registry](../../internal/backend/registry.go) | | `internal/backend` | Constructs each engine's immutable registry from the maintained built-in definitions and consumer additions, validates and defensively copies definitions through the shared JSON-value package, and consumes the LLM-owned OpenAI-compatible reserved request-field rule. | [Backend registry](../../internal/backend/registry.go) |
| `internal/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/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. | [Domain declarations](../../internal/domain/domain.go) | | `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/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/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/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` | Loads strictly decoded, locally validated execution profiles from filesystem and `fs.FS` sources, overlays raw sources with error-preserving fallback, and resolves inherited profiles. | [Framework formats](../formats.md), [profile repositories](../../internal/profile/filesystem_repository.go), [internal sources](sources.md#profiles-and-built-ins) |
| `internal/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 maintained built-in backends. | [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/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/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/llm` | Defines the internal generation boundary and implements outbound OpenAI-compatible chat requests from resolved execution targets, including bounded structured non-success response decoding, successful-response decoding, authentication, deadline handling, and ownership of the OpenAI-compatible reserved request-field policy. | [Internal model client](llm.md) |
| `internal/usecase` | Resolves 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 The root package assembles these internal components without exposing their
representations. Consumers depend only on the root facade. representations. Consumers depend only on the root facade.

View File

@@ -20,14 +20,43 @@ and override semantics consumed by the runner.
profiles, backend resolution, artifacts, rendering, model generation, and profiles, backend resolution, artifacts, rendering, model generation, and
validation. The root engine supplies one immutable registry containing the validation. The root engine supplies one immutable registry containing the
built-in backend and validated consumer additions, one engine-local run built-in backend and validated consumer additions, one engine-local run
admitter, and a model client wrapped by the same capacity manager. Schema admitter, and a model client wrapped by the same capacity manager. Validation
documents are loaded through the validator's optional schema-loader interface. plans and provider-facing schema metadata come from the validator's preparation
interface.
An output repairer can be injected internally, but the ordinary runner An output repairer can be injected internally, but the ordinary runner
constructor does not enable one. constructor does not enable one.
Each invocation carries its state in request, prepared-run, and result values. Each invocation carries its state in request, prepared-run, and result values.
The runner has no durable run or session store. 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 ## Shared Preparation Pipeline
`Prepare` and `Run` share one private preparation pipeline split at the point `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, 5. resolve application-neutral defaults, backend defaults, profile values,
and explicit request overrides in that order; and explicit request overrides in that order;
6. validate endpoint, model, numeric overrides, and credential requirements; 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 8. retain the definition, source identities, effective settings, output
contract, and preparation start time in invocation-local state. contract, and preparation start time in invocation-local state.
The completion phase consumes that state without reloading the prompt, The completion phase consumes that state without reloading the prompt,
profile, or backend: 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; 2. load and hash input artifacts;
3. render messages and the prompt-defined session; 3. render messages and the prompt-defined session;
4. apply any direct session ID; 4. apply any direct session ID;
@@ -59,7 +90,10 @@ profile, or backend:
`Run` performs backend admission between the phases. This structure preserves `Run` performs backend admission between the phases. This structure preserves
one execution-precedence and error-ordering implementation while allowing a one execution-precedence and error-ordering implementation while allowing a
full backend pool to reject work before expensive schema, artifact, and 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 Pointer-based numeric overrides preserve an explicit zero. Invalid negative or
out-of-range values fail as invalid requests. Endpoint overrides do not change 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` admitter is an internal unlimited fallback. After successful admission, `Run`
immediately defers the returned release function, performs the completion immediately defers the returned release function, performs the completion
phase, makes one initial generation call, builds the named output artifact, phase, makes one initial generation call, builds the named output artifact,
and validates that artifact. Invalid generated content remains a validation and validates that artifact with the plan compiled during completion. Invalid
result; an inability to perform validation is an operational error. 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, The admission lease covers completion-phase preparation, initial generation,
validation, every repair, and every exit. It bounds accepted work without 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 When an internal repairer is present, a JSON or JSON Schema content failure can
trigger bounded repair attempts. Repair receives the effective execution trigger bounded repair attempts. Repair receives the effective execution
target and session ID, validation errors, prior output, and structured-output target, explicit numeric-presence bits, credential, backend identity, session
specification. The default repairer uses the same wrapped client as initial ID, validation errors, prior output, and structured-output specification. One
generation, so each repair reacquires the selected backend's active permit request constructor supplies those common fields to initial and repair
while remaining inside its original admission lease. Repair never performs a generation while their rendered prompts remain intentionally distinct. The
second bounded admission. This capability remains internal and is not a public default repairer uses the same wrapped client as initial generation, so each
option. 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 A successful result includes the output artifact and raw output, validation
state, effective session ID, prompt and rendered-prompt hashes, selected state, effective session ID, prompt and rendered-prompt hashes, selected
profile and backend, effective settings, input hashes, token usage, a generated profile and backend, effective settings, input hashes, token usage, a generated
run identifier, and UTC timing. The same effective session reaches initial run identifier, and UTC timing. The same effective session reaches initial
generation and any repair attempt through the rendered prompt. The same generation and any repair attempt through the rendered prompt. The same
effective target, including backend identity, reaches generation and any effective target and presence metadata, including backend identity and direct
repair attempt. 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 ## 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 public facade and retains collaborator identities where they are part of the
internal contract. internal contract.
Admission capacity exhaustion retains the internal capacity identity and adds Admission capacity exhaustion retains the internal capacity identity. At the
the selected backend ID as context. It is not recategorized as an invalid use-case boundary, the runner attaches the selected backend ID in an internal
request or generation failure, and no partial result is returned. A context typed error, and the root facade copies that value into the public
already done at admission retains its context identity directly. Cancellation [`CapacityError`](../../capacity_error.go) without parsing diagnostic text. It
while waiting for an active generation permit prevents client invocation when is not recategorized as an invalid request or generation failure, and no
it wins the grant race; the model-client boundary then preserves the context partial result is returned. A context already done at admission retains its
error through the generation-failure category. Deferred release restores the context identity directly. Cancellation while waiting for an active generation
admission lease on preparation, generation, validation, repair, and permit prevents client invocation when it wins the grant race; the model-client
cancellation failures. 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 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 overlong direct session is an invalid request before source loading, while
an invalid or overlong prompt session template remains a prompt-render failure. an invalid or overlong prompt session template remains a prompt-render failure.
An unknown selected backend, or a selected backend with no configured resolver, 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, selection and override precedence, the two-phase boundary, early admission,
lease lifetime and release, direct-session resolution, schema-before-generation lease lifetime and release, direct-session resolution, schema-before-generation
behavior, hashing, generation and validation outcomes, backend propagation, behavior, hashing, generation and validation outcomes, backend propagation,
bounded repair, shared initial/repair capacity, credentials and redaction, bounded repair progression, initial/repair request parity, cumulative usage,
error categories, artifact metadata, usage, and timing. The shared initial/repair capacity, credentials and redaction, error categories,
artifact metadata, and timing. The
[capacity subsystem document](capacity.md) identifies the focused pool, [capacity subsystem document](capacity.md) identifies the focused pool,
waiter, and wrapped-client tests. waiter, and wrapped-client tests.

View File

@@ -12,9 +12,26 @@ validation modes, built-in catalog, and source precedence.
## Prompt Definitions ## Prompt Definitions
`internal/promptdef` discovers YAML deterministically, decodes and validates `internal/promptdef` uses one source-neutral flow for prompt selection and
definitions, selects an ID and optional version, and resolves file-backed normalization. That flow scans normalized YAML ID and version metadata,
message content within the selected operating-system or `fs.FS` source. 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, Its package tests own prompt selection, strict decoding, definition validation,
duplicate detection, and source containment: duplicate detection, and source containment:
@@ -22,30 +39,69 @@ duplicate detection, and source containment:
## Profiles And Built-Ins ## Profiles And Built-Ins
`internal/profile` loads and validates execution profiles from an `internal/profile` loads, locally validates, overlays, and resolves execution
operating-system filesystem or an `fs.FS`. It supports a primary repository profiles from an operating-system filesystem or an `fs.FS`. A file contains
with fallback only when the primary reports that a profile is absent. Strict exactly one YAML document and its trimmed YAML `id` is its only selection
YAML decoding recognizes the optional `backend` field, trims its value, and identity; filenames do not confer authority. Each point lookup reads discovered
requires a model plus at least one non-blank backend or endpoint. Loading does files once for their metadata and reuses the selected file's bytes for strict
not check registry membership because the available registry belongs to the decoding; unrelated profiles are not fully decoded. Strict selected decoding
assembled engine; the runner checks membership during preparation. recognizes `base_profile` and the optional `backend` field, trims their values,
and permits inherited target fields only when a base is named. File-backed
`extra_params` values are validated and defensively copied through the shared
bounded JSON-value owner before a profile is published. OpenAI-compatible
reserved-field policy remains with the model-client and backend-registry owners.
`internal/profile/builtin` embeds the maintained built-in profile catalog and The overlay repository consults the next repository only when the
can place a caller-selected repository ahead of that catalog. Every embedded higher-precedence repository reports that a profile is absent. A reliably
profile selects `openrouter` and inherits its endpoint and credential selected malformed profile stops fallback, while an unrelated malformed file
environment-variable name from the built-in backend registry rather than does not become authoritative through its filename. Loading does not check
repeating those values. Profile behavior is owned by the backend registry membership because the available registry belongs to the
[profile repository tests](../../internal/profile/repository_test.go), while assembled engine; the runner checks membership during preparation and exact
catalog completeness, the backend-selection invariant, duplicate IDs, and profile inspection.
The root engine assembles one raw composite catalog 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. One outer resolving repository wraps that complete raw catalog, so
each base lookup observes the same precedence and shadowing rules.
The resolving repository traverses every selected chain afresh, retains no
cache, detects cycles, limits a chain to 32 profiles, merges root-to-leaf into a
new caller-owned value, and validates the final target before publishing it. It
does not check backend registry membership. Exact `base_profile` syntax, merge
rules, and consumer-visible failure behavior belong to the [framework format
reference](../formats.md#profile-inheritance).
Exact profile inspection performs one point-in-time resolved lookup through
those profile sources and checks the final target without reading prompt, input,
or schema sources. It does not retain that lookup for a later execution.
Prepared execution instead freezes the fully resolved target; a later ordinary
operation performs a fresh traversal.
`internal/profile/builtin` embeds the maintained built-in profile catalog.
Every embedded profile selects a maintained built-in backend and inherits that
backend's 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 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). [built-in repository tests](../../internal/profile/builtin/repository_test.go).
## Ordinary Artifacts ## Ordinary Artifacts
`internal/artifact` resolves inline references and unrestricted, `internal/artifact` accepts explicitly typed inline references even when their
caller-selected file paths. It copies content into an artifact, records body is empty. It also resolves unrestricted, caller-selected paths only when
metadata and a content hash, applies a content-type fallback, and honors they identify regular operating-system files, checking that condition before
context cancellation. 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 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 particular, it does not constrain files to an application root or impose an
@@ -57,8 +113,17 @@ implemented reader behavior and failures.
## Rendering ## Rendering
`internal/prompt` renders definition messages as Go templates using named `internal/prompt` renders definition messages as Go templates using named
artifacts and variables. It carries message roles, session IDs, and cache artifacts and variables. Within one render, each referenced artifact body is
control into the rendered prompt. The 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. [renderer tests](../../internal/prompt/renderer_test.go) own rendering behavior.
## Schemas And Output Validation ## Schemas And Output Validation
@@ -68,6 +133,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 result; inability to load, register, or compile a schema is an operational
error. 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 The [validator tests](../../internal/validate/standard_validator_test.go) own
basic, JSON, JSON Schema, source resolution, schema loading, compilation, and basic, JSON, JSON Schema, source resolution, schema loading, compilation,
content-failure behavior. 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).

View File

@@ -19,10 +19,10 @@ results, public values, extension interfaces, profiles, and error sentinels.
The implemented internal components consist of: The implemented internal components consist of:
- `internal/domain`, which owns framework data values shared by later internal - `internal/domain`, which owns framework data values and source-neutral
components; invariants shared by later internal components;
- `internal/backend`, which owns validated immutable OpenAI-compatible backend - `internal/backend`, which owns validated immutable OpenAI-compatible backend
definitions and the built-in OpenRouter definition; definitions and the maintained built-in definitions;
- `internal/capacity`, which owns engine-local bounded run admission and - `internal/capacity`, which owns engine-local bounded run admission and
model-generation scheduling for limited backends; model-generation scheduling for limited backends;
- `internal/defaults`, which owns application-neutral framework defaults and - `internal/defaults`, which owns application-neutral framework defaults and
@@ -92,6 +92,13 @@ coordinates internal components and adapts the supported public extension
interfaces to narrow internal abstractions. Internal components must not depend interfaces to narrow internal abstractions. Internal components must not depend
on consumers or on Scriptorium. 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 ## Repository And Consumer Boundary
Scriptorium is a downstream application that consumes Promptkit through Scriptorium is a downstream application that consumes Promptkit through

View File

@@ -50,19 +50,18 @@ Examples of appropriate seams include clocks, randomness, subprocesses, remote A
## Test execution requirements ## Test execution requirements
Promptkit currently uses maintainer-run validation rather than hosted CI. Promptkit currently uses maintainer-run validation rather than hosted CI.
Maintainers run the repository-documented test, vet, build, formatting, Maintainers run the complete local workflow in the
documentation-link, and repository-hygiene checks before accepting changes. [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 Introducing hosted CI later would supplement, not silently redefine, this
documented validation model. documented validation model.
The complete test sequence includes ordinary and race-enabled package tests. Maintainer validation must include ordinary and race-enabled package tests,
The maintained offline consumer workflow is also run from the repository root: static analysis, a complete build, and execution of both maintained offline
consumer examples. The preparation example protects assembled preparation and
```sh inspection behavior. The execution example separately protects assembled
go test ./... `Run`, injected-client, validation, usage, and result behavior.
go test -race ./...
go run ./examples/go-library/prepare
```
Tests in the default suite must be deterministic, offline, and independent of Tests in the default suite must be deterministic, offline, and independent of
real credentials. They must not invoke paid APIs, use live network 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. interaction, while replacing live or nondeterministic external boundaries.
- External-package root tests exercise the public facade as a Go consumer, - External-package root tests exercise the public facade as a Go consumer,
while internal package tests own focused implementation behavior. while internal package tests own focused implementation behavior.
- The maintained offline preparation example protects one representative - The maintained offline preparation and execution examples protect distinct
assembled consumer workflow without contacting a model provider. representative assembled consumer workflows without contacting a model
provider.
- Fixtures should be minimal, synthetic, versioned with the behavior they - Fixtures should be minimal, synthetic, versioned with the behavior they
exercise, and free of credentials or private data. exercise, and free of credentials or private data.
- Golden files are appropriate only when the complete output is intentionally - Golden files are appropriate only when the complete output is intentionally

View File

@@ -106,48 +106,12 @@ gitea.maximumdirect.net/eric/promptkit 1.25.5
promptkit gitea.maximumdirect.net/eric/promptkit promptkit gitea.maximumdirect.net/eric/promptkit
``` ```
Run the complete maintainer validation required by the As a release prerequisite, run the complete
[development guide](development.md): [maintainer validation workflow](development.md#maintainer-validation) against
the clean candidate. Do not substitute a partial command list: the development
```sh guide owns the tests, race checks, analysis, build, both offline examples,
go test ./... formatting, Markdown links, generated-output and credential review, and
go test -race ./... repository hygiene. Record the successful workflow result with the candidate.
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)"
```
## Write The Release Note ## Write The Release Note

View File

@@ -90,7 +90,7 @@ engine, err := promptkit.NewEngine(
``` ```
See the 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 for task-oriented usage. The
[`Backend` and `WithBackend` GoDoc](../../backends.go) owns exact registration, [`Backend` and `WithBackend` GoDoc](../../backends.go) owns exact registration,
validation, copying, defaulting, and uniqueness semantics. The validation, copying, defaulting, and uniqueness semantics. The

74
docs/releases/v0.3.0.md Normal file
View 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
View 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
View 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
View 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.

155
docs/releases/v0.7.0.md Normal file
View File

@@ -0,0 +1,155 @@
# Promptkit v0.7.0
This supplemental changelog and migration guide summarizes the consumer-facing
changes from `v0.6.0` to `v0.7.0`. The annotated `v0.7.0` tag is the
authoritative release record. Exact current contracts belong to the linked
GoDoc and durable documentation.
## Summary
`v0.7.0` expands provider integration and profile composition while making
credential and generation-failure handling more flexible:
- Promptkit now includes the `rakestrawhome` backend and its Gemma profile;
- built-in generation failures expose bounded structured provider details;
- an unavailable optional API-key environment source no longer prevents a
request from reaching an upstream that permits unauthenticated access; and
- profiles can inherit from and selectively refine another profile.
## Compatibility
This release adds public declarations and fields but removes none. Existing
keyed configuration literals and ordinary `errors.Is` handling continue to
work.
Adding `BaseProfileID` to `Profile` and `OpenAICompatibleProfileConfig` changes
their struct shape. Consumers using positional composite literals for either
type must convert them to keyed literals. Existing keyed literals require no
change.
The `rakestrawhome` backend ID is now built in and reserved. A consumer that
previously registered that exact ID with `WithBackend` must remove its manual
registration before upgrading. Other custom backend registrations are
unchanged.
When an optional backend, profile, or request `APIKeyEnv` is unset, empty, or
whitespace-only, the built-in client now omits `Authorization` and sends the
request. Previously this condition could fail before transport. Set
`Profile.APIKeyRequired` when missing credentials must remain a local
preflight error.
Provider non-success responses continue to match `ErrLLMGenerate`. Their
rendered wording is not a compatibility contract; consumers can now use
`errors.As` with `*GenerationError` when structured status information is
needed.
## Upgrade
Update the module dependency with:
```sh
go get gitea.maximumdirect.net/eric/promptkit@v0.7.0
go mod tidy
```
Remove any manual `rakestrawhome` backend registration, convert positional
profile literals to keyed literals, and run the consuming project's ordinary
and race-enabled tests.
## Rakestrawhome Built-In Backend And Profile
Every engine now includes the reserved `rakestrawhome` backend, identified by
`BackendRakestrawHome`. The built-in `rakestrawhome-gemma-4-31b` profile
selects that backend. Consumers can use the maintained endpoint, credential,
capacity, and model defaults without registering either definition themselves.
See the [built-in backend and profile catalogs](../formats.md#built-in-backends)
and the [consumer adoption example](../consumers/pkg-promptkit.md#use-the-rakestrawhome-built-in-profile)
for the current contracts.
## Structured Generation Errors
Non-2xx responses from the built-in OpenAI-compatible client now return an
immutable `*GenerationError`. Consumers can inspect the HTTP status and any
safely extracted provider code, type, or message while retaining the ordinary
generation-error category:
```go
var generationErr *promptkit.GenerationError
if errors.As(err, &generationErr) {
status := generationErr.StatusCode()
_ = status
}
```
Provider fields are bounded and normalized but remain untrusted and may
contain sensitive request or schema details. Default and Go-syntax formatting
omit those fields. Applications must apply their own disclosure policy before
logging or presenting accessor values.
See the [`GenerationError` GoDoc](../../generation_error.go), the
[consumer error-handling guide](../consumers/pkg-promptkit.md#handle-errors),
and the [OpenAI-compatible response contract](../integrations/openai-compatible-chat.md#response-handling).
## Optional Credential Sources
`APIKeyEnv` names an optional environment lookup source unless the selected
profile explicitly sets `APIKeyRequired`. When neither a direct request key nor
a usable environment value exists, the built-in client omits the bearer header
and handles the upstream response normally. This supports local and other
OpenAI-compatible providers that permit unauthenticated requests without
hiding an authentication error returned by a provider that requires one.
The [credential format reference](../formats.md#credentials), the
[`Backend` GoDoc](../../backends.go), the
[`ExecutionTargetOverride` GoDoc](../../types.go), and the
[authentication integration contract](../integrations/openai-compatible-chat.md#authentication)
define the current precedence and availability rules.
## Profile Inheritance
YAML profiles can name one parent with `base_profile`; in-memory profiles use
`Profile.BaseProfileID`, and `OpenAICompatibleProfileConfig` forwards the same
field. A profile can act as an application-owned alias of a built-in or refine
selected inherited settings:
```go
promptkit.WithProfiles(promptkit.Profile{
ID: "weather-light",
BaseProfileID: "deepseek-4-flash",
ReasoningEffort: "high",
})
```
Base lookup observes the existing source precedence. Chains are linear,
cycle-safe, and resolved afresh for ordinary operations. Prepared execution
freezes the fully resolved target. The selected leaf ID remains public while
effective execution settings reflect the resolved chain.
See the [profile inheritance format reference](../formats.md#profile-inheritance),
the [consumer alias example](../consumers/pkg-promptkit.md#alias-a-built-in-profile),
and the [`Profile` GoDoc](../../types.go) for exact merge and validation
behavior.
## Public API Changes
The release adds:
- `BackendRakestrawHome`;
- `GenerationError`, including `StatusCode`, `ProviderCode`, `ProviderType`,
`ProviderMessage`, `Error`, `GoString`, and `Unwrap`;
- `Profile.BaseProfileID`; and
- `OpenAICompatibleProfileConfig.BaseProfileID`.
No public declaration was removed.
## Consumer Action
- Remove a manual backend registration whose ID is exactly `rakestrawhome`.
- Convert positional `Profile` or `OpenAICompatibleProfileConfig` literals to
keyed literals.
- Set `Profile.APIKeyRequired` where a missing credential must fail locally
instead of reaching the provider unauthenticated.
- Treat `GenerationError` provider fields as untrusted and potentially
sensitive when adopting the new accessors.
- Run consumer ordinary and race-enabled tests after updating the module.

View File

@@ -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
View 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.

View File

@@ -12,6 +12,9 @@ consumer value, and important scope boundaries. Defer API design,
implementation details, sequencing, and acceptance criteria until an idea is implementation details, sequencing, and acceptance criteria until an idea is
selected. selected.
Ideas that have been deliberately postponed rather than left available for
ordinary selection belong in the [deferred catalog](deferred.md).
## Using This Catalog ## Using This Catalog
- Add an idea when its purpose and likely value can be stated clearly. - 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 an idea is selected, move its active planning to a focused roadmap or,
when it requires a durable architectural decision, an ADR. Update when it requires a durable architectural decision, an ADR. Update
current-state documentation only when implementation lands. 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 - Remove ideas that are no longer relevant. Retain a rejected idea only when
its rationale is likely to prevent repeated reconsideration. its rationale is likely to prevent repeated reconsideration.
@@ -33,9 +38,28 @@ consumers.
## Ideas ## Ideas
No ideas are currently cataloged. Backend-specific concurrency management has ### Public bounded output repair
been selected for active planning in the
[focused concurrency roadmap](concurrency.md). 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.
## Entry Format ## Entry Format

View File

@@ -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.

375
engine.go
View File

@@ -53,9 +53,9 @@ var (
// an execution profile or resolve its backend, except for the profile // an execution profile or resolve its backend, except for the profile
// not-found case represented by ErrProfileNotFound. // not-found case represented by ErrProfileNotFound.
ErrProfileLoad = errors.New("failed to load execution profile") ErrProfileLoad = errors.New("failed to load execution profile")
// ErrAPIKeyEnvMissing identifies an APIKeyEnv whose environment variable is // ErrAPIKeyEnvMissing identifies an explicitly required APIKeyEnv whose
// unset or empty when no direct RunRequest.APIKey takes precedence. Such an // environment variable is unset or empty after direct RunRequest.APIKey
// error also matches ErrInvalidRequest. // precedence is applied. Such an error also matches ErrInvalidRequest.
ErrAPIKeyEnvMissing = errors.New("api_key_env points to an unset environment variable") ErrAPIKeyEnvMissing = errors.New("api_key_env points to an unset environment variable")
// ErrArtifactLoad identifies a failure to resolve an input artifact. Errors // ErrArtifactLoad identifies a failure to resolve an input artifact. Errors
// returned by an injected ArtifactReader remain available through errors.Is. // returned by an injected ArtifactReader remain available through errors.Is.
@@ -63,13 +63,15 @@ var (
// ErrPromptRender identifies a failure to render prompt messages or the // ErrPromptRender identifies a failure to render prompt messages or the
// session ID from the resolved inputs and variables. // session ID from the resolved inputs and variables.
ErrPromptRender = errors.New("failed to render prompt") ErrPromptRender = errors.New("failed to render prompt")
// ErrCapacityExceeded identifies a Run rejected because the selected backend // ErrCapacityExceeded identifies a Run or RunPrepared rejected because the
// already admitted ConcurrencyLimit + QueueCapacity calls. It is not an // selected backend already admitted ConcurrencyLimit + QueueCapacity calls.
// invalid request, an LLM or provider rate-limit response, or ErrLLMGenerate. // 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") ErrCapacityExceeded = errors.New("backend capacity exceeded")
// ErrLLMGenerate identifies a model-client failure or a nil successful // ErrLLMGenerate identifies a model-client failure or a nil successful
// response. Errors returned by an injected LLMClient remain available // response. A built-in OpenAI-compatible non-2xx response is available as a
// through errors.Is. // [GenerationError]. Errors returned by an injected LLMClient remain
// available through errors.Is.
ErrLLMGenerate = errors.New("failed to generate output") ErrLLMGenerate = errors.New("failed to generate output")
// ErrValidation identifies an operational failure to load or compile a // ErrValidation identifies an operational failure to load or compile a
// schema or validate output. A completed validation whose Status is // schema or validate output. A completed validation whose Status is
@@ -77,12 +79,15 @@ var (
ErrValidation = errors.New("failed to validate output") 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]. // An Engine is safe for concurrent calls to [Engine.InspectPrompt],
// Each Engine owns independent backend-capacity pools that coordinate Run // [Engine.InspectProfile], [Engine.Prepare], [Engine.PrepareExecution],
// admission and model generation. Injected collaborators may still be invoked // [Engine.Run], and [Engine.RunPrepared]. Each Engine owns independent
// concurrently across different backend pools or for unlimited backends. // 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 { type Engine struct {
runner *usecase.Runner runner *usecase.Runner
} }
@@ -94,9 +99,10 @@ type Config struct {
// It is required unless a WithPromptFS or WithPromptFile option supplies the // It is required unless a WithPromptFS or WithPromptFile option supplies the
// prompt source. // prompt source.
PromptDir string PromptDir string
// ProfileDir is an optional directory whose profiles take precedence over // ProfileDir is an optional ordinary configured source whose profiles take
// embedded built-in profiles. An empty value selects only built-ins unless // precedence over application fallback and embedded built-in profiles. An
// profile options are also supplied. // empty value selects the lower-precedence sources unless a profile-source
// option supplies the ordinary source.
ProfileDir string ProfileDir string
// SchemaDir is the root for JSON Schema files. An empty value uses the // SchemaDir is the root for JSON Schema files. An empty value uses the
// current directory. WithSchemaFS or WithSchemaFile replaces this source. // current directory. WithSchemaFS or WithSchemaFile replaces this source.
@@ -115,12 +121,12 @@ type Config struct {
// Option customizes engine construction. // Option customizes engine construction.
// //
// NewEngine applies options in argument order and ignores nil options. Within // NewEngine applies options in argument order and ignores nil options. Within
// each prompt-source, profile-source, in-memory-profile, schema-source, // each prompt-source, ordinary-profile-source, fallback-profile-source,
// model-client, and artifact-reader category, the last non-nil valid option // in-memory-profile, schema-source, model-client, and artifact-reader
// replaces earlier options in that category. WithBackend is the additive // category, the last non-nil valid option replaces earlier options in that
// exception: unique registrations accumulate, and a repeated backend ID is an // category. WithBackend is the additive exception: unique registrations
// error rather than a replacement. An invalid option fails construction even // accumulate, and a repeated backend ID is an error rather than a replacement.
// if a later option would replace it. // An invalid option fails construction even if a later option would replace it.
type Option interface { type Option interface {
apply(*engineOptions) error apply(*engineOptions) error
} }
@@ -132,26 +138,30 @@ func (f optionFunc) apply(options *engineOptions) error {
} }
type engineOptions struct { type engineOptions struct {
llmClient llm.Client llmClient llm.Client
artifactReader artifactadapter.Reader artifactReader artifactadapter.Reader
promptDefs promptdef.Repository promptDefs promptdef.Repository
profiles profile.Repository profiles profile.Repository
memoryProfiles profile.Repository fallbackProfiles profile.Repository
backends []domain.Backend memoryProfiles profile.Repository
validator validate.Validator backends []domain.Backend
promptSource bool validator validate.Validator
profileSource bool promptSource bool
memorySource bool profileSource bool
validatorSource bool fallbackProfileSource bool
artifactSource 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 // A nil client makes NewEngine fail with ErrInvalidConfig. The Engine schedules
// Generate calls according to the selected backend's capacity policy, but the // Generate calls according to the selected backend's capacity policy, but the
// client may still be called concurrently across different backend pools or for // 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 { func WithLLMClient(client LLMClient) Option {
return optionFunc(func(options *engineOptions) error { return optionFunc(func(options *engineOptions) error {
if client == nil { if client == nil {
@@ -210,7 +220,7 @@ func WithPromptFile(path string) Option {
if err != nil { if err != nil {
return err return err
} }
options.promptDefs = promptdef.NewFSRepository(fsys, root) options.promptDefs = promptdef.NewFileRepository(fsys, root, filepath.Dir(path))
options.promptSource = true options.promptSource = true
return nil return nil
}) })
@@ -218,12 +228,12 @@ func WithPromptFile(path string) Option {
// WithProfileFS loads execution profiles from fsys under root. // WithProfileFS loads execution profiles from fsys under root.
// //
// Profiles from this source overlay built-in profiles. Profile YAML must use // Profiles from this ordinary configured source take precedence over
// api_key_env for environment-based credentials; raw API keys are rejected. // application fallback and built-in profiles. Profile YAML must use api_key_env
// fsys must be non-nil and root must be non-empty; otherwise NewEngine fails // for environment-based credentials; raw API keys are rejected. fsys must be
// with ErrInvalidConfig. This option replaces Config.ProfileDir and earlier // non-nil and root must be non-empty; otherwise NewEngine fails with
// file or FS profile-source options, but remains below WithProfiles in // ErrInvalidConfig. This option replaces Config.ProfileDir and earlier file or
// precedence. // FS profile-source options, but remains below WithProfiles in precedence.
func WithProfileFS(fsys fs.FS, root string) Option { func WithProfileFS(fsys fs.FS, root string) Option {
return optionFunc(func(options *engineOptions) error { return optionFunc(func(options *engineOptions) error {
if fsys == nil { if fsys == nil {
@@ -240,11 +250,12 @@ func WithProfileFS(fsys fs.FS, root string) Option {
// WithProfileFile loads execution profiles from the single profile file at path. // 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 // The profile takes precedence over application fallback and built-in profiles.
// environment-based credentials; raw API keys are rejected. path must name an // Profile YAML must use api_key_env for environment-based credentials; raw API
// existing non-directory file when NewEngine applies the option. This option // keys are rejected. path must name an existing non-directory file when
// replaces Config.ProfileDir and earlier file or FS profile-source options, // NewEngine applies the option. This option replaces Config.ProfileDir and
// but remains below WithProfiles in precedence. // earlier file or FS profile-source options, but remains below WithProfiles in
// precedence.
func WithProfileFile(path string) Option { func WithProfileFile(path string) Option {
return optionFunc(func(options *engineOptions) error { return optionFunc(func(options *engineOptions) error {
fsys, root, err := fileSource(path) fsys, root, err := fileSource(path)
@@ -257,13 +268,48 @@ func WithProfileFile(path string) Option {
}) })
} }
// WithProfiles configures in-memory profiles that take precedence over // WithFallbackProfileFS supplies application-owned fallback profile
// configured profile files and built-in profiles. // definitions from fsys under root.
// //
// NewEngine validates and copies every profile. IDs must be unique within one // Profile lookup checks, in order, profiles supplied by WithProfiles; the
// call. An invalid profile, duplicate ID, or unsupported ExtraParams value // ordinary configured source selected by WithProfileFile, WithProfileFS, or
// makes construction fail with ErrInvalidConfig. Repeating WithProfiles // Config.ProfileDir; this fallback source; and Promptkit's embedded built-in
// replaces the complete earlier in-memory set rather than merging it. // 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
// ordinary configured, application fallback, and built-in profiles.
//
// NewEngine locally validates and copies every profile. IDs must be unique
// within one call. An invalid local definition, duplicate ID, or unsupported
// ExtraParams value makes construction fail with ErrInvalidConfig. A derived
// profile's base reference and resolved target completeness are checked when it
// is selected or inspected. Repeating WithProfiles replaces the complete
// earlier in-memory set rather than merging it.
func WithProfiles(profiles ...Profile) Option { func WithProfiles(profiles ...Profile) Option {
return optionFunc(func(options *engineOptions) error { return optionFunc(func(options *engineOptions) error {
repo, err := newMemoryProfileRepository(profiles) repo, err := newMemoryProfileRepository(profiles)
@@ -344,13 +390,7 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
promptDefs = promptdef.NewFilesystemRepository(cfg.PromptDir) promptDefs = promptdef.NewFilesystemRepository(cfg.PromptDir)
} }
profiles := builtin.NewRepositoryWithDirectory(cfg.ProfileDir) profiles := newProfileRepository(cfg.ProfileDir, options)
if options.profileSource {
profiles = builtin.NewRepositoryWithPrimary(options.profiles)
}
if options.memorySource {
profiles = profile.NewOverlayRepository(options.memoryProfiles, profiles)
}
backendRegistry, err := backend.NewRegistry(options.backends) backendRegistry, err := backend.NewRegistry(options.backends)
if err != nil { if err != nil {
@@ -403,26 +443,132 @@ func NewEngine(cfg Config, opts ...Option) (*Engine, error) {
}, nil }, 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 profile.NewResolvingRepository(repository)
}
func fileSource(name string) (fs.FS, string, error) { func fileSource(name string) (fs.FS, string, error) {
cleanName := strings.TrimSpace(name) if strings.TrimSpace(name) == "" {
if cleanName == "" {
return nil, "", ErrInvalidConfig return nil, "", ErrInvalidConfig
} }
dir := filepath.Dir(cleanName) dir := filepath.Dir(name)
base := filepath.Base(cleanName) base := filepath.Base(name)
if base == "." || base == string(filepath.Separator) || strings.TrimSpace(base) == "" { if base == "." || base == string(filepath.Separator) {
return nil, "", ErrInvalidConfig return nil, "", ErrInvalidConfig
} }
info, err := os.Stat(cleanName) info, err := os.Stat(name)
if err != nil { 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() { 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 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 resolves and renders a prompt request without calling an LLM.
// //
// Prepare selects the prompt and profile, resolves any selected backend and // Prepare selects the prompt and profile, resolves any selected backend and
@@ -457,6 +603,38 @@ func (e *Engine) Prepare(ctx context.Context, req RunRequest) (*PreparedRun, err
return fromDomainPreparedRun(prepared), nil 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 // Run prepares a request, invokes the configured LLMClient, and validates the
// generated output. // generated output.
// //
@@ -467,14 +645,17 @@ func (e *Engine) Prepare(ctx context.Context, req RunRequest) (*PreparedRun, err
// single-pass even when OutputContract.RepairAttempts is positive. // single-pass even when OutputContract.RepairAttempts is positive.
// //
// Run can return every error category documented by [Engine.Prepare], plus // Run can return every error category documented by [Engine.Prepare], plus
// ErrCapacityExceeded and ErrLLMGenerate. ErrCapacityExceeded identifies // ErrCapacityExceeded and ErrLLMGenerate. An engine admission rejection is
// rejection before artifacts, schemas, rendering, or model generation because // discoverable as [CapacityError] and still matches ErrCapacityExceeded. It
// the selected backend's admission capacity is full; it does not match // occurs before artifacts, schemas, rendering, or model generation because the
// ErrInvalidRequest or ErrLLMGenerate. Errors from injected clients remain // selected backend's admission capacity is full; it does not match
// available through errors.Is. Cancellation while waiting for model-generation // ErrInvalidRequest or ErrLLMGenerate. A built-in OpenAI-compatible non-2xx
// capacity matches both ErrLLMGenerate and the context error. Cancellation // response is discoverable as [GenerationError]. Errors from injected clients
// otherwise follows the active collaborator's documented behavior. A nil // remain available through errors.Is. Cancellation while waiting for
// Engine returns ErrInvalidConfig. Run returns no partial result on error. // model-generation capacity matches both ErrLLMGenerate and the context error.
// Cancellation otherwise follows the active collaborator's documented
// behavior. A nil Engine returns ErrInvalidConfig. Run returns no partial
// result on error.
func (e *Engine) Run(ctx context.Context, req RunRequest) (*RunResult, error) { func (e *Engine) Run(ctx context.Context, req RunRequest) (*RunResult, error) {
if e == nil || e.runner == nil { if e == nil || e.runner == nil {
return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig) return nil, fmt.Errorf("%w: engine is nil", ErrInvalidConfig)
@@ -491,3 +672,43 @@ func (e *Engine) Run(ctx context.Context, req RunRequest) (*RunResult, error) {
} }
return fromDomainRunResult(result), nil 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 built-in OpenAI-compatible non-2xx response is
// discoverable as [GenerationError]. 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
}

File diff suppressed because it is too large Load Diff

View File

@@ -3,8 +3,10 @@ package promptkit
import ( import (
"errors" "errors"
"fmt" "fmt"
"strings"
"gitea.maximumdirect.net/eric/promptkit/internal/capacity" "gitea.maximumdirect.net/eric/promptkit/internal/capacity"
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
"gitea.maximumdirect.net/eric/promptkit/internal/profile" "gitea.maximumdirect.net/eric/promptkit/internal/profile"
"gitea.maximumdirect.net/eric/promptkit/internal/promptdef" "gitea.maximumdirect.net/eric/promptkit/internal/promptdef"
"gitea.maximumdirect.net/eric/promptkit/internal/usecase" "gitea.maximumdirect.net/eric/promptkit/internal/usecase"
@@ -14,7 +16,25 @@ func mapPublicError(err error) error {
if err == nil { if err == nil {
return 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) publicErr := publicErrorFor(err)
var providerHTTPError *llm.ProviderHTTPError
if errors.As(err, &providerHTTPError) && providerHTTPError != nil {
generationErr := newGenerationError(
providerHTTPError.StatusCode(),
providerHTTPError.ProviderCode(),
providerHTTPError.ProviderType(),
providerHTTPError.ProviderMessage(),
)
if publicErr != nil && !errors.Is(publicErr, ErrLLMGenerate) {
return fmt.Errorf("%w: %w", publicErr, generationErr)
}
return generationErr
}
if publicErr == nil { if publicErr == nil {
return err return err
} }

View File

@@ -6,6 +6,7 @@ import (
"fmt" "fmt"
"testing" "testing"
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
"gitea.maximumdirect.net/eric/promptkit/internal/usecase" "gitea.maximumdirect.net/eric/promptkit/internal/usecase"
) )
@@ -20,3 +21,55 @@ func TestMapPublicErrorPreservesGenerationCancellation(t *testing.T) {
t.Fatalf("mapped error=%v, want context.Canceled", err) 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)
}
}
func TestMapPublicErrorPreservesValidationAroundGenerationError(t *testing.T) {
internalErr := fmt.Errorf(
"%w: %w",
usecase.ErrValidation,
&llm.ProviderHTTPError{},
)
err := mapPublicError(internalErr)
if !errors.Is(err, ErrValidation) {
t.Fatalf("mapped error=%v, want ErrValidation", err)
}
if !errors.Is(err, ErrLLMGenerate) {
t.Fatalf("mapped error=%v, want ErrLLMGenerate", err)
}
var generationErr *GenerationError
if !errors.As(err, &generationErr) || generationErr == nil {
t.Fatalf("mapped error=%v, want GenerationError", err)
}
var leakedInternalErr *llm.ProviderHTTPError
if errors.As(err, &leakedInternalErr) {
t.Fatalf("mapped error exposes internal ProviderHTTPError: %v", err)
}
}

87
generation_error.go Normal file
View File

@@ -0,0 +1,87 @@
package promptkit
import "fmt"
// GenerationError reports a non-2xx response from Promptkit's built-in
// OpenAI-compatible client during [Engine.Run] or [Engine.RunPrepared].
//
// Engine-produced values are immutable, caller-owned values. Use errors.Is to
// match [ErrLLMGenerate] and errors.As with a *GenerationError target to obtain
// this type. The four provider accessors expose untrusted provider-controlled
// values that can contain sensitive request or schema fragments. Applications
// must apply their own disclosure policy before logging, displaying, or
// returning them to another caller.
//
// Accessors, Error, GoString, and Unwrap are safe on a nil receiver and a zero
// value. Default and Go-syntax formatting deliberately redact provider details.
// GenerationError has no stable JSON representation.
type GenerationError struct {
statusCode int
providerCode string
providerType string
providerMessage string
}
func newGenerationError(statusCode int, providerCode, providerType, providerMessage string) *GenerationError {
return &GenerationError{
statusCode: statusCode,
providerCode: providerCode,
providerType: providerType,
providerMessage: providerMessage,
}
}
// StatusCode returns the received provider HTTP status code, or zero for a nil
// receiver or zero value.
func (e *GenerationError) StatusCode() int {
if e == nil {
return 0
}
return e.statusCode
}
// ProviderCode returns the normalized provider error code, if present. Its
// value is untrusted and may contain sensitive data.
func (e *GenerationError) ProviderCode() string {
if e == nil {
return ""
}
return e.providerCode
}
// ProviderType returns the normalized provider error type, if present. Its
// value is untrusted and may contain sensitive data.
func (e *GenerationError) ProviderType() string {
if e == nil {
return ""
}
return e.providerType
}
// ProviderMessage returns the bounded normalized provider diagnostic, if
// present. Its value is untrusted and may contain sensitive data.
func (e *GenerationError) ProviderMessage() string {
if e == nil {
return ""
}
return e.providerMessage
}
// Error returns a redacted diagnostic that is not a parsing contract.
func (e *GenerationError) Error() string {
if e == nil || e.statusCode == 0 {
return ErrLLMGenerate.Error()
}
return fmt.Sprintf("%s: provider returned HTTP status %d", ErrLLMGenerate, e.statusCode)
}
// GoString returns the same redacted diagnostic as Error.
func (e *GenerationError) GoString() string {
return e.Error()
}
// Unwrap returns ErrLLMGenerate. It is safe to call on a nil receiver or zero
// value.
func (e *GenerationError) Unwrap() error {
return ErrLLMGenerate
}

View File

@@ -0,0 +1,95 @@
package promptkit_test
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"strings"
"testing"
"gitea.maximumdirect.net/eric/promptkit"
)
func TestBuiltInGenerationError(t *testing.T) {
const (
codeMarker = "provider-code-marker"
typeMarker = "provider-type-marker"
messageMarker = "provider-message-marker"
)
engine := newBuiltInGenerationErrorEngine(t, http.StatusUnprocessableEntity,
`{"error":{"code":"`+codeMarker+`","type":"`+typeMarker+`","message":"`+messageMarker+`"}}`)
result, err := engine.Run(context.Background(), generationErrorRunRequest())
if result != nil {
t.Fatalf("Run result = %#v, want nil", result)
}
assertGenerationError(t, err, http.StatusUnprocessableEntity, codeMarker, typeMarker, messageMarker)
preparedEngine := newBuiltInGenerationErrorEngine(t, http.StatusServiceUnavailable, `{"error":{}}`)
prepared, err := preparedEngine.PrepareExecution(context.Background(), generationErrorRunRequest())
if err != nil {
t.Fatalf("PrepareExecution: %v", err)
}
result, err = preparedEngine.RunPrepared(context.Background(), prepared)
if result != nil {
t.Fatalf("RunPrepared result = %#v, want nil", result)
}
assertGenerationError(t, err, http.StatusServiceUnavailable, "", "", "")
}
func assertGenerationError(t *testing.T, err error, statusCode int, code, providerType, message string) {
t.Helper()
if !errors.Is(err, promptkit.ErrLLMGenerate) {
t.Fatalf("errors.Is(%v, ErrLLMGenerate) = false", err)
}
var generationErr *promptkit.GenerationError
if !errors.As(err, &generationErr) || generationErr == nil {
t.Fatalf("error = %T, want *GenerationError", err)
}
if generationErr.StatusCode() != statusCode || generationErr.ProviderCode() != code || generationErr.ProviderType() != providerType || generationErr.ProviderMessage() != message {
t.Fatalf("GenerationError = %#v", generationErr)
}
wantFormatted := fmt.Sprintf("failed to generate output: provider returned HTTP status %d", statusCode)
for _, rendered := range []string{fmt.Sprintf("%v", generationErr), fmt.Sprintf("%+v", generationErr), fmt.Sprintf("%#v", generationErr)} {
if rendered != wantFormatted {
t.Fatalf("formatted error = %q, want %q", rendered, wantFormatted)
}
for _, marker := range []string{code, providerType, message} {
if marker != "" && strings.Contains(rendered, marker) {
t.Fatalf("formatted error exposed provider marker %q: %q", marker, rendered)
}
}
}
}
func newBuiltInGenerationErrorEngine(t *testing.T, statusCode int, body string) *promptkit.Engine {
t.Helper()
config := contractConfig(frameworkSchemaDir)
config.HTTPClient = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: statusCode,
ContentLength: int64(len(body)),
Body: io.NopCloser(strings.NewReader(body)),
}, nil
})}
engine, err := promptkit.NewEngine(config)
if err != nil {
t.Fatalf("NewEngine: %v", err)
}
return engine
}
func generationErrorRunRequest() promptkit.RunRequest {
return promptkit.RunRequest{
PromptID: frameworkMarkdownSummaryPromptID,
Inputs: map[string]promptkit.ArtifactRef{
"transcript": promptkit.Inline("Rin opens the gate."),
"glossary": promptkit.Inline("gate: A guarded passage."),
},
}
}

View File

@@ -0,0 +1,32 @@
package promptkit
import (
"errors"
"fmt"
"testing"
)
func TestGenerationErrorNilAndZeroValue(t *testing.T) {
var nilError *GenerationError
zeroError := &GenerationError{}
for name, err := range map[string]*GenerationError{
"nil": nilError,
"zero": zeroError,
} {
t.Run(name, func(t *testing.T) {
if err.StatusCode() != 0 || err.ProviderCode() != "" || err.ProviderType() != "" || err.ProviderMessage() != "" {
t.Fatalf("accessors returned provider details: %#v", err)
}
if err.Error() != "failed to generate output" || err.GoString() != "failed to generate output" {
t.Fatalf("redacted formatting = (%q, %q)", err.Error(), err.GoString())
}
if fmt.Sprintf("%v", err) != "failed to generate output" || fmt.Sprintf("%#v", err) != "failed to generate output" {
t.Fatalf("formatted error = (%q, %q)", fmt.Sprintf("%v", err), fmt.Sprintf("%#v", err))
}
if !errors.Is(err, ErrLLMGenerate) {
t.Fatalf("errors.Is(%v, ErrLLMGenerate) = false", err)
}
})
}
}

View File

@@ -16,10 +16,12 @@ import (
var ( var (
ErrUnsupportedRefType = errors.New("unsupported artifact reference type") ErrUnsupportedRefType = errors.New("unsupported artifact reference type")
ErrMissingInlineBody = errors.New("missing body for inline artifact")
ErrMissingFilePath = errors.New("missing file path for file 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. // Reader resolves artifact references into actual artifacts.
type Reader interface { type Reader interface {
Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error)
@@ -34,7 +36,7 @@ type CompositeReader struct {
func NewCompositeReader() Reader { func NewCompositeReader() Reader {
return &CompositeReader{ return &CompositeReader{
inlineReader: &inlineReader{}, inlineReader: &inlineReader{},
fileReader: &fileReader{}, fileReader: &fileReader{open: openArtifactFile},
} }
} }
@@ -64,10 +66,6 @@ func (r *inlineReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domai
default: default:
} }
if ref.Body == "" {
return nil, ErrMissingInlineBody
}
body := []byte(ref.Body) body := []byte(ref.Body)
return &domain.Artifact{ return &domain.Artifact{
ContentType: defaults.ContentTypeTextPlain, ContentType: defaults.ContentTypeTextPlain,
@@ -78,7 +76,15 @@ func (r *inlineReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domai
}, nil }, 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) { func (r *fileReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error) {
select { select {
@@ -91,25 +97,71 @@ func (r *fileReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.
return nil, ErrMissingFilePath return nil, ErrMissingFilePath
} }
return readFileArtifact(ref.URI) return readFileArtifact(ctx, ref.URI, r.open)
} }
func readFileArtifact(path string) (*domain.Artifact, error) { func openArtifactFile(path string) (artifactFile, error) {
file, err := os.Open(path) 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 { if err != nil {
return nil, fmt.Errorf("failed to read file %s: %w", path, err) return nil, fmt.Errorf("failed to read file %s: %w", path, err)
} }
defer file.Close() defer file.Close()
data, err := io.ReadAll(file) openedInfo, err := file.Stat()
if err != nil { 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)) contentType := mime.TypeByExtension(filepath.Ext(path))
if contentType == "" { if contentType == "" {
contentType = defaults.ContentTypeTextPlain contentType = defaults.ContentTypeTextPlain
} }
hash := fmt.Sprintf("%x", sha256.Sum256(data))
if err := ctx.Err(); err != nil {
return nil, err
}
return &domain.Artifact{ return &domain.Artifact{
Name: filepath.Base(path), Name: filepath.Base(path),
@@ -117,6 +169,6 @@ func readFileArtifact(path string) (*domain.Artifact, error) {
Body: data, Body: data,
URI: path, URI: path,
Size: int64(len(data)), Size: int64(len(data)),
Hash: fmt.Sprintf("%x", sha256.Sum256(data)), Hash: hash,
}, nil }, nil
} }

View 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")
}
}

View File

@@ -1,6 +1,7 @@
package artifact package artifact
import ( import (
"bytes"
"context" "context"
"errors" "errors"
"os" "os"
@@ -11,54 +12,97 @@ import (
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "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() 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) { for _, tc := range tests {
ref := domain.ArtifactRef{ t.Run(tc.name, func(t *testing.T) {
Type: domain.ArtifactRefInline, filePath := filepath.Join(t.TempDir(), "artifact.txt")
Body: "hello world", if err := os.WriteFile(filePath, []byte(tc.content), 0o600); err != nil {
} t.Fatal(err)
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)
}
})
t.Run("inline artifact missing body", func(t *testing.T) { sources := []struct {
ref := domain.ArtifactRef{ name string
Type: domain.ArtifactRefInline, ref domain.ArtifactRef
Body: "", wantURI string
} }{
_, err := reader.Read(ctx, ref) {
if !errors.Is(err, ErrMissingInlineBody) { name: "inline",
t.Errorf("expected ErrMissingInlineBody, got %v", err) 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) { var sourceHash string
ref := domain.ArtifactRef{ for _, source := range sources {
Type: domain.ArtifactRefType("unsupported"), t.Run(source.name, func(t *testing.T) {
URI: "unsupported://bucket/key", first, err := reader.Read(context.Background(), source.ref)
} if err != nil {
_, err := reader.Read(ctx, ref) t.Fatalf("first read: %v", err)
if !errors.Is(err, ErrUnsupportedRefType) { }
t.Error("expected error for unsupported type") 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) { func TestCompositeReaderCopiesInlineData(t *testing.T) {
@@ -87,94 +131,154 @@ func TestCompositeReaderCopiesInlineData(t *testing.T) {
} }
} }
func TestCompositeReaderHonorsCancellation(t *testing.T) { func TestCompositeReaderHonorsPreCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background()) filePath := filepath.Join(t.TempDir(), "artifact.txt")
cancel() 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{ for _, tc := range tests {
Type: domain.ArtifactRefInline, t.Run(tc.name, func(t *testing.T) {
Body: "ignored", ctx, cancel := context.WithCancel(context.Background())
}) cancel()
if !errors.Is(err, context.Canceled) {
t.Fatalf("expected context cancellation, got %v", err) 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) { func TestFileReaderFailuresAndMetadata(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)
}
reader := NewCompositeReader() 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) { t.Run("missing file path", func(t *testing.T) {
ref := domain.ArtifactRef{ _, err := reader.Read(context.Background(), domain.ArtifactRef{Type: domain.ArtifactRefFile})
Type: domain.ArtifactRefFile,
URI: "",
}
_, err := reader.Read(ctx, ref)
if !errors.Is(err, ErrMissingFilePath) { 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) { t.Run("missing file", func(t *testing.T) {
ref := domain.ArtifactRef{ _, err := reader.Read(context.Background(), domain.ArtifactRef{
Type: domain.ArtifactRefFile, Type: domain.ArtifactRefFile,
URI: filepath.Join(t.TempDir(), "missing.txt"), 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.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) { t.Run("unknown extension uses text fallback", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "artifact.unknownextension") filePath := filepath.Join(t.TempDir(), "artifact.unknownextension")
if err := os.WriteFile(path, content, 0o600); err != nil { if err := os.WriteFile(filePath, []byte("content"), 0o600); err != nil {
t.Fatal(err) t.Fatal(err)
} }
art, err := reader.Read(ctx, domain.ArtifactRef{ artifact, err := reader.Read(context.Background(), domain.ArtifactRef{
Type: domain.ArtifactRefFile, Type: domain.ArtifactRefFile,
URI: path, URI: filePath,
}) })
if err != nil { if err != nil {
t.Fatalf("unexpected error: %v", err) t.Fatalf("read artifact: %v", err)
} }
if art.ContentType != "text/plain" { if artifact.ContentType != "text/plain" {
t.Errorf("expected text/plain fallback, got %q", art.ContentType) 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
}

View File

@@ -5,7 +5,6 @@ package backend
import ( import (
"errors" "errors"
"fmt" "fmt"
"net/url"
"regexp" "regexp"
"sort" "sort"
"strings" "strings"
@@ -19,12 +18,20 @@ const (
// OpenRouterID is the reserved ID of Promptkit's built-in OpenRouter // OpenRouterID is the reserved ID of Promptkit's built-in OpenRouter
// backend. // backend.
OpenRouterID = "openrouter" OpenRouterID = "openrouter"
// RakestrawHomeID is the reserved ID of Promptkit's built-in Rakestrawhome
// backend.
RakestrawHomeID = "rakestrawhome"
openRouterEndpoint = "https://openrouter.ai/api/v1" openRouterEndpoint = "https://openrouter.ai/api/v1"
openRouterAPIKeyEnv = "OPENROUTER_API_KEY" openRouterAPIKeyEnv = "OPENROUTER_API_KEY"
openRouterConcurrencyLimit = 16 openRouterConcurrencyLimit = 16
defaultQueueCapacity = 1024
rakestrawHomeEndpoint = "https://inference.ai.rakestrawhome.com/v1"
rakestrawHomeAPIKeyEnv = "RAKESTRAWHOME_INFERENCE_API_KEY"
rakestrawHomeConcurrencyLimit = 4
defaultQueueCapacity = 1024
) )
// ErrBackendNotFound identifies a registry lookup for an unknown backend ID. // ErrBackendNotFound identifies a registry lookup for an unknown backend ID.
@@ -37,20 +44,16 @@ type Registry struct {
backends map[string]domain.Backend backends map[string]domain.Backend
} }
// NewRegistry constructs a registry containing the built-in OpenRouter // NewRegistry constructs a registry containing the built-in definitions
// definition followed by the supplied additions. Every ID must be unique. // followed by the supplied additions. Every ID must be unique.
func NewRegistry(additions []domain.Backend) (*Registry, error) { func NewRegistry(additions []domain.Backend) (*Registry, error) {
builtIns := builtInBackends()
registry := &Registry{ registry := &Registry{
backends: make(map[string]domain.Backend, len(additions)+1), backends: make(map[string]domain.Backend, len(builtIns)+len(additions)),
} }
definitions := make([]domain.Backend, 0, len(additions)+1) definitions := make([]domain.Backend, 0, len(builtIns)+len(additions))
definitions = append(definitions, domain.Backend{ definitions = append(definitions, builtIns...)
ID: OpenRouterID,
Endpoint: openRouterEndpoint,
APIKeyEnv: openRouterAPIKeyEnv,
ConcurrencyLimit: openRouterConcurrencyLimit,
})
definitions = append(definitions, additions...) definitions = append(definitions, additions...)
for _, definition := range definitions { for _, definition := range definitions {
@@ -72,6 +75,23 @@ func NewRegistry(additions []domain.Backend) (*Registry, error) {
return registry, nil return registry, nil
} }
func builtInBackends() []domain.Backend {
return []domain.Backend{
{
ID: OpenRouterID,
Endpoint: openRouterEndpoint,
APIKeyEnv: openRouterAPIKeyEnv,
ConcurrencyLimit: openRouterConcurrencyLimit,
},
{
ID: RakestrawHomeID,
Endpoint: rakestrawHomeEndpoint,
APIKeyEnv: rakestrawHomeAPIKeyEnv,
ConcurrencyLimit: rakestrawHomeConcurrencyLimit,
},
}
}
// GetBackend returns a defensive copy of the backend registered with id. // GetBackend returns a defensive copy of the backend registered with id.
func (r *Registry) GetBackend(id string) (domain.Backend, error) { func (r *Registry) GetBackend(id string) (domain.Backend, error) {
if r == nil { if r == nil {
@@ -109,10 +129,11 @@ func (r *Registry) CapacityPolicies() map[string]domain.BackendCapacityPolicy {
} }
func normalizeBackend(definition domain.Backend) (domain.Backend, error) { func normalizeBackend(definition domain.Backend) (domain.Backend, error) {
definition.Endpoint = strings.TrimSpace(definition.Endpoint) endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(definition.Endpoint)
if err := validateEndpoint(definition.Endpoint); err != nil { if err != nil {
return domain.Backend{}, fmt.Errorf("backend %q endpoint: %w", definition.ID, err) return domain.Backend{}, fmt.Errorf("backend %q endpoint: %w", definition.ID, err)
} }
definition.Endpoint = endpoint
definition.APIKeyEnv = strings.TrimSpace(definition.APIKeyEnv) definition.APIKeyEnv = strings.TrimSpace(definition.APIKeyEnv)
if definition.APIKeyEnv != "" && !environmentVariableName.MatchString(definition.APIKeyEnv) { if definition.APIKeyEnv != "" && !environmentVariableName.MatchString(definition.APIKeyEnv) {
@@ -182,31 +203,3 @@ func normalizeBackend(definition domain.Backend) (domain.Backend, error) {
definition.ExtraParams = extraParams definition.ExtraParams = extraParams
return definition, nil 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
}

View File

@@ -11,32 +11,63 @@ import (
const validEndpoint = "https://backend.example/v1" const validEndpoint = "https://backend.example/v1"
func TestRegistryIncludesExactOpenRouterDefinition(t *testing.T) { func TestRegistryIncludesExactBuiltInDefinitions(t *testing.T) {
registry, err := backend.NewRegistry(nil) registry, err := backend.NewRegistry(nil)
if err != nil { if err != nil {
t.Fatalf("construct registry: %v", err) t.Fatalf("construct registry: %v", err)
} }
definition, err := registry.GetBackend(backend.OpenRouterID) tests := []struct {
if err != nil { name string
t.Fatalf("look up OpenRouter: %v", err) id string
endpoint string
apiKeyEnv string
concurrent int
}{
{
name: "OpenRouter",
id: backend.OpenRouterID,
endpoint: "https://openrouter.ai/api/v1",
apiKeyEnv: "OPENROUTER_API_KEY",
concurrent: 16,
},
{
name: "Rakestrawhome",
id: backend.RakestrawHomeID,
endpoint: "https://inference.ai.rakestrawhome.com/v1",
apiKeyEnv: "RAKESTRAWHOME_INFERENCE_API_KEY",
concurrent: 4,
},
} }
if definition.ID != "openrouter" || for _, tc := range tests {
definition.Endpoint != "https://openrouter.ai/api/v1" || t.Run(tc.name, func(t *testing.T) {
definition.APIKeyEnv != "OPENROUTER_API_KEY" || definition, err := registry.GetBackend(tc.id)
definition.ConcurrencyLimit != 16 || if err != nil {
definition.QueueCapacity != 1024 || t.Fatalf("look up built-in: %v", err)
!definition.QueueCapacitySet || }
definition.ExtraParams != nil { if definition.ID != tc.id ||
t.Fatalf("unexpected OpenRouter definition: %#v", definition) definition.Endpoint != tc.endpoint ||
definition.APIKeyEnv != tc.apiKeyEnv ||
definition.ConcurrencyLimit != tc.concurrent ||
definition.QueueCapacity != 1024 ||
!definition.QueueCapacitySet ||
definition.ExtraParams != nil {
t.Fatalf("unexpected built-in definition: %#v", definition)
}
})
} }
policies := registry.CapacityPolicies() policies := registry.CapacityPolicies()
if len(policies) != 1 || if len(policies) != 2 ||
policies["openrouter"] != (domain.BackendCapacityPolicy{ policies[backend.OpenRouterID] != (domain.BackendCapacityPolicy{
ConcurrencyLimit: 16, ConcurrencyLimit: 16,
QueueCapacity: 1024, QueueCapacity: 1024,
}) ||
policies[backend.RakestrawHomeID] != (domain.BackendCapacityPolicy{
ConcurrencyLimit: 4,
QueueCapacity: 1024,
}) { }) {
t.Fatalf("unexpected OpenRouter capacity policies: %#v", policies) t.Fatalf("unexpected built-in capacity policies: %#v", policies)
} }
} }
@@ -109,11 +140,12 @@ func TestRegistryNormalizesUniqueAdditionsAndIsolatesMutations(t *testing.T) {
} }
policies := registry.CapacityPolicies() policies := registry.CapacityPolicies()
if len(policies) != 2 { if len(policies) != 3 {
t.Fatalf("unexpected capacity policy count: %#v", policies) t.Fatalf("unexpected capacity policy count: %#v", policies)
} }
policies["custom"] = domain.BackendCapacityPolicy{} policies["custom"] = domain.BackendCapacityPolicy{}
delete(policies, backend.OpenRouterID) delete(policies, backend.OpenRouterID)
delete(policies, backend.RakestrawHomeID)
againPolicies := registry.CapacityPolicies() againPolicies := registry.CapacityPolicies()
if againPolicies["custom"] != (domain.BackendCapacityPolicy{ if againPolicies["custom"] != (domain.BackendCapacityPolicy{
ConcurrencyLimit: 3, ConcurrencyLimit: 3,
@@ -122,7 +154,10 @@ func TestRegistryNormalizesUniqueAdditionsAndIsolatesMutations(t *testing.T) {
t.Fatalf("capacity policy map mutated registry state: %#v", againPolicies) t.Fatalf("capacity policy map mutated registry state: %#v", againPolicies)
} }
if _, ok := againPolicies[backend.OpenRouterID]; !ok { if _, ok := againPolicies[backend.OpenRouterID]; !ok {
t.Fatalf("capacity policy deletion mutated registry state: %#v", againPolicies) t.Fatalf("OpenRouter capacity policy deletion mutated registry state: %#v", againPolicies)
}
if _, ok := againPolicies[backend.RakestrawHomeID]; !ok {
t.Fatalf("Rakestrawhome capacity policy deletion mutated registry state: %#v", againPolicies)
} }
} }
@@ -242,11 +277,14 @@ func TestNewRegistryRejectsDuplicateIDs(t *testing.T) {
wantID string wantID string
}{ }{
{ {
name: "built-in collision after normalization", name: "OpenRouter collision after normalization",
additions: []domain.Backend{{ additions: []domain.Backend{{ID: " openrouter "}},
ID: " openrouter ", wantID: backend.OpenRouterID,
}}, },
wantID: "openrouter", {
name: "Rakestrawhome collision after normalization",
additions: []domain.Backend{{ID: " rakestrawhome "}},
wantID: backend.RakestrawHomeID,
}, },
{ {
name: "consumer collision after normalization", name: "consumer collision after normalization",

View File

@@ -124,6 +124,11 @@ func TestManagerAdmissionHonorsContextAndUnlimitedBackends(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("construct manager: %v", err) 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()) ctx, cancel := context.WithCancel(context.Background())
cancel() cancel()

View File

@@ -14,21 +14,12 @@ const (
ContentTypeApplicationJSON = "application/json" ContentTypeApplicationJSON = "application/json"
OpenAIChatCompletionsPath = "/chat/completions" OpenAIChatCompletionsPath = "/chat/completions"
ExecutionDefaultTemperature = 0.0
ExecutionDefaultMaxTokens = 0
ExecutionDefaultTopP = 1.0
ExecutionDefaultTimeoutSeconds = 600 ExecutionDefaultTimeoutSeconds = 600
) LLMRequestTimeoutDefault = 10 * time.Minute
var (
LLMRequestTimeoutDefault = 10 * time.Minute
) )
func ExecutionTargetDefault() domain.ExecutionTarget { func ExecutionTargetDefault() domain.ExecutionTarget {
return domain.ExecutionTarget{ return domain.ExecutionTarget{
Temperature: ExecutionDefaultTemperature,
MaxTokens: ExecutionDefaultMaxTokens,
TopP: ExecutionDefaultTopP,
TimeoutSeconds: ExecutionDefaultTimeoutSeconds, TimeoutSeconds: ExecutionDefaultTimeoutSeconds,
} }
} }

View File

@@ -97,22 +97,22 @@ type RunResult struct {
// PreparedRun contains pre-LLM execution state from the prepare/render phase. // PreparedRun contains pre-LLM execution state from the prepare/render phase.
// It must never include resolved API key values, model output, or validation data. // It must never include resolved API key values, model output, or validation data.
type PreparedRun struct { type PreparedRun struct {
PromptID string `json:"prompt_id"` PromptID string
PromptVersion string `json:"prompt_version,omitempty"` PromptVersion string
PromptHash string `json:"prompt_hash,omitempty"` PromptHash string
SelectedProfileID string `json:"selected_profile_id"` SelectedProfileID string
SelectedBackendID string `json:"selected_backend_id,omitempty"` SelectedBackendID string
EffectiveModelParams ExecutionTarget `json:"effective_model_params"` EffectiveModelParams ExecutionTarget
TargetPresence ExecutionTargetPresence `json:"-"` TargetPresence ExecutionTargetPresence
OutputContract OutputContract `json:"output_contract"` OutputContract OutputContract
StructuredOutput *StructuredOutputSpec `json:"structured_output,omitempty"` StructuredOutput *StructuredOutputSpec
InputHashes map[string]string `json:"input_hashes,omitempty"` InputHashes map[string]string
SessionID string `json:"session_id,omitempty"` SessionID string
RenderedPromptHash string `json:"rendered_prompt_hash"` RenderedPromptHash string
Messages []RenderedMessage `json:"messages"` Messages []RenderedMessage
StartTime time.Time `json:"start_time,omitempty"` StartTime time.Time
EndTime time.Time `json:"end_time,omitempty"` EndTime time.Time
DurationMS int64 `json:"duration_ms,omitempty"` DurationMS int64
} }
// ArtifactRef represents a reference to an input artifact. // ArtifactRef represents a reference to an input artifact.
@@ -145,6 +145,16 @@ type PromptDefinition struct {
Validation OutputContract `yaml:"validation"` 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. // PromptInput describes one named input expected by a prompt definition.
type PromptInput struct { type PromptInput struct {
Name string `yaml:"name"` Name string `yaml:"name"`
@@ -182,6 +192,7 @@ type BackendCapacityPolicy struct {
// ExecutionProfile describes how and where to execute a model. // ExecutionProfile describes how and where to execute a model.
type ExecutionProfile struct { type ExecutionProfile struct {
ID string `yaml:"id"` ID string `yaml:"id"`
BaseProfileID string `yaml:"base_profile"`
BackendID string `yaml:"backend"` BackendID string `yaml:"backend"`
Endpoint string `yaml:"endpoint"` Endpoint string `yaml:"endpoint"`
Model string `yaml:"model"` Model string `yaml:"model"`
@@ -236,6 +247,13 @@ type ExecutionTarget struct {
ExtraParams map[string]any `yaml:"extra_params" json:"extra_params"` 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. // OutputContract defines the requirements for the output artifact.
type OutputContract struct { type OutputContract struct {
Format OutputFormat `yaml:"format"` Format OutputFormat `yaml:"format"`

View 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
}

View 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)
}
})
}
}

View 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)
}

View 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)
}
})
}
}

View 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
}

View 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)
}
})
}
}

View File

@@ -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)
}
}

View File

@@ -8,6 +8,9 @@ import (
// NormalizeSessionID applies the shared session identifier rule. // NormalizeSessionID applies the shared session identifier rule.
func NormalizeSessionID(raw string) (string, error) { func NormalizeSessionID(raw string) (string, error) {
if !utf8.ValidString(raw) {
return "", fmt.Errorf("session_id must contain valid UTF-8")
}
normalized := strings.TrimSpace(raw) normalized := strings.TrimSpace(raw)
if normalized == "" { if normalized == "" {
return "", nil return "", nil

View File

@@ -7,10 +7,10 @@ import (
func TestNormalizeSessionID(t *testing.T) { func TestNormalizeSessionID(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
raw string raw string
want string want string
wantErr bool wantErrContains string
}{ }{
{ {
name: "trims surrounding Unicode whitespace", name: "trims surrounding Unicode whitespace",
@@ -28,21 +28,24 @@ func TestNormalizeSessionID(t *testing.T) {
want: strings.Repeat("界", SessionIDMaxLength), want: strings.Repeat("界", SessionIDMaxLength),
}, },
{ {
name: "one Unicode code point over maximum is rejected", name: "one Unicode code point over maximum is rejected",
raw: strings.Repeat("界", SessionIDMaxLength+1), raw: strings.Repeat("界", SessionIDMaxLength+1),
wantErr: true, 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 { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
got, err := NormalizeSessionID(tt.raw) got, err := NormalizeSessionID(tt.raw)
if tt.wantErr { if tt.wantErrContains != "" {
if err == nil { if err == nil {
t.Fatal("expected normalization error") t.Fatal("expected normalization error")
} }
if !strings.Contains(err.Error(), "exceeds maximum") { if !strings.Contains(err.Error(), tt.wantErrContains) {
t.Fatalf("expected useful length diagnostic, got %v", err) t.Fatalf("expected diagnostic containing %q, got %v", tt.wantErrContains, err)
} }
return return
} }

View File

@@ -36,7 +36,8 @@ func FindYAMLFiles(ctx context.Context, root string) ([]string, error) {
return files, err 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) { func FindFSYAMLFiles(ctx context.Context, fsys fs.FS, root string) ([]string, error) {
cleanRoot := CleanFSRoot(root) cleanRoot := CleanFSRoot(root)
var files []string var files []string
@@ -52,6 +53,10 @@ func FindFSYAMLFiles(ctx context.Context, fsys fs.FS, root string) ([]string, er
if d.IsDir() { if d.IsDir() {
return nil return nil
} }
if name == cleanRoot {
files = append(files, name)
return nil
}
if !IsYAMLFile(d.Name()) { if !IsYAMLFile(d.Name()) {
return nil return nil
} }
@@ -71,10 +76,10 @@ func RelativePath(root string, filePath string) string {
return filepath.Clean(rel) 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 { func CleanFSRoot(root string) string {
root = strings.TrimSpace(root) if strings.TrimSpace(root) == "" || root == "." {
if root == "" || root == "." {
return "." return "."
} }
return path.Clean(root) 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. // ResolveFSPath resolves userPath from baseDir and keeps it inside root.
func ResolveFSPath(root string, baseDir string, userPath string) (string, string, error) { func ResolveFSPath(root string, baseDir string, userPath string) (string, string, error) {
cleanRoot := CleanFSRoot(root) cleanRoot := CleanFSRoot(root)
cleanBase := path.Clean(strings.TrimSpace(baseDir)) cleanBase := path.Clean(baseDir)
if cleanBase == "" { if strings.TrimSpace(baseDir) == "" {
cleanBase = cleanRoot cleanBase = cleanRoot
} }
if !containsFSPath(cleanRoot, cleanBase) { if !containsFSPath(cleanRoot, cleanBase) {
return "", "", fmt.Errorf("base path %q is outside source root %q", cleanBase, cleanRoot) return "", "", fmt.Errorf("base path %q is outside source root %q", cleanBase, cleanRoot)
} }
cleanUserPath := strings.TrimSpace(userPath) if strings.TrimSpace(userPath) == "" {
if cleanUserPath == "" {
return "", "", fmt.Errorf("path is required") return "", "", fmt.Errorf("path is required")
} }
cleanUserPath = path.Clean(cleanUserPath) if path.IsAbs(userPath) {
if path.IsAbs(cleanUserPath) {
return "", "", fmt.Errorf("path %q must be relative", userPath) return "", "", fmt.Errorf("path %q must be relative", userPath)
} }
cleanUserPath := path.Clean(userPath)
resolved := path.Clean(path.Join(cleanBase, cleanUserPath)) resolved := path.Clean(path.Join(cleanBase, cleanUserPath))
if !containsFSPath(cleanRoot, resolved) { if !containsFSPath(cleanRoot, resolved) {
@@ -130,13 +134,6 @@ func containsFSPath(root string, name string) bool {
return name == root || strings.HasPrefix(name, strings.TrimSuffix(root, "/")+"/") 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 { func IsYAMLFile(name string) bool {
return strings.HasSuffix(name, ".yaml") || strings.HasSuffix(name, ".yml") return strings.HasSuffix(name, ".yaml") || strings.HasSuffix(name, ".yml")
} }

View File

@@ -54,7 +54,7 @@ func TestFindFSYAMLFilesNestedSortedAndFiltered(t *testing.T) {
"other/ignored.yaml": &fstest.MapFile{Data: []byte("id: ignored")}, "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 { if err != nil {
t.Fatalf("expected no error, got %v", err) t.Fatalf("expected no error, got %v", err)
} }
@@ -98,8 +98,11 @@ func TestCleanFSRoot(t *testing.T) {
want string want string
}{ }{
{name: "empty", root: "", want: "."}, {name: "empty", root: "", want: "."},
{name: "whitespace only", root: " \t ", want: "."},
{name: "dot", root: ".", 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 { for _, tc := range tests {
@@ -158,6 +161,22 @@ func TestResolveFSPath(t *testing.T) {
wantPath: "prompts/shared/user.tmpl", wantPath: "prompts/shared/user.tmpl",
wantDisplay: "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", name: "escape rejected",
root: "prompts", 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) { func TestIsYAMLFile(t *testing.T) {
tests := []struct { tests := []struct {
name string name string

View File

@@ -1,5 +1,5 @@
// Package jsonvalue validates and defensively copies JSON-compatible value // Package jsonvalue validates and defensively copies bounded JSON-compatible
// trees used by public configuration and request boundaries. // value trees used by configuration, request, and prepared-state boundaries.
package jsonvalue package jsonvalue
import ( import (
@@ -8,23 +8,39 @@ import (
"math" "math"
"reflect" "reflect"
"sort" "sort"
"strconv"
) )
const maxSafeJSONInteger = 1<<53 - 1 const (
maxContainerDepth = 100
maxProducedNodes = 100_000
)
type visit struct { type visit struct {
typ reflect.Type typ reflect.Type
ptr uintptr 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 // 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) { func CopyMap(src map[string]any) (map[string]any, error) {
if src == nil { if src == nil {
return nil, 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 { if err != nil {
return nil, err return nil, err
} }
@@ -35,88 +51,121 @@ func CopyMap(src map[string]any) (map[string]any, error) {
return out, nil return out, nil
} }
func copyValue(value reflect.Value, path string, seen map[visit]struct{}) (any, error) { func copyValue(
if !value.IsValid() { 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 return nil, nil
} }
if value.Kind() == reflect.Interface { value = resolved
if value.IsNil() {
return nil, nil
}
return copyValue(value.Elem(), path, seen)
}
if !value.CanInterface() { if !value.CanInterface() {
return nil, fmt.Errorf("%s: value cannot be copied", path) return nil, fmt.Errorf("%s: value cannot be copied", path)
} }
if number, ok := value.Interface().(json.Number); ok { 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) return nil, fmt.Errorf("%s: invalid JSON number", path)
} }
f, err := strconv.ParseFloat(number.String(), 64) if err := state.produceNode(path); err != nil {
if err != nil || math.IsNaN(f) || math.IsInf(f, 0) { return nil, err
return nil, fmt.Errorf("%s: invalid JSON number", path)
} }
return number, nil return number, nil
} }
switch value.Kind() { switch value.Kind() {
case reflect.Bool, reflect.String: case reflect.Bool, reflect.String:
if err := state.produceNode(path); err != nil {
return nil, err
}
return value.Interface(), nil return value.Interface(), nil
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
if value.Int() < -maxSafeJSONInteger || value.Int() > maxSafeJSONInteger { if err := state.produceNode(path); err != nil {
return nil, fmt.Errorf("%s: integer is outside the JSON-safe range", path) return nil, err
} }
return value.Interface(), nil return value.Interface(), nil
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr: case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
if value.Uint() > maxSafeJSONInteger { if err := state.produceNode(path); err != nil {
return nil, fmt.Errorf("%s: integer is outside the JSON-safe range", path) return nil, err
} }
return value.Interface(), nil return value.Interface(), nil
case reflect.Float32, reflect.Float64: case reflect.Float32, reflect.Float64:
number := value.Convert(reflect.TypeOf(float64(0))).Float() number := value.Float()
if math.IsNaN(number) || math.IsInf(number, 0) { if math.IsNaN(number) || math.IsInf(number, 0) {
return nil, fmt.Errorf("%s: floating-point value must be finite", path) 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 return value.Interface(), nil
case reflect.Pointer: case reflect.Map:
if value.IsNil() { if value.IsNil() {
if err := state.produceNode(path); err != nil {
return nil, err
}
return nil, nil return nil, nil
} }
current := visit{typ: value.Type(), ptr: value.Pointer()} nextDepth, err := state.enterContainer(path, containerDepth)
if _, ok := seen[current]; ok { if err != nil {
return nil, fmt.Errorf("%s: cyclic value is not supported", path) return nil, err
} }
seen[current] = struct{}{} return copyMapValue(value, path, state, allowEmptyMapKeys, nextDepth)
defer delete(seen, current)
return copyValue(value.Elem(), path, seen)
case reflect.Map:
return copyMapValue(value, path, seen)
case reflect.Slice: case reflect.Slice:
if value.IsNil() { if value.IsNil() {
if err := state.produceNode(path); err != nil {
return nil, err
}
return nil, nil 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: 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: default:
return nil, fmt.Errorf("%s: unsupported JSON value type %s", path, value.Type()) 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) { func copyMapValue(
if value.IsNil() { value reflect.Value,
return nil, nil path string,
} state *traversalState,
allowEmptyMapKeys bool,
containerDepth int,
) (any, error) {
if value.Type().Key().Kind() != reflect.String { if value.Type().Key().Kind() != reflect.String {
return nil, fmt.Errorf("%s: map key type %s is not supported", path, value.Type().Key()) 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()} 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) return nil, fmt.Errorf("%s: cyclic value is not supported", path)
} }
seen[current] = struct{}{} state.active[current] = struct{}{}
defer delete(seen, current) defer delete(state.active, current)
keys := value.MapKeys() keys := value.MapKeys()
sort.Slice(keys, func(i, j int) bool { 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() elementType := value.Type().Elem()
for _, key := range keys { for _, key := range keys {
name := key.String() name := key.String()
if name == "" { if name == "" && !allowEmptyMapKeys {
return nil, fmt.Errorf("%s: map key must not be empty", path) 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 { if err != nil {
return nil, err return nil, err
} }
@@ -171,22 +226,41 @@ func copyMapValue(value reflect.Value, path string, seen map[visit]struct{}) (an
return out, nil 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 var current visit
if value.Kind() == reflect.Slice { if value.Kind() == reflect.Slice {
current = visit{typ: value.Type(), ptr: value.Pointer()} 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) return nil, fmt.Errorf("%s: cyclic value is not supported", path)
} }
seen[current] = struct{}{} state.active[current] = struct{}{}
defer delete(seen, current) defer delete(state.active, current)
} }
values := make([]any, value.Len()) values := make([]any, value.Len())
preserveType := true preserveType := true
elementType := value.Type().Elem() elementType := value.Type().Elem()
for i := 0; i < value.Len(); i++ { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -222,6 +296,73 @@ func copySequenceValue(value reflect.Value, path string, seen map[visit]struct{}
return out, nil 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 { func canAssignNil(typ reflect.Type) bool {
switch typ.Kind() { switch typ.Kind() {
case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice: case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice:

View File

@@ -1,97 +1,332 @@
package jsonvalue_test package jsonvalue
import ( import (
"encoding/json" "encoding/json"
"math" "math"
"reflect" "reflect"
"strings"
"testing" "testing"
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
) )
func TestCopyMapPreservesTypesAndIsolatesMutations(t *testing.T) { type (
nested := map[string]int{"limit": 2} namedBool bool
sequence := []string{"one", "two"} namedString string
input := map[string]any{ namedInt64 int64
"count": int64(7), namedUint64 uint64
"number": json.Number("-1.25e+2"), namedFloat32 float32
"nested": nested, namedFloat64 float64
"sequence": sequence, namedKey string
} namedMap map[namedKey]namedInt64
namedSlice []namedString
copied, err := jsonvalue.CopyMap(input) namedArray [1]map[string]int
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
func TestCopyPreservesSupportedScalarAndNumberTypes(t *testing.T) {
maxInt := int(^uint(0) >> 1)
minInt := -maxInt - 1
tests := []struct { tests := []struct {
name string name string
value any value any
}{ }{
{name: "empty nested key", value: map[string]int{"": 1}}, {name: "bool", value: true},
{name: "non-string map key", value: map[int]string{1: "one"}}, {name: "named bool", value: namedBool(true)},
{name: "unsupported value", value: make(chan int)}, {name: "string", value: "value"},
{name: "cyclic map", value: cyclicMap}, {name: "named string", value: namedString("value")},
{name: "cyclic slice", value: cyclicSlice}, {name: "int", value: minInt},
{name: "NaN", value: math.NaN()}, {name: "int8", value: int8(-1 << 7)},
{name: "positive infinity", value: math.Inf(1)}, {name: "int16", value: int16(-1 << 15)},
{name: "unsafe signed integer", value: int64(1 << 53)}, {name: "int32", value: int32(-1 << 31)},
{name: "unsafe unsigned integer", value: uint64(1 << 53)}, {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 { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { 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") t.Fatal("expected validation error")
} }
}) })
} }
} }
func TestCopyMapValidatesJSONNumberSyntaxAndRange(t *testing.T) { func TestCopyPreservesCompatibleCollectionsAndNilEmptyDistinctions(t *testing.T) {
for _, number := range []json.Number{"0", "-1", "1.25", "-1.25e+2"} { collections := []struct {
t.Run("valid "+number.String(), func(t *testing.T) { name string
got, err := jsonvalue.CopyMap(map[string]any{"value": number}) 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 { if err != nil {
t.Fatalf("copy valid JSON number: %v", err) t.Fatalf("copy collection: %v", err)
} }
if !reflect.DeepEqual(got["value"], number) { if !reflect.DeepEqual(got, tc.value) || reflect.TypeOf(got) != reflect.TypeOf(tc.value) {
t.Fatalf("JSON number changed: got %#v want %#v", got["value"], number) 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"} { var nilMap map[string]int
t.Run("invalid "+number.String(), func(t *testing.T) { var nilSlice []string
if _, err := jsonvalue.CopyMap(map[string]any{"value": number}); err == nil { var nilPointer *namedInt64
t.Fatal("expected invalid JSON number error") 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
}

View File

@@ -25,6 +25,20 @@ var (
ErrMalformedResponse = errors.New("malformed llm response") 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 { type OpenAICompatibleConfig struct {
BaseURL string BaseURL string
Model string Model string
@@ -39,9 +53,11 @@ type OpenAICompatibleClient struct {
} }
func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) { func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) {
baseURL := strings.TrimSpace(cfg.BaseURL) baseURL := ""
if baseURL != "" { if strings.TrimSpace(cfg.BaseURL) != "" {
if _, err := url.ParseRequestURI(baseURL); err != nil { var err error
baseURL, err = domain.NormalizeOpenAICompatibleBaseEndpoint(cfg.BaseURL)
if err != nil {
return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err) return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err)
} }
} }
@@ -63,25 +79,29 @@ func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleCli
} }
return &OpenAICompatibleClient{ return &OpenAICompatibleClient{
baseURL: strings.TrimRight(baseURL, "/"), baseURL: baseURL,
defaultModel: cfg.Model, defaultModel: cfg.Model,
httpClient: client, httpClient: client,
}, nil }, nil
} }
func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) { func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) {
if req.Target.TimeoutSeconds < 0 { if err := domain.ValidateExecutionTargetSettings(req.Target); err != nil {
return nil, fmt.Errorf("%w: timeout_seconds must be greater than or equal to 0", ErrInvalidRequest) return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
} }
endpoint := strings.TrimSpace(req.Target.Endpoint) selectedEndpoint := req.Target.Endpoint
if endpoint == "" { if strings.TrimSpace(selectedEndpoint) == "" {
endpoint = c.baseURL selectedEndpoint = c.baseURL
} }
if endpoint == "" { endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(selectedEndpoint)
return nil, fmt.Errorf("%w: endpoint is required", ErrInvalidRequest) 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) wireReq, err := openAIChatRequestFromGenerateRequest(req, c.defaultModel)
if err != nil { if err != nil {
@@ -113,13 +133,18 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
return nil, fmt.Errorf("%w: failed to create request: %v", ErrRequestFailed, err) return nil, fmt.Errorf("%w: failed to create request: %v", ErrRequestFailed, err)
} }
httpReq.Header.Set("Content-Type", "application/json") httpReq.Header.Set("Content-Type", "application/json")
if apiKey := strings.TrimSpace(req.Target.APIKey); apiKey != "" { apiKey := strings.TrimSpace(req.Target.APIKey)
httpReq.Header.Set("Authorization", "Bearer "+apiKey) envName := strings.TrimSpace(req.Target.APIKeyEnv)
} else if envName := strings.TrimSpace(req.Target.APIKeyEnv); envName != "" { if apiKey == "" && envName != "" {
apiKey := strings.TrimSpace(os.Getenv(envName)) apiKey = strings.TrimSpace(os.Getenv(envName))
if apiKey == "" { }
if apiKey == "" && req.Target.APIKeyRequired {
if envName != "" {
return nil, fmt.Errorf("%w: api key environment variable %q is not set", ErrInvalidRequest, envName) return nil, fmt.Errorf("%w: api key environment variable %q is not set", ErrInvalidRequest, envName)
} }
return nil, fmt.Errorf("%w: api key is required", ErrInvalidRequest)
}
if apiKey != "" {
httpReq.Header.Set("Authorization", "Bearer "+apiKey) httpReq.Header.Set("Authorization", "Bearer "+apiKey)
} }
@@ -130,18 +155,24 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
httpResp, err := httpClient.Do(httpReq) httpResp, err := httpClient.Do(httpReq)
if err != nil { if err != nil {
return nil, fmt.Errorf("%w: %v", ErrRequestFailed, err) return nil, &requestFailedError{cause: err}
} }
defer httpResp.Body.Close() defer httpResp.Body.Close()
if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 { if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 {
_, _ = io.Copy(io.Discard, io.LimitReader(httpResp.Body, 4096)) return nil, providerHTTPErrorFromBody(
return nil, fmt.Errorf("%w: status=%d", ErrUnexpectedStatus, httpResp.StatusCode) httpResp.StatusCode,
httpResp.ContentLength,
httpResp.Body,
)
}
if httpResp.ContentLength > maxOpenAIChatResponseBytes {
return nil, openAIChatResponseTooLargeError()
} }
var wireResp openAIChatResponse wireResp, err := decodeOpenAIChatResponse(httpResp.Body)
if err := json.NewDecoder(httpResp.Body).Decode(&wireResp); err != nil { if err != nil {
return nil, fmt.Errorf("%w: failed to decode response: %v", ErrMalformedResponse, err) return nil, err
} }
if len(wireResp.Choices) == 0 { if len(wireResp.Choices) == 0 {
@@ -164,6 +195,46 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
}, nil }, 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) { func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultModel string) (openAIChatRequest, error) {
model := strings.TrimSpace(req.Target.Model) model := strings.TrimSpace(req.Target.Model)
if model == "" { if model == "" {

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,189 @@
package llm
import (
"encoding/json"
"errors"
"fmt"
"io"
"strings"
"unicode"
)
const (
maxProviderErrorResponseBytes int64 = 64 << 10
maxProviderErrorIdentifierRunes = 256
maxProviderErrorMessageRunes = 4096
)
// ProviderHTTPError describes a non-success response from an LLM provider.
type ProviderHTTPError struct {
statusCode int
providerCode string
providerType string
providerMessage string
}
func (e *ProviderHTTPError) StatusCode() int {
if e == nil {
return 0
}
return e.statusCode
}
func (e *ProviderHTTPError) ProviderCode() string {
if e == nil {
return ""
}
return e.providerCode
}
func (e *ProviderHTTPError) ProviderType() string {
if e == nil {
return ""
}
return e.providerType
}
func (e *ProviderHTTPError) ProviderMessage() string {
if e == nil {
return ""
}
return e.providerMessage
}
func (e *ProviderHTTPError) Error() string {
if e == nil || e.statusCode == 0 {
return ErrUnexpectedStatus.Error()
}
return fmt.Sprintf("%s: status=%d", ErrUnexpectedStatus, e.statusCode)
}
func (e *ProviderHTTPError) GoString() string {
return e.Error()
}
func (e *ProviderHTTPError) Unwrap() error {
return ErrUnexpectedStatus
}
type providerErrorDetails struct {
providerCode string
providerType string
providerMessage string
}
func newProviderHTTPError(statusCode int, details providerErrorDetails) *ProviderHTTPError {
return &ProviderHTTPError{
statusCode: statusCode,
providerCode: details.providerCode,
providerType: details.providerType,
providerMessage: details.providerMessage,
}
}
func providerHTTPErrorFromBody(statusCode int, contentLength int64, body io.Reader) *ProviderHTTPError {
if contentLength > maxProviderErrorResponseBytes {
return newProviderHTTPError(statusCode, providerErrorDetails{})
}
limited := &io.LimitedReader{
R: body,
N: maxProviderErrorResponseBytes + 1,
}
contents, err := io.ReadAll(limited)
if err != nil || limited.N == 0 {
return newProviderHTTPError(statusCode, providerErrorDetails{})
}
return newProviderHTTPError(statusCode, parseProviderErrorEnvelope(contents))
}
func parseProviderErrorEnvelope(body []byte) providerErrorDetails {
decoder := json.NewDecoder(strings.NewReader(string(body)))
decoder.UseNumber()
var envelope map[string]json.RawMessage
if err := decoder.Decode(&envelope); err != nil {
return providerErrorDetails{}
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return providerErrorDetails{}
}
rawError, ok := envelope["error"]
if !ok {
return providerErrorDetails{}
}
var providerError map[string]json.RawMessage
if err := json.Unmarshal(rawError, &providerError); err != nil || providerError == nil {
return providerErrorDetails{}
}
var details providerErrorDetails
if raw, ok := providerError["message"]; ok {
var value string
if json.Unmarshal(raw, &value) == nil {
details.providerMessage = normalizeProviderErrorMessage(value)
}
}
if raw, ok := providerError["type"]; ok {
var value string
if json.Unmarshal(raw, &value) == nil {
details.providerType = normalizeProviderErrorIdentifier(value)
}
}
if raw, ok := providerError["code"]; ok {
var value any
fieldDecoder := json.NewDecoder(strings.NewReader(string(raw)))
fieldDecoder.UseNumber()
if fieldDecoder.Decode(&value) == nil {
switch value := value.(type) {
case string:
details.providerCode = normalizeProviderErrorIdentifier(value)
case json.Number:
details.providerCode = normalizeProviderErrorIdentifier(value.String())
}
}
}
return details
}
func normalizeProviderErrorIdentifier(value string) string {
normalized := normalizeProviderErrorText(value)
if len([]rune(normalized)) > maxProviderErrorIdentifierRunes {
return ""
}
return normalized
}
func normalizeProviderErrorMessage(value string) string {
normalized := normalizeProviderErrorText(value)
runes := []rune(normalized)
if len(runes) <= maxProviderErrorMessageRunes {
return normalized
}
return string(runes[:maxProviderErrorMessageRunes-1]) + "…"
}
func normalizeProviderErrorText(value string) string {
value = strings.ToValidUTF8(value, "<22>")
var result strings.Builder
result.Grow(len(value))
separatorPending := false
for _, r := range value {
if unicode.IsSpace(r) || unicode.IsControl(r) || unicode.In(r, unicode.Cf) {
if result.Len() > 0 {
separatorPending = true
}
continue
}
if separatorPending {
result.WriteByte(' ')
separatorPending = false
}
result.WriteRune(r)
}
return result.String()
}

View File

@@ -0,0 +1,255 @@
package llm
import (
"errors"
"fmt"
"io"
"reflect"
"strings"
"testing"
"unicode/utf8"
)
type guardedReader struct {
reader io.Reader
remaining int64
bytes int64
violated bool
}
func (r *guardedReader) Read(buffer []byte) (int, error) {
if int64(len(buffer)) > r.remaining {
r.violated = true
return 0, errors.New("reader was read past its allowed boundary")
}
n, err := r.reader.Read(buffer)
r.bytes += int64(n)
r.remaining -= int64(n)
return n, err
}
type failingReader struct {
err error
}
func (r failingReader) Read([]byte) (int, error) {
return 0, r.err
}
func TestProviderHTTPErrorEnvelopeParsing(t *testing.T) {
tests := []struct {
name string
body string
want providerErrorDetails
}{
{
name: "all supported string fields",
body: `{"error":{"message":"diagnostic","type":"invalid_request_error","code":"unsupported_parameter"}}`,
want: providerErrorDetails{providerMessage: "diagnostic", providerType: "invalid_request_error", providerCode: "unsupported_parameter"},
},
{
name: "integer code",
body: `{"error":{"code":17}}`,
want: providerErrorDetails{providerCode: "17"},
},
{
name: "fractional code",
body: `{"error":{"code":1.25}}`,
want: providerErrorDetails{providerCode: "1.25"},
},
{
name: "exponent code",
body: `{"error":{"code":6.02e+23}}`,
want: providerErrorDetails{providerCode: "6.02e+23"},
},
{
name: "invalid fields do not discard valid fields",
body: `{"error":{"message":null,"type":"invalid_request_error","code":false}}`,
want: providerErrorDetails{providerType: "invalid_request_error"},
},
{
name: "unknown fields are ignored",
body: `{"trace":"do not retain","error":{"param":"temperature","metadata":{"secret":"x"}}}`,
want: providerErrorDetails{},
},
{name: "missing error", body: `{}`, want: providerErrorDetails{}},
{name: "null error", body: `{"error":null}`, want: providerErrorDetails{}},
{name: "scalar error", body: `{"error":"nope"}`, want: providerErrorDetails{}},
{name: "empty error", body: `{"error":{}}`, want: providerErrorDetails{}},
{name: "malformed", body: `{"error":`, want: providerErrorDetails{}},
{name: "truncated", body: `{"error":{"message":"x"`, want: providerErrorDetails{}},
{name: "trailing garbage", body: `{"error":{"message":"x"}} garbage`, want: providerErrorDetails{}},
{name: "second document", body: `{"error":{"message":"x"}} {}`, want: providerErrorDetails{}},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
if got := parseProviderErrorEnvelope([]byte(tc.body)); !reflect.DeepEqual(got, tc.want) {
t.Fatalf("parseProviderErrorEnvelope() = %#v, want %#v", got, tc.want)
}
})
}
}
func TestProviderErrorTextNormalizationAndLimits(t *testing.T) {
validIdentifier := strings.Repeat("界", maxProviderErrorIdentifierRunes)
validMessage := strings.Repeat("界", maxProviderErrorMessageRunes)
tests := []struct {
name string
got string
want string
}{
{name: "multibyte text", got: "Grüße 世界", want: "Grüße 世界"},
{name: "invalid UTF-8", got: string([]byte{'a', 0xff, 'b'}), want: "a<>b"},
{name: "whitespace control and format runs", got: " \n\talpha\x00\u200b\u200bbeta \r ", want: "alpha beta"},
{name: "blank normalization", got: "\t\u200b\n", want: ""},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
if got := normalizeProviderErrorText(tc.got); got != tc.want {
t.Fatalf("normalizeProviderErrorText() = %q, want %q", got, tc.want)
}
})
}
if got := normalizeProviderErrorIdentifier(validIdentifier); got != validIdentifier {
t.Fatalf("exact identifier boundary = %q, want retained value", got)
}
if got := normalizeProviderErrorIdentifier(validIdentifier + "界"); got != "" {
t.Fatalf("overlong identifier = %q, want empty", got)
}
if got := normalizeProviderErrorMessage(validMessage); got != validMessage {
t.Fatalf("exact message boundary = %q, want retained value", got)
}
wantTruncatedMessage := strings.Repeat("界", maxProviderErrorMessageRunes-1) + "…"
if got := normalizeProviderErrorMessage(validMessage + "界"); got != wantTruncatedMessage {
t.Fatalf("overlong message length = %d, want %d", utf8.RuneCountInString(got), maxProviderErrorMessageRunes)
}
}
func TestProviderHTTPErrorIdentityAndFormatting(t *testing.T) {
const marker = "provider-secret-marker"
err := newProviderHTTPError(429, providerErrorDetails{
providerCode: marker + "-code",
providerType: marker + "-type",
providerMessage: marker + "-message",
})
if err.StatusCode() != 429 || err.ProviderCode() != marker+"-code" || err.ProviderType() != marker+"-type" || err.ProviderMessage() != marker+"-message" {
t.Fatalf("accessors returned unexpected values: %#v", err)
}
if !errors.Is(err, ErrUnexpectedStatus) {
t.Fatalf("errors.Is(%v, ErrUnexpectedStatus) = false", err)
}
for _, rendered := range []string{fmt.Sprintf("%v", err), fmt.Sprintf("%+v", err), fmt.Sprintf("%#v", err)} {
if rendered != "llm returned non-success status: status=429" {
t.Fatalf("formatted error = %q", rendered)
}
if strings.Contains(rendered, marker) {
t.Fatalf("formatted error exposed provider marker: %q", rendered)
}
}
var nilError *ProviderHTTPError
if nilError.StatusCode() != 0 || nilError.ProviderCode() != "" || nilError.ProviderType() != "" || nilError.ProviderMessage() != "" {
t.Fatal("nil accessors returned provider values")
}
if nilError.Error() != "llm returned non-success status" || nilError.GoString() != "llm returned non-success status" || !errors.Is(nilError, ErrUnexpectedStatus) {
t.Fatalf("nil error behavior is not safe: %v", nilError)
}
zero := &ProviderHTTPError{}
if zero.Error() != "llm returned non-success status" || zero.GoString() != "llm returned non-success status" || !errors.Is(zero, ErrUnexpectedStatus) {
t.Fatalf("zero error behavior is not safe: %v", zero)
}
}
func TestProviderHTTPErrorBodyBounds(t *testing.T) {
const (
statusCode = 502
marker = "provider-body-marker"
)
ordinaryBody := `{"error":{"message":"` + marker + `"}}`
exactLimitBody := ordinaryBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)-len(ordinaryBody))
overLimitBody := ordinaryBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)+1-len(ordinaryBody))
tests := []struct {
name string
contentLength int64
reader io.Reader
wantRead int64
wantMessage string
}{
{
name: "recognized envelope",
contentLength: int64(len(ordinaryBody)),
reader: strings.NewReader(ordinaryBody),
wantRead: int64(len(ordinaryBody)),
wantMessage: marker,
},
{
name: "exact limit",
contentLength: maxProviderErrorResponseBytes,
reader: strings.NewReader(exactLimitBody),
wantRead: maxProviderErrorResponseBytes,
wantMessage: marker,
},
{
name: "declared oversize does not read",
contentLength: maxProviderErrorResponseBytes + 1,
reader: strings.NewReader(ordinaryBody),
wantRead: 0,
},
{
name: "unknown length oversize",
contentLength: -1,
reader: strings.NewReader(overLimitBody),
wantRead: maxProviderErrorResponseBytes + 1,
},
{
name: "underreported oversize",
contentLength: maxProviderErrorResponseBytes,
reader: strings.NewReader(overLimitBody),
wantRead: maxProviderErrorResponseBytes + 1,
},
{
name: "read failure",
contentLength: -1,
reader: failingReader{err: errors.New("read failure")},
wantRead: 0,
},
{
name: "empty body",
contentLength: 0,
reader: strings.NewReader(""),
wantRead: 0,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
reader := &guardedReader{
reader: tc.reader,
remaining: maxProviderErrorResponseBytes + 1,
}
err := providerHTTPErrorFromBody(statusCode, tc.contentLength, reader)
if err == nil || err.StatusCode() != statusCode {
t.Fatalf("error status = %v, want %d", err, statusCode)
}
if reader.bytes != tc.wantRead {
t.Fatalf("body bytes read = %d, want %d", reader.bytes, tc.wantRead)
}
if reader.violated {
t.Fatal("body reader was asked to read beyond the overflow probe")
}
if got := err.ProviderMessage(); got != tc.wantMessage {
t.Fatalf("provider message = %q, want %q", got, tc.wantMessage)
}
if tc.wantMessage == "" {
if err.ProviderCode() != "" || err.ProviderType() != "" || strings.Contains(err.Error(), marker) {
t.Fatalf("discarded details were retained: %#v", err)
}
}
})
}
}

View File

@@ -0,0 +1,145 @@
package llm
import (
"context"
"errors"
"io"
"net/http"
"strings"
"testing"
)
func TestOpenAICompatibleClientStructuredNonSuccessResponse(t *testing.T) {
body := `{"error":{"message":" provider\nmessage\u200b","type":"invalid\ttype","code":1.5e+4}}`
responseBody := &countingReadCloser{reader: strings.NewReader(body)}
client := newNonSuccessResponseClient(t, http.StatusBadRequest, int64(len(body)), responseBody)
response, err := client.Generate(context.Background(), ordinaryGenerateRequest())
if response != nil {
t.Fatalf("response = %#v, want nil", response)
}
if !errors.Is(err, ErrUnexpectedStatus) {
t.Fatalf("errors.Is(%v, ErrUnexpectedStatus) = false", err)
}
var providerHTTPError *ProviderHTTPError
if !errors.As(err, &providerHTTPError) {
t.Fatalf("error = %T, want *ProviderHTTPError", err)
}
if providerHTTPError.StatusCode() != http.StatusBadRequest || providerHTTPError.ProviderCode() != "1.5e+4" || providerHTTPError.ProviderType() != "invalid type" || providerHTTPError.ProviderMessage() != "provider message" {
t.Fatalf("provider error = %#v", providerHTTPError)
}
if !responseBody.closed {
t.Fatal("non-success response body was not closed")
}
}
func TestOpenAICompatibleClientNonSuccessBodyOwnership(t *testing.T) {
const marker = "provider-body-marker"
normalBody := `{"error":{"message":"` + marker + `"}}`
overLimitBody := normalBody + strings.Repeat(" ", int(maxProviderErrorResponseBytes)+1-len(normalBody))
tests := []struct {
name string
contentLength int64
reader io.Reader
wantRead int64
wantMessage string
}{
{
name: "normal",
contentLength: int64(len(normalBody)),
reader: strings.NewReader(normalBody),
wantRead: int64(len(normalBody)),
wantMessage: marker,
},
{
name: "declared oversize",
contentLength: maxProviderErrorResponseBytes + 1,
reader: strings.NewReader(normalBody),
wantRead: 0,
},
{
name: "streamed oversize",
contentLength: -1,
reader: &guardedReader{
reader: strings.NewReader(overLimitBody),
remaining: maxProviderErrorResponseBytes + 1,
},
wantRead: maxProviderErrorResponseBytes + 1,
},
{
name: "underreported oversize",
contentLength: maxProviderErrorResponseBytes,
reader: &guardedReader{
reader: strings.NewReader(overLimitBody),
remaining: maxProviderErrorResponseBytes + 1,
},
wantRead: maxProviderErrorResponseBytes + 1,
},
{
name: "malformed",
contentLength: 1,
reader: strings.NewReader("{"),
wantRead: 1,
},
{
name: "read failure",
contentLength: -1,
reader: failingReader{err: errors.New("response read failed")},
wantRead: 0,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
body := &countingReadCloser{reader: tc.reader}
client := newNonSuccessResponseClient(t, http.StatusBadGateway, tc.contentLength, body)
response, err := client.Generate(context.Background(), ordinaryGenerateRequest())
if response != nil {
t.Fatalf("response = %#v, want nil", response)
}
var providerHTTPError *ProviderHTTPError
if !errors.As(err, &providerHTTPError) {
t.Fatalf("error = %T, want *ProviderHTTPError", err)
}
if !body.closed {
t.Fatal("response body was not closed")
}
if body.bytesRead != tc.wantRead {
t.Fatalf("body bytes read = %d, want %d", body.bytesRead, tc.wantRead)
}
if body.bytesRead > maxProviderErrorResponseBytes+1 {
t.Fatalf("body bytes read = %d, exceeds overflow probe", body.bytesRead)
}
if guarded, ok := tc.reader.(*guardedReader); ok && guarded.violated {
t.Fatal("body reader was asked to read beyond the overflow probe")
}
if got := providerHTTPError.ProviderMessage(); got != tc.wantMessage {
t.Fatalf("provider message = %q, want %q", got, tc.wantMessage)
}
if tc.wantMessage == "" && (providerHTTPError.ProviderCode() != "" || providerHTTPError.ProviderType() != "" || strings.Contains(providerHTTPError.Error(), marker)) {
t.Fatalf("discarded details were retained: %#v", providerHTTPError)
}
})
}
}
func newNonSuccessResponseClient(t *testing.T, statusCode int, contentLength int64, body io.ReadCloser) *OpenAICompatibleClient {
t.Helper()
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "https://provider.example/v1",
Model: "m",
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: statusCode,
ContentLength: contentLength,
Body: body,
}, nil
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
return client
}

View File

@@ -0,0 +1,3 @@
id: rakestrawhome-gemma-4-31b
backend: rakestrawhome
model: google/gemma-4-31b-it

View File

@@ -2,7 +2,6 @@ package builtin
import ( import (
"embed" "embed"
"strings"
"gitea.maximumdirect.net/eric/promptkit/internal/profile" "gitea.maximumdirect.net/eric/promptkit/internal/profile"
) )
@@ -15,17 +14,3 @@ var assets embed.FS
func NewRepository() profile.Repository { func NewRepository() profile.Repository {
return profile.NewFSRepository(assets, assetRoot) 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))
}

View File

@@ -2,14 +2,11 @@ package builtin
import ( import (
"context" "context"
"errors"
"io/fs" "io/fs"
"strings" "strings"
"testing" "testing"
"gitea.maximumdirect.net/eric/promptkit/internal/backend" "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" "gopkg.in/yaml.v3"
) )
@@ -29,8 +26,8 @@ func TestBuiltInProfilesValidateThroughRepository(t *testing.T) {
if p.ID != id { if p.ID != id {
t.Fatalf("expected profile id %q, got %q", id, p.ID) t.Fatalf("expected profile id %q, got %q", id, p.ID)
} }
if p.BackendID != backend.OpenRouterID { if !builtInBackendIDs[p.BackendID] {
t.Fatalf("expected profile %q to select %q, got %q", id, backend.OpenRouterID, p.BackendID) t.Fatalf("expected profile %q to select a maintained built-in, got %q", id, p.BackendID)
} }
if p.Endpoint != "" || p.APIKeyEnv != "" { if p.Endpoint != "" || p.APIKeyEnv != "" {
t.Fatalf("expected profile %q to inherit backend connection settings, got endpoint=%q api_key_env=%q", id, p.Endpoint, p.APIKeyEnv) t.Fatalf("expected profile %q to inherit backend connection settings, got endpoint=%q api_key_env=%q", id, p.Endpoint, p.APIKeyEnv)
@@ -43,6 +40,33 @@ func TestBuiltInProfilesDoNotContainDuplicateIDsOrRawAPIKeys(t *testing.T) {
loadBuiltInProfileIDs(t) loadBuiltInProfileIDs(t)
} }
func TestRakestrawhomeGemmaProfileUsesNativeDefaults(t *testing.T) {
p, err := NewRepository().GetProfile(context.Background(), "rakestrawhome-gemma-4-31b")
if err != nil {
t.Fatalf("load Rakestrawhome Gemma profile: %v", err)
}
if p.ID != "rakestrawhome-gemma-4-31b" ||
p.BackendID != backend.RakestrawHomeID ||
p.Model != "google/gemma-4-31b-it" ||
p.Endpoint != "" ||
p.Temperature != 0 ||
p.MaxTokens != 0 ||
p.TopP != 0 ||
p.TimeoutSeconds != 0 ||
p.ServiceTier != "" ||
p.ReasoningEffort != "" ||
p.APIKeyEnv != "" ||
p.APIKeyRequired ||
p.ExtraParams != nil {
t.Fatalf("unexpected Rakestrawhome Gemma profile: %#v", p)
}
}
var builtInBackendIDs = map[string]bool{
backend.OpenRouterID: true,
backend.RakestrawHomeID: true,
}
func loadBuiltInProfileIDs(t *testing.T) map[string]string { func loadBuiltInProfileIDs(t *testing.T) map[string]string {
t.Helper() t.Helper()
@@ -67,8 +91,9 @@ func loadBuiltInProfileIDs(t *testing.T) map[string]string {
if _, ok := raw["api_key"]; ok { if _, ok := raw["api_key"]; ok {
t.Fatalf("built-in profile %s contains raw api_key", name) t.Fatalf("built-in profile %s contains raw api_key", name)
} }
if raw["backend"] != backend.OpenRouterID { backendID, ok := raw["backend"].(string)
t.Fatalf("built-in profile %s does not select %q", name, backend.OpenRouterID) if !ok || !builtInBackendIDs[backendID] {
t.Fatalf("built-in profile %s does not select a maintained built-in: %#v", name, raw["backend"])
} }
if _, ok := raw["endpoint"]; ok { if _, ok := raw["endpoint"]; ok {
t.Fatalf("built-in profile %s repeats endpoint", name) t.Fatalf("built-in profile %s repeats endpoint", name)
@@ -91,53 +116,3 @@ func loadBuiltInProfileIDs(t *testing.T) map[string]string {
} }
return ids 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
}

View File

@@ -0,0 +1,47 @@
package profile
import (
"errors"
"strings"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
)
// NormalizeAndValidateDefinition normalizes and validates one source-local
// profile definition without resolving a base profile.
func NormalizeAndValidateDefinition(profile *domain.ExecutionProfile) error {
if profile == nil {
return errors.New("profile is required")
}
profile.ID = strings.TrimSpace(profile.ID)
profile.BaseProfileID = strings.TrimSpace(profile.BaseProfileID)
profile.BackendID = strings.TrimSpace(profile.BackendID)
profile.Endpoint = strings.TrimSpace(profile.Endpoint)
if profile.ID == "" {
return errors.New("id is required")
}
if profile.Endpoint != "" {
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(profile.Endpoint)
if err != nil {
return err
}
profile.Endpoint = endpoint
}
if profile.BaseProfileID == "" {
if profile.BackendID == "" && profile.Endpoint == "" {
return errors.New("backend or endpoint is required")
}
if strings.TrimSpace(profile.Model) == "" {
return errors.New("model is required")
}
}
return domain.ValidateExecutionTargetSettings(domain.ExecutionTarget{
Temperature: profile.Temperature,
MaxTokens: profile.MaxTokens,
TopP: profile.TopP,
TimeoutSeconds: profile.TimeoutSeconds,
})
}

View File

@@ -5,13 +5,14 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"io"
"io/fs" "io/fs"
"os" "os"
"path"
"strings" "strings"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/domain"
"gitea.maximumdirect.net/eric/promptkit/internal/filecatalog" "gitea.maximumdirect.net/eric/promptkit/internal/filecatalog"
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
"gopkg.in/yaml.v3" "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) { 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) return nil, fmt.Errorf("%w: profile id is required", ErrInvalidProfile)
} }
if fsys == nil { if fsys == nil {
@@ -94,42 +96,49 @@ func loadProfile(ctx context.Context, fsys fs.FS, root string, id string) (*doma
} }
relPath := filecatalog.DisplayPath(root, fullPath) relPath := filecatalog.DisplayPath(root, fullPath)
fileMatch := filecatalog.Stem(path.Base(fullPath)) == id
data, err := fs.ReadFile(fsys, fullPath) data, err := fs.ReadFile(fsys, fullPath)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to read profile file %s: %w", relPath, err) return nil, fmt.Errorf("failed to read profile file %s: %w", relPath, err)
} }
metadata := readProfileFileMetadata(data) metadata, metadataErr := readProfileFileMetadata(data)
idMatch := fileMatch || metadata.id == id idMatch := metadata.matchesID(id)
if metadataErr != nil {
if idMatch {
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, metadataErr)
}
continue
}
if metadata.hasRawAPIKey { if metadata.hasRawAPIKey {
if idMatch { if idMatch {
return nil, fmt.Errorf("%w: %s", ErrRawAPIKeyNotAllowed, relPath) return nil, fmt.Errorf("%w: %s", ErrRawAPIKeyNotAllowed, relPath)
} }
continue continue
} }
if !idMatch {
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)
}
continue 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 { if prof.ID != id {
continue continue
} }
prof.BackendID = strings.TrimSpace(prof.BackendID) prof.ExtraParams, err = jsonvalue.CopyMap(prof.ExtraParams)
if err := validateProfile(&prof); err != nil { if err != nil {
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidProfile, relPath, err)
}
if err := NormalizeAndValidateDefinition(prof); err != nil {
if errors.Is(err, ErrRawAPIKeyNotAllowed) { if errors.Is(err, ErrRawAPIKeyNotAllowed) {
return nil, fmt.Errorf("%w: %s", err, relPath) return nil, fmt.Errorf("%w: %s", err, relPath)
} }
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidProfile, relPath, err) return nil, fmt.Errorf("%w: %s: %v", ErrInvalidProfile, relPath, err)
} }
matches = append(matches, profileMatch{ matches = append(matches, profileMatch{
profile: &prof, profile: prof,
path: relPath, path: relPath,
}) })
} }
@@ -155,15 +164,36 @@ type profileMatch struct {
} }
type profileFileMetadata struct { type profileFileMetadata struct {
id string ids []string
hasRawAPIKey bool hasRawAPIKey bool
} }
func readProfileFileMetadata(data []byte) profileFileMetadata { func readProfileFileMetadata(data []byte) (profileFileMetadata, error) {
decoder := yaml.NewDecoder(bytes.NewReader(data))
var node yaml.Node var node yaml.Node
if err := yaml.NewDecoder(bytes.NewReader(data)).Decode(&node); err != nil { if err := decoder.Decode(&node); err != nil {
return profileFileMetadata{} 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 { if node.Kind != yaml.DocumentNode || len(node.Content) == 0 {
return profileFileMetadata{} return profileFileMetadata{}
} }
@@ -178,7 +208,7 @@ func readProfileFileMetadata(data []byte) profileFileMetadata {
value := mapping.Content[i+1] value := mapping.Content[i+1]
switch key.Value { switch key.Value {
case "id": case "id":
metadata.id = strings.TrimSpace(value.Value) metadata.ids = append(metadata.ids, strings.TrimSpace(value.Value))
case "api_key": case "api_key":
metadata.hasRawAPIKey = true metadata.hasRawAPIKey = true
} }
@@ -186,29 +216,41 @@ func readProfileFileMetadata(data []byte) profileFileMetadata {
return metadata return metadata
} }
func validateProfile(p *domain.ExecutionProfile) error { func (m profileFileMetadata) matchesID(id string) bool {
if strings.TrimSpace(p.ID) == "" { for _, candidate := range m.ids {
return errors.New("id is required") if candidate == id {
return true
}
} }
if strings.TrimSpace(p.BackendID) == "" && strings.TrimSpace(p.Endpoint) == "" { return false
return errors.New("backend or endpoint is required") }
}
if strings.TrimSpace(p.Model) == "" { func (m *profileFileMetadata) merge(other profileFileMetadata) {
return errors.New("model is required") m.ids = append(m.ids, other.ids...)
} m.hasRawAPIKey = m.hasRawAPIKey || other.hasRawAPIKey
}
if p.Temperature < 0 || p.Temperature > 2 {
return errors.New("temperature must be between 0 and 2") func decodeProfile(data []byte) (*domain.ExecutionProfile, error) {
} var prof domain.ExecutionProfile
if p.MaxTokens < 0 { decoder := yaml.NewDecoder(bytes.NewReader(data))
return errors.New("max_tokens must be greater than or equal to 0") decoder.KnownFields(true)
} if err := decoder.Decode(&prof); err != nil {
if p.TopP < 0 || p.TopP > 1 { return nil, err
return errors.New("top_p must be between 0 and 1") }
} if err := requireYAMLStreamEnd(decoder); err != nil {
if p.TimeoutSeconds < 0 { return nil, err
return errors.New("timeout_seconds must be greater than or equal to 0") }
} return &prof, nil
}
return 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")
} }

View 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)
}
}
})
}
}

View File

@@ -2,11 +2,13 @@ package profile
import ( import (
"context" "context"
"encoding/json"
"errors" "errors"
"fmt"
"io/fs"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"sync"
"testing" "testing"
"testing/fstest" "testing/fstest"
@@ -61,7 +63,7 @@ func TestFilesystemRepository_GetProfile(t *testing.T) {
wantErr bool wantErr bool
}{ }{
{name: "backend only", connection: "backend: ' openrouter '", wantBackend: "openrouter"}, {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: "both", connection: "backend: openrouter\nendpoint: http://localhost:8000/v1", wantBackend: "openrouter", wantEndpoint: "http://localhost:8000/v1"},
{name: "neither", wantErr: true}, {name: "neither", wantErr: true},
{name: "blank backend", connection: "backend: ' '", 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) { t.Run("duplicate profile IDs fail as ambiguous", func(t *testing.T) {
writeProfileTestFile(t, filepath.Join(tmpDir, "duplicate-profile-a.yaml"), ` writeProfileTestFile(t, filepath.Join(tmpDir, "duplicate-profile-a.yaml"), `
id: duplicate-profile 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") _, err := repo.GetProfile(ctx, "invalid_yaml")
if !errors.Is(err, ErrInvalidYAML) { if !errors.Is(err, ErrProfileNotFound) {
t.Fatalf("expected ErrInvalidYAML, got %v", err) t.Fatalf("expected ErrProfileNotFound, got %v", err)
} }
}) })
@@ -274,14 +219,14 @@ api_key: secret
}) })
t.Run("unknown field", func(t *testing.T) { 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) { if !errors.Is(err, ErrInvalidYAML) {
t.Fatalf("expected ErrInvalidYAML for strict decode unknown field, got %v", err) t.Fatalf("expected ErrInvalidYAML for strict decode unknown field, got %v", err)
} }
}) })
t.Run("raw api_key rejected", func(t *testing.T) { 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) { if !errors.Is(err, ErrRawAPIKeyNotAllowed) {
t.Fatalf("expected ErrRawAPIKeyNotAllowed, got %v", err) t.Fatalf("expected ErrRawAPIKeyNotAllowed, got %v", err)
} }
@@ -406,6 +351,614 @@ 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 TestProfileRepositoriesValidateDerivedDefinitions(t *testing.T) {
tests := []struct {
name string
files map[string]string
wantError error
wantBaseID string
wantProfile bool
}{
{
name: "alias is locally valid and normalizes base id",
files: map[string]string{"alias.yaml": `
id: selected-profile
base_profile: " base-profile "
`},
wantBaseID: "base-profile",
wantProfile: true,
},
{
name: "derived endpoint remains valid",
files: map[string]string{"invalid.yaml": `
id: selected-profile
base_profile: base-profile
endpoint: /v1
`},
wantError: ErrInvalidProfile,
},
{
name: "derived settings remain valid",
files: map[string]string{"invalid.yaml": `
id: selected-profile
base_profile: base-profile
top_p: 1.1
`},
wantError: ErrInvalidProfile,
},
{
name: "derived extra params remain valid",
files: map[string]string{"invalid.yaml": `
id: selected-profile
base_profile: base-profile
extra_params:
timestamp: 2026-08-11T12:34:56Z
`},
wantError: ErrInvalidProfile,
},
{
name: "derived raw key remains prohibited",
files: map[string]string{"invalid.yaml": `
id: selected-profile
base_profile: base-profile
api_key: secret
`},
wantError: ErrRawAPIKeyNotAllowed,
},
{
name: "derived duplicate id remains invalid",
files: map[string]string{
"first.yaml": "id: selected-profile\nbase_profile: first-base\n",
"second.yaml": "id: selected-profile\nbase_profile: second-base\n",
},
wantError: ErrInvalidProfile,
},
{
name: "derived extra document remains invalid",
files: map[string]string{"invalid.yaml": `
id: selected-profile
base_profile: base-profile
---
id: other
`},
wantError: ErrInvalidYAML,
},
{
name: "standalone profile remains complete",
files: map[string]string{"invalid.yaml": "id: selected-profile\n"},
wantError: ErrInvalidProfile,
},
}
for _, source := range profileRepositorySources() {
for _, tc := range tests {
t.Run(source.name+"/"+tc.name, func(t *testing.T) {
repo := source.newRepository(t, tc.files)
got, err := repo.GetProfile(context.Background(), "selected-profile")
if tc.wantError != nil {
if !errors.Is(err, tc.wantError) {
t.Fatalf("error = %v, want %v", err, tc.wantError)
}
return
}
if err != nil || !tc.wantProfile {
t.Fatalf("profile = %+v, error = %v, want valid derived definition", got, err)
}
if got.BaseProfileID != tc.wantBaseID {
t.Fatalf("BaseProfileID = %q, want %q", got.BaseProfileID, tc.wantBaseID)
}
})
}
}
}
func TestOverlayRepository(t *testing.T) { func TestOverlayRepository(t *testing.T) {
ctx := context.Background() ctx := context.Background()
primaryProfile := &domain.ExecutionProfile{ID: "shared", Endpoint: "http://primary", Model: "primary"} primaryProfile := &domain.ExecutionProfile{ID: "shared", Endpoint: "http://primary", Model: "primary"}
@@ -495,6 +1048,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 { func profileMapFile(content string) *fstest.MapFile {
return &fstest.MapFile{Data: []byte(strings.TrimLeft(content, "\n"))} return &fstest.MapFile{Data: []byte(strings.TrimLeft(content, "\n"))}
} }

View File

@@ -0,0 +1,160 @@
package profile
import (
"context"
"errors"
"fmt"
"strings"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
)
const maximumProfileChainLength = 32
type resolvingRepository struct {
source Repository
}
// NewResolvingRepository resolves inherited profile definitions from source.
func NewResolvingRepository(source Repository) Repository {
return &resolvingRepository{source: source}
}
func (r *resolvingRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
if r == nil || r.source == nil {
return nil, fmt.Errorf("%w: profile repository is required", ErrInvalidProfile)
}
requestedID := strings.TrimSpace(id)
if requestedID == "" {
return nil, fmt.Errorf("%w: profile id is required", ErrInvalidProfile)
}
if err := ctx.Err(); err != nil {
return nil, err
}
profile, err := r.getRawProfile(ctx, requestedID)
if err != nil {
return nil, err
}
if profile == nil {
return nil, fmt.Errorf("%w: selected profile %q is nil", ErrInvalidProfile, requestedID)
}
chain := []*domain.ExecutionProfile{profile}
chainIDs := []string{requestedID}
visited := map[string]struct{}{requestedID: {}}
current := profile
for {
baseID := strings.TrimSpace(current.BaseProfileID)
if baseID == "" {
break
}
if err := ctx.Err(); err != nil {
return nil, err
}
if _, seen := visited[baseID]; seen {
return nil, fmt.Errorf("%w: profile inheritance cycle %s", ErrInvalidProfile, joinProfileChain(chainIDs, baseID))
}
if len(chain) >= maximumProfileChainLength {
return nil, fmt.Errorf("%w: profile inheritance chain exceeds %d profiles: %s", ErrInvalidProfile, maximumProfileChainLength, joinProfileChain(chainIDs, baseID))
}
base, err := r.getRawProfile(ctx, baseID)
if err != nil {
if errors.Is(err, ErrProfileNotFound) {
return nil, fmt.Errorf("%w: base profile %q is missing in chain %s", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID))
}
return nil, fmt.Errorf("%w: failed to load base profile %q in chain %s: %w", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID), err)
}
if base == nil {
return nil, fmt.Errorf("%w: base profile %q is nil in chain %s", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID))
}
chain = append(chain, base)
chainIDs = append(chainIDs, baseID)
visited[baseID] = struct{}{}
current = base
}
resolved, err := mergeProfileChain(chain)
if err != nil {
return nil, fmt.Errorf("%w: resolved profile chain %s: %w", ErrInvalidProfile, strings.Join(chainIDs, " -> "), err)
}
if err := validateResolvedProfile(resolved); err != nil {
return nil, fmt.Errorf("%w: resolved profile chain %s: %w", ErrInvalidProfile, strings.Join(chainIDs, " -> "), err)
}
return resolved, nil
}
func (r *resolvingRepository) getRawProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
profile, err := r.source.GetProfile(ctx, id)
if err != nil {
return nil, err
}
if err := ctx.Err(); err != nil {
return nil, err
}
return profile, nil
}
func joinProfileChain(chain []string, next string) string {
return strings.Join(append(append([]string(nil), chain...), next), " -> ")
}
func mergeProfileChain(chain []*domain.ExecutionProfile) (*domain.ExecutionProfile, error) {
resolved := &domain.ExecutionProfile{ID: chain[0].ID}
for index := len(chain) - 1; index >= 0; index-- {
definition := chain[index]
if strings.TrimSpace(definition.BackendID) != "" {
resolved.BackendID = definition.BackendID
}
if strings.TrimSpace(definition.Endpoint) != "" {
resolved.Endpoint = definition.Endpoint
}
if strings.TrimSpace(definition.Model) != "" {
resolved.Model = definition.Model
}
if definition.Temperature != 0 {
resolved.Temperature = definition.Temperature
}
if definition.MaxTokens != 0 {
resolved.MaxTokens = definition.MaxTokens
}
if definition.TopP != 0 {
resolved.TopP = definition.TopP
}
if definition.TimeoutSeconds != 0 {
resolved.TimeoutSeconds = definition.TimeoutSeconds
}
if strings.TrimSpace(definition.ServiceTier) != "" {
resolved.ServiceTier = definition.ServiceTier
}
if strings.TrimSpace(definition.ReasoningEffort) != "" {
resolved.ReasoningEffort = definition.ReasoningEffort
}
if strings.TrimSpace(definition.APIKeyEnv) != "" {
resolved.APIKeyEnv = definition.APIKeyEnv
}
resolved.APIKeyRequired = resolved.APIKeyRequired || definition.APIKeyRequired
if len(definition.ExtraParams) != 0 {
extraParams, err := jsonvalue.CopyMap(definition.ExtraParams)
if err != nil {
return nil, err
}
resolved.ExtraParams = extraParams
}
}
resolved.ID = chain[0].ID
resolved.BaseProfileID = ""
return resolved, nil
}
func validateResolvedProfile(profile *domain.ExecutionProfile) error {
if profile == nil {
return errors.New("resolved profile is required")
}
profile.BaseProfileID = ""
return NormalizeAndValidateDefinition(profile)
}

View File

@@ -0,0 +1,418 @@
package profile
import (
"context"
"errors"
"fmt"
"io/fs"
"reflect"
"strings"
"sync"
"testing"
"testing/fstest"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
)
func TestResolvingRepositoryMergesProfileChain(t *testing.T) {
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"leaf": {
ID: "leaf",
BaseProfileID: "middle",
BackendID: "leaf-backend",
TopP: 0.8,
TimeoutSeconds: 45,
ReasoningEffort: "high",
},
"middle": {
ID: "middle",
BaseProfileID: "root",
Endpoint: "https://middle.example/v1",
Model: "middle-model",
MaxTokens: 256,
APIKeyEnv: "MIDDLE_API_KEY",
APIKeyRequired: true,
ExtraParams: map[string]any{"middle": map[string]any{"value": "middle"}},
},
"root": {
ID: "root",
BackendID: "root-backend",
Endpoint: "https://root.example/v1",
Model: "root-model",
Temperature: 0.3,
ServiceTier: "priority",
ExtraParams: map[string]any{"root": "value"},
},
}}
got, err := NewResolvingRepository(repo).GetProfile(context.Background(), "leaf")
if err != nil {
t.Fatalf("resolve profile: %v", err)
}
want := &domain.ExecutionProfile{
ID: "leaf",
BackendID: "leaf-backend",
Endpoint: "https://middle.example/v1",
Model: "middle-model",
Temperature: 0.3,
MaxTokens: 256,
TopP: 0.8,
TimeoutSeconds: 45,
ServiceTier: "priority",
ReasoningEffort: "high",
APIKeyEnv: "MIDDLE_API_KEY",
APIKeyRequired: true,
ExtraParams: map[string]any{"middle": map[string]any{"value": "middle"}},
}
if !reflect.DeepEqual(got, want) {
t.Fatalf("resolved profile:\n got %#v\nwant %#v", got, want)
}
}
func TestResolvingRepositoryRejectsMissingSourceAndProfileID(t *testing.T) {
if _, err := NewResolvingRepository(nil).GetProfile(context.Background(), "profile"); !errors.Is(err, ErrInvalidProfile) {
t.Fatalf("nil source error = %v, want ErrInvalidProfile", err)
}
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{}}
if _, err := NewResolvingRepository(repo).GetProfile(context.Background(), " \t "); !errors.Is(err, ErrInvalidProfile) {
t.Fatalf("blank id error = %v, want ErrInvalidProfile", err)
}
if got := repo.callCount(" "); got != 0 {
t.Fatalf("blank id looked up source %d times", got)
}
}
func TestResolvingRepositoryCopiesExtraParams(t *testing.T) {
baseParams := map[string]any{"nested": map[string]any{"value": "base"}}
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"child": {ID: "child", BaseProfileID: "base"},
"base": {
ID: "base",
Endpoint: "https://base.example/v1",
Model: "model",
ExtraParams: baseParams,
},
}}
resolver := NewResolvingRepository(repo)
first, err := resolver.GetProfile(context.Background(), "child")
if err != nil {
t.Fatalf("resolve inherited map: %v", err)
}
first.ExtraParams["nested"].(map[string]any)["value"] = "mutated"
second, err := resolver.GetProfile(context.Background(), "child")
if err != nil {
t.Fatalf("resolve inherited map again: %v", err)
}
if got := second.ExtraParams["nested"].(map[string]any)["value"]; got != "base" {
t.Fatalf("later result retained mutation: %v", got)
}
if got := baseParams["nested"].(map[string]any)["value"]; got != "base" {
t.Fatalf("source map retained mutation: %v", got)
}
repo.set("child", &domain.ExecutionProfile{
ID: "child",
BaseProfileID: "base",
ExtraParams: map[string]any{"child": "replacement"},
})
replaced, err := resolver.GetProfile(context.Background(), "child")
if err != nil {
t.Fatalf("resolve replacement map: %v", err)
}
if !reflect.DeepEqual(replaced.ExtraParams, map[string]any{"child": "replacement"}) {
t.Fatalf("extra params = %#v, want complete child replacement", replaced.ExtraParams)
}
}
func TestResolvingRepositoryUsesRawOverlayForEachLookup(t *testing.T) {
leafSource := NewFSRepository(profileTestFS(map[string]string{
"leaf.yaml": "id: leaf\nbase_profile: base\n",
}), ".")
fallback := NewFSRepository(profileTestFS(map[string]string{
"base.yaml": "id: base\nendpoint: https://fallback.example/v1\nmodel: fallback-model\n",
}), ".")
overlay := NewOverlayRepository(leafSource, fallback)
resolver := NewResolvingRepository(overlay)
got, err := resolver.GetProfile(context.Background(), "leaf")
if err != nil {
t.Fatalf("resolve fallback base: %v", err)
}
if got.Model != "fallback-model" {
t.Fatalf("fallback base model = %q", got.Model)
}
shadowing := NewOverlayRepository(NewFSRepository(profileTestFS(map[string]string{
"leaf.yaml": "id: leaf\nbase_profile: base\n",
"base.yaml": "id: base\nendpoint: https://primary.example/v1\nmodel: primary-model\n",
}), "."), fallback)
got, err = NewResolvingRepository(shadowing).GetProfile(context.Background(), "leaf")
if err != nil {
t.Fatalf("resolve shadowed base: %v", err)
}
if got.Model != "primary-model" || got.Endpoint != "https://primary.example/v1" {
t.Fatalf("shadowed base = %+v", got)
}
}
func TestResolvingRepositoryReportsSafetyAndSourceErrors(t *testing.T) {
sourceErr := errors.New("source failure")
tests := []struct {
name string
repo *resolvingTestRepository
id string
want []error
wantNot error
contains []string
}{
{
name: "missing selected profile preserves not found",
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{}},
id: "missing",
want: []error{ErrProfileNotFound},
wantNot: ErrInvalidProfile,
},
{
name: "missing base is invalid but not not found",
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"leaf": {ID: "leaf", BaseProfileID: "missing"},
}},
id: "leaf",
want: []error{ErrInvalidProfile},
wantNot: ErrProfileNotFound,
contains: []string{"missing", "leaf -> missing"},
},
{
name: "direct cycle",
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"a": {ID: "a", BaseProfileID: "a"},
}},
id: "a",
want: []error{ErrInvalidProfile},
contains: []string{"a -> a"},
},
{
name: "indirect cycle",
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"a": {ID: "a", BaseProfileID: "b"},
"b": {ID: "b", BaseProfileID: "c"},
"c": {ID: "c", BaseProfileID: "a"},
}},
id: "a",
want: []error{ErrInvalidProfile},
contains: []string{"a -> b -> c -> a"},
},
{
name: "nil result",
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"leaf": nil,
}},
id: "leaf",
want: []error{ErrInvalidProfile},
},
{
name: "incomplete resolved profile",
repo: &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"leaf": {ID: "leaf", BaseProfileID: "base"},
"base": {ID: "base", Model: "model"},
}},
id: "leaf",
want: []error{ErrInvalidProfile},
},
{
name: "base source error is retained",
repo: &resolvingTestRepository{
profiles: map[string]*domain.ExecutionProfile{"leaf": {ID: "leaf", BaseProfileID: "base"}},
errors: map[string]error{"base": sourceErr},
},
id: "leaf",
want: []error{ErrInvalidProfile, sourceErr},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
_, err := NewResolvingRepository(tc.repo).GetProfile(context.Background(), tc.id)
for _, want := range tc.want {
if !errors.Is(err, want) {
t.Fatalf("error = %v, want %v", err, want)
}
}
if tc.wantNot != nil && errors.Is(err, tc.wantNot) {
t.Fatalf("error = %v, must not match %v", err, tc.wantNot)
}
for _, fragment := range tc.contains {
if !strings.Contains(err.Error(), fragment) {
t.Fatalf("error = %v, want %q", err, fragment)
}
}
})
}
}
func TestResolvingRepositoryEnforcesChainLength(t *testing.T) {
for _, count := range []int{maximumProfileChainLength, maximumProfileChainLength + 1} {
t.Run(fmt.Sprintf("%d profiles", count), func(t *testing.T) {
profiles := make(map[string]*domain.ExecutionProfile, count)
for index := 1; index <= count; index++ {
id := fmt.Sprintf("profile-%d", index)
definition := &domain.ExecutionProfile{ID: id}
if index == count {
definition.Endpoint = "https://root.example/v1"
definition.Model = "model"
} else {
definition.BaseProfileID = fmt.Sprintf("profile-%d", index+1)
}
profiles[id] = definition
}
got, err := NewResolvingRepository(&resolvingTestRepository{profiles: profiles}).GetProfile(context.Background(), "profile-1")
if count == maximumProfileChainLength {
if err != nil || got == nil {
t.Fatalf("profile = %+v, error = %v, want accepted chain", got, err)
}
return
}
if !errors.Is(err, ErrInvalidProfile) {
t.Fatalf("error = %v, want ErrInvalidProfile", err)
}
})
}
}
func TestResolvingRepositoryIsFreshAndCancellationAware(t *testing.T) {
repo := &resolvingTestRepository{profiles: map[string]*domain.ExecutionProfile{
"leaf": {ID: "leaf", BaseProfileID: "base"},
"base": {ID: "base", Endpoint: "https://base.example/v1", Model: "first", ExtraParams: map[string]any{"nested": map[string]any{"value": "first"}}},
}}
resolver := NewResolvingRepository(repo)
first, err := resolver.GetProfile(context.Background(), "leaf")
if err != nil || first.Model != "first" {
t.Fatalf("first result=(%+v, %v)", first, err)
}
repo.set("base", &domain.ExecutionProfile{ID: "base", Endpoint: "https://base.example/v1", Model: "second", ExtraParams: map[string]any{"nested": map[string]any{"value": "second"}}})
second, err := resolver.GetProfile(context.Background(), "leaf")
if err != nil || second.Model != "second" {
t.Fatalf("second result=(%+v, %v)", second, err)
}
canceled, cancel := context.WithCancel(context.Background())
cancel()
if _, err := resolver.GetProfile(canceled, "leaf"); !errors.Is(err, context.Canceled) {
t.Fatalf("canceled lookup error = %v", err)
}
if got := repo.callCount("leaf"); got != 2 {
t.Fatalf("calls after canceled lookup = %d, want 2", got)
}
duringTraversal, cancelDuringTraversal := context.WithCancel(context.Background())
repo.afterGet = func(id string) {
if id == "leaf" {
cancelDuringTraversal()
}
}
if _, err := resolver.GetProfile(duringTraversal, "leaf"); !errors.Is(err, context.Canceled) {
t.Fatalf("during traversal error = %v", err)
}
if got := repo.callCount("base"); got != 2 {
t.Fatalf("base calls after cancellation = %d, want 2", got)
}
terminalLookup, cancelTerminalLookup := context.WithCancel(context.Background())
repo.afterGet = func(id string) {
if id == "base" {
cancelTerminalLookup()
}
}
if _, err := resolver.GetProfile(terminalLookup, "leaf"); !errors.Is(err, context.Canceled) {
t.Fatalf("terminal lookup cancellation error = %v", err)
}
if got := repo.callCount("base"); got != 3 {
t.Fatalf("base calls after terminal cancellation = %d, want 3", got)
}
repo.afterGet = nil
var wg sync.WaitGroup
errors := make(chan error, 8)
for index := 0; index < cap(errors); index++ {
wg.Add(1)
go func() {
defer wg.Done()
resolved, err := resolver.GetProfile(context.Background(), "leaf")
if err != nil {
errors <- err
return
}
resolved.ExtraParams["nested"].(map[string]any)["value"] = "mutated"
}()
}
wg.Wait()
close(errors)
for err := range errors {
t.Errorf("concurrent resolution: %v", err)
}
latest, err := resolver.GetProfile(context.Background(), "leaf")
if err != nil || latest.ExtraParams["nested"].(map[string]any)["value"] != "second" {
t.Fatalf("latest result=(%+v, %v)", latest, err)
}
}
type resolvingTestRepository struct {
mu sync.Mutex
profiles map[string]*domain.ExecutionProfile
errors map[string]error
calls map[string]int
afterGet func(string)
}
func (r *resolvingTestRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
r.mu.Lock()
if r.calls == nil {
r.calls = make(map[string]int)
}
r.calls[id]++
err := r.errors[id]
profile := r.profiles[id]
afterGet := r.afterGet
r.mu.Unlock()
if afterGet != nil {
afterGet(id)
}
if err != nil {
return nil, err
}
if profile == nil {
if _, exists := r.profiles[id]; exists {
return nil, nil
}
return nil, ErrProfileNotFound
}
copy := *profile
return &copy, nil
}
func (r *resolvingTestRepository) set(id string, profile *domain.ExecutionProfile) {
r.mu.Lock()
defer r.mu.Unlock()
r.profiles[id] = profile
}
func (r *resolvingTestRepository) callCount(id string) int {
r.mu.Lock()
defer r.mu.Unlock()
return r.calls[id]
}
func profileTestFS(files map[string]string) fs.FS {
fsys := make(fstest.MapFS, len(files))
for name, content := range files {
fsys[name] = profileMapFile(content)
}
return fsys
}

View File

@@ -5,6 +5,7 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"strings"
"text/template" "text/template"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/domain"
@@ -18,6 +19,8 @@ var (
ErrInvalidMessageRole = errors.New("invalid or empty message role") ErrInvalidMessageRole = errors.New("invalid or empty message role")
) )
const artifactTextChunkSize = 64 * 1024
type goRenderer struct{} type goRenderer struct{}
func NewGoRenderer() Renderer { 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) { 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 { if definition == nil {
return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure) return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure)
} }
// 1. Verify required inputs
for _, in := range definition.Inputs { for _, in := range definition.Inputs {
if !in.Required { if !in.Required {
continue 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{ funcs := template.FuncMap{
"input": func(name string) (string, error) { "input": resolver.resolve,
art, ok := inputs[name]
if !ok || art == nil {
return "", fmt.Errorf("%w: %s", ErrUnknownInput, name)
}
return string(art.Body), nil
},
} }
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 { if err != nil {
return nil, err 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 { for i, tmplMsg := range definition.Templates {
select { if err := ctx.Err(); err != nil {
case <-ctx.Done(): return nil, err
return nil, ctx.Err()
default:
} }
if tmplMsg.Role == "" { if tmplMsg.Role == "" {
return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i) return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i)
} }
// Parse and execute template if err := ctx.Err(); err != nil {
tmpl, err := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content) return nil, err
if err != nil { }
return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, 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 var buf bytes.Buffer
if err := tmpl.Execute(&buf, vars); err != nil { executeErr := tmpl.Execute(&buf, vars)
return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, err) 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{ renderedMessages = append(renderedMessages, domain.RenderedMessage{
@@ -85,6 +100,13 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
Content: buf.String(), Content: buf.String(),
CacheControl: cloneCacheControl(tmplMsg.CacheControl), 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{ return &domain.RenderedPrompt{
@@ -93,15 +115,75 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
}, nil }, nil
} }
func renderSessionID(raw string, funcs template.FuncMap, vars map[string]string) (string, error) { type artifactTextResolver struct {
tmpl, err := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw) ctx context.Context
if err != nil { inputs map[string]*domain.Artifact
return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, err) 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 var buf bytes.Buffer
if err := tmpl.Execute(&buf, vars); err != nil { executeErr := tmpl.Execute(&buf, vars)
return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, err) 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()) sessionID, err := domain.NormalizeSessionID(buf.String())

View File

@@ -1,6 +1,7 @@
package prompt package prompt
import ( import (
"bytes"
"context" "context"
"errors" "errors"
"strings" "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) { t.Run("inserting required input artifact", func(t *testing.T) {
def := &domain.PromptDefinition{ def := &domain.PromptDefinition{
Inputs: []domain.PromptInput{{Name: "transcript", Required: true}}, 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()
}

View 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))
}

View File

@@ -5,14 +5,11 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"io"
"io/fs" "io/fs"
"os"
"path"
"path/filepath"
"strings" "strings"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/domain"
"gitea.maximumdirect.net/eric/promptkit/internal/filecatalog"
"gopkg.in/yaml.v3" "gopkg.in/yaml.v3"
) )
@@ -22,13 +19,8 @@ var (
ErrInvalidPromptDefinition = errors.New("invalid prompt definition configuration") ErrInvalidPromptDefinition = errors.New("invalid prompt definition configuration")
) )
type filesystemRepository struct { type sourceRepository struct {
dir string source promptDefinitionSource
}
type fsRepository struct {
fsys fs.FS
root string
} }
type promptDefinitionFile struct { type promptDefinitionFile struct {
@@ -69,19 +61,47 @@ type promptOutputContractFile struct {
} }
func NewFilesystemRepository(dir string) Repository { 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 { 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) == "" { if strings.TrimSpace(id) == "" {
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidPromptDefinition) 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 { if err != nil {
return nil, fmt.Errorf("failed to read prompt definition directory: %w", err) 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: default:
} }
relPath := filecatalog.RelativePath(r.dir, fullPath) relPath := r.source.displayPath(fullPath)
fileMatch := filecatalog.Stem(filepath.Base(fullPath)) == id data, err := r.source.readDefinition(fullPath)
raw, err := loadPromptDefinitionFile(fullPath)
if err != nil { 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) return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, err)
} }
continue continue
} }
if !promptDefinitionMatches(raw, id, version) {
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 {
continue continue
} }
matches = append(matches, promptDefinitionMatch{ matches = append(matches, promptDefinitionMatch{
def: def, raw: raw,
path: relPath, sourcePath: fullPath,
path: relPath,
}) })
} }
@@ -137,130 +148,21 @@ func (r *filesystemRepository) GetPromptDefinition(ctx context.Context, id strin
} }
if len(matches) == 1 { 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 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 { type promptDefinitionMatch struct {
def *domain.PromptDefinition raw *promptDefinitionFile
path string sourcePath string
} 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
} }
func decodePromptDefinition(data []byte) (*promptDefinitionFile, error) { func decodePromptDefinition(data []byte) (*promptDefinitionFile, error) {
@@ -270,59 +172,44 @@ func decodePromptDefinition(data []byte) (*promptDefinitionFile, error) {
if err := decoder.Decode(&raw); err != nil { if err := decoder.Decode(&raw); err != nil {
return nil, err 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 return &raw, nil
} }
func promptDefinitionDataHasID(data []byte, id string) bool { func promptDefinitionDataMatches(data []byte, id string, version string) bool {
var raw struct { 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 { if err := yaml.NewDecoder(bytes.NewReader(data)).Decode(&raw); err != nil {
return false 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) { func promptDefinitionMatches(raw *promptDefinitionFile, id string, version string) bool {
promptDir := filepath.Dir(sourcePath) if raw == nil {
return normalizePromptDefinitionWithContent(raw, func(contentFile string) (string, string, error) { return false
resolvedPath := strings.TrimSpace(contentFile) }
if !filepath.IsAbs(resolvedPath) { return promptSelectorMatches(raw.ID, raw.Version, id, version)
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 normalizePromptDefinitionFromFS(raw *promptDefinitionFile, fsys fs.FS, root string, sourcePath string, rootIsDir bool) (*domain.PromptDefinition, error) { func promptSelectorMatches(rawID string, rawVersion string, id string, version string) bool {
promptDir := path.Dir(sourcePath) if strings.TrimSpace(rawID) != id {
return normalizePromptDefinitionWithContent(raw, func(contentFile string) (string, string, error) { return false
var resolvedPath string }
if rootIsDir { return version == "" || strings.TrimSpace(rawVersion) == version
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), "/")
}
body, err := fs.ReadFile(fsys, resolvedPath) func normalizePromptDefinition(raw *promptDefinitionFile, sourceRoot contentSourceRoot, sourcePath string) (*domain.PromptDefinition, error) {
if err != nil { return normalizePromptDefinitionWithContent(raw, func(contentFile string) (string, string, error) {
return "", "", err return sourceRoot.readContentFile(sourcePath, contentFile)
}
return string(body), resolvedPath, nil
}) })
} }
@@ -402,17 +289,14 @@ func normalizePromptDefinitionWithContent(raw *promptDefinitionFile, readContent
}) })
} }
if !isValidOutputFormat(raw.Output.Format) { outputContract := domain.OutputContract{
return nil, fmt.Errorf("invalid output format: %q", raw.Output.Format) Format: raw.Output.Format,
ValidationMode: raw.Output.ValidationMode,
SchemaPath: strings.TrimSpace(raw.Output.SchemaPath),
RepairAttempts: raw.Output.RepairAttempts,
} }
if !isValidValidationMode(raw.Output.ValidationMode) { if err := domain.ValidateOutputContract(outputContract); err != nil {
return nil, fmt.Errorf("invalid validation mode: %q", raw.Output.ValidationMode) return nil, fmt.Errorf("output: %w", err)
}
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")
} }
defaultProfile := "" defaultProfile := ""
@@ -432,12 +316,7 @@ func normalizePromptDefinitionWithContent(raw *promptDefinitionFile, readContent
Inputs: inputs, Inputs: inputs,
Templates: templates, Templates: templates,
OutputFormat: raw.Output.Format, OutputFormat: raw.Output.Format,
Validation: domain.OutputContract{ Validation: outputContract,
Format: raw.Output.Format,
ValidationMode: raw.Output.ValidationMode,
SchemaPath: strings.TrimSpace(raw.Output.SchemaPath),
RepairAttempts: raw.Output.RepairAttempts,
},
}, nil }, nil
} }
@@ -464,21 +343,3 @@ func normalizeCacheControl(raw *cacheControlFile) (*domain.CacheControl, error)
TTL: ttl, TTL: ttl,
}, nil }, 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
}
}

View File

@@ -3,23 +3,51 @@ package promptdef
import ( import (
"context" "context"
"errors" "errors"
"fmt"
"io/fs" "io/fs"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"sync"
"testing" "testing"
"testing/fstest" "testing/fstest"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "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() tmpDir := t.TempDir()
if err := copyTree("testdata", tmpDir); err != nil { if err := copyTree("testdata", tmpDir); err != nil {
t.Fatalf("failed to copy testdata: %v", err) t.Fatalf("failed to copy testdata: %v", err)
} }
repo := NewFilesystemRepository(tmpDir) repo := newRepository(tmpDir)
ctx := context.Background() ctx := context.Background()
t.Run("valid inline prompt", func(t *testing.T) { t.Run("valid inline prompt", func(t *testing.T) {
@@ -64,8 +92,8 @@ func TestFilesystemRepository_GetPromptDefinition(t *testing.T) {
if p.Templates[1].ContentFile == "" { if p.Templates[1].ContentFile == "" {
t.Fatal("expected ContentFile source metadata to be preserved") t.Fatal("expected ContentFile source metadata to be preserved")
} }
if !filepath.IsAbs(p.Templates[1].ContentFile) { if filepath.IsAbs(p.Templates[1].ContentFile) != contentPathsAreFull {
t.Fatalf("expected resolved content_file path to be absolute, got %q", p.Templates[1].ContentFile) t.Fatalf("unexpected content_file path representation: %q", p.Templates[1].ContentFile)
} }
}) })
@@ -287,20 +315,20 @@ output:
targetErr error targetErr error
errSubstrs []string errSubstrs []string
}{ }{
{name: "invalid YAML", id: "invalid_yaml", targetErr: ErrInvalidYAML}, {name: "unidentifiable invalid YAML is unrelated", id: "invalid-yaml", targetErr: ErrPromptDefinitionNotFound},
{name: "missing id", id: "missing_id", targetErr: ErrInvalidPromptDefinition, errSubstrs: []string{"id is required"}}, {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: "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: "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: "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: "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: "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: "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: "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: "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: "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 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: "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: "unknown cache control field", id: "unknown-cache-control-field", targetErr: ErrInvalidYAML, errSubstrs: []string{"field unexpected not found"}},
} }
for _, tc := range cases { for _, tc := range cases {
@@ -396,7 +424,7 @@ output:
for _, tc := range tests { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
repo := NewFSRepository(fstest.MapFS{ fsys := &recordingFS{FS: fstest.MapFS{
"prompts/prompt.yaml": &fstest.MapFile{Data: []byte(` "prompts/prompt.yaml": &fstest.MapFile{Data: []byte(`
id: fs-escaped-prompt id: fs-escaped-prompt
version: "1.0.0" version: "1.0.0"
@@ -409,7 +437,8 @@ output:
repair_attempts: 0 repair_attempts: 0
`)}, `)},
"outside.tmpl": &fstest.MapFile{Data: []byte(`Outside root.`)}, "outside.tmpl": &fstest.MapFile{Data: []byte(`Outside root.`)},
}, "prompts") }}
repo := NewFSRepository(fsys, "prompts")
_, err := repo.GetPromptDefinition(context.Background(), "fs-escaped-prompt", "") _, err := repo.GetPromptDefinition(context.Background(), "fs-escaped-prompt", "")
if !errors.Is(err, ErrInvalidPromptDefinition) { if !errors.Is(err, ErrInvalidPromptDefinition) {
@@ -418,64 +447,611 @@ output:
if !strings.Contains(err.Error(), tc.wantErr) { if !strings.Contains(err.Error(), tc.wantErr) {
t.Fatalf("expected error to contain %q, got %v", tc.wantErr, err) 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) { type recordingFS struct {
repo := NewFSRepository(fstest.MapFS{ fs.FS
"one.yaml": &fstest.MapFile{Data: []byte(` mu sync.Mutex
id: duplicate-fs-prompt opened []string
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
`)},
}, ".")
_, err := repo.GetPromptDefinition(context.Background(), "duplicate-fs-prompt", "") func TestPromptRepositoryReturnsDefinitionReadFailures(t *testing.T) {
if !errors.Is(err, ErrInvalidPromptDefinition) { readErr := errors.New("definition read failed")
t.Fatalf("expected ErrInvalidPromptDefinition, got %v", err) 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") { repo := NewFSRepository(fsys, "prompts")
t.Fatalf("expected duplicate paths in error, got %v", err)
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) { type definitionReadFailureFS struct {
repo := NewFSRepository(fstest.MapFS{ fs.FS
"not_named_like_id.yaml": &fstest.MapFile{Data: []byte(` target string
id: strict-fs-prompt err error
version: "1.0.0" }
unknown: true
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: messages:
- role: user - role: user
content: Invalid. content: selected
output: output:
format: text format: text
validation_mode: none 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", "") for _, source := range promptRepositorySources() {
if !errors.Is(err, ErrInvalidYAML) { for _, tc := range tests {
t.Fatalf("expected ErrInvalidYAML, got %v", err) 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)
}
}
})
} }
} }

View 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)
}

View 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
}

View 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)
}
}

View 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,
}
}

View 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,
)
}
}

View 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)
}

View File

@@ -0,0 +1,632 @@
package usecase
import (
"context"
"errors"
"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 TestRunnerRunPreparedCredentialAvailabilityBeforeAdmission(t *testing.T) {
const environmentName = "PROMPTKIT_PREPARED_EXECUTION_TEST_KEY"
tests := []struct {
name string
apiKeyRequired bool
profileEnv bool
overrideEnv bool
wantFailure bool
}{
{name: "optional environment becomes unavailable", profileEnv: true},
{
name: "required request environment becomes unavailable",
apiKeyRequired: true,
overrideEnv: true,
wantFailure: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Setenv(environmentName, "available-during-preparation")
profile := defaultExecutionProfile()
profile.APIKeyRequired = tc.apiKeyRequired
if tc.profileEnv {
profile.APIKeyEnv = environmentName
}
validator := &recordingValidationPreparer{plan: &recordingPreparedValidation{}}
admitter := &fakeRunAdmitter{}
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}
runner := NewRunner(
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": profile}},
nil,
defaultArtifactReader(),
defaultRenderer(),
llmClient,
validator,
admitter,
)
request := domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}
if tc.overrideEnv {
request.Execution = &domain.ExecutionTargetOverride{APIKeyEnv: environmentName}
}
prepared, err := runner.PrepareExecution(context.Background(), request)
if err != nil {
t.Fatalf("prepare execution: %v", err)
}
t.Setenv(environmentName, "")
result, err := runner.RunPrepared(context.Background(), prepared)
if tc.wantFailure {
if result != nil || !errors.Is(err, ErrInvalidRequest) || !errors.Is(err, ErrAPIKeyEnvMissing) {
t.Fatalf("required credential result = (%+v, %v)", result, err)
}
if len(admitter.backendIDs) != 0 || llmClient.calls != 0 {
t.Fatalf("required credential reached admission or generation: admission=%v generation=%d", admitter.backendIDs, llmClient.calls)
}
} else {
if result == nil || err != nil {
t.Fatalf("optional credential result = (%+v, %v), want success", result, err)
}
if len(admitter.backendIDs) != 1 || llmClient.calls != 1 {
t.Fatalf("optional credential admission=%v generation=%d, want one each", admitter.backendIDs, llmClient.calls)
}
}
if _, err := runner.RunPrepared(context.Background(), prepared); !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("execution outcome did not consume handle: %v", err)
}
})
}
}
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)
}
}

View 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
}

View 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)
}
})
}
})
}

View 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
}

View 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)
}
}

View File

@@ -19,6 +19,7 @@ type RepairRequest struct {
ValidationErrors []string ValidationErrors []string
SessionID string SessionID string
Target domain.ExecutionTarget Target domain.ExecutionTarget
TargetPresence domain.ExecutionTargetPresence
StructuredOutput *domain.StructuredOutputSpec StructuredOutput *domain.StructuredOutputSpec
Attempt int Attempt int
MaxAttempts int MaxAttempts int
@@ -44,7 +45,6 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
} }
prompt := domain.RenderedPrompt{ prompt := domain.RenderedPrompt{
SessionID: req.SessionID,
Messages: []domain.RenderedMessage{ Messages: []domain.RenderedMessage{
{ {
Role: "system", Role: "system",
@@ -64,11 +64,13 @@ func (r *defaultOutputRepairer) Repair(ctx context.Context, req RepairRequest) (
}, },
} }
resp, err := r.llm.Generate(ctx, domain.GenerateRequest{ resp, err := r.llm.Generate(ctx, newGenerationRequest(
Prompt: prompt, prompt,
Target: req.Target, req.SessionID,
StructuredOutput: req.StructuredOutput, req.Target,
}) req.TargetPresence,
req.StructuredOutput,
))
if err != nil { if err != nil {
return nil, err return nil, err
} }

View File

@@ -17,6 +17,7 @@ import (
"gitea.maximumdirect.net/eric/promptkit/internal/capacity" "gitea.maximumdirect.net/eric/promptkit/internal/capacity"
"gitea.maximumdirect.net/eric/promptkit/internal/defaults" "gitea.maximumdirect.net/eric/promptkit/internal/defaults"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "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/llm"
"gitea.maximumdirect.net/eric/promptkit/internal/profile" "gitea.maximumdirect.net/eric/promptkit/internal/profile"
"gitea.maximumdirect.net/eric/promptkit/internal/prompt" "gitea.maximumdirect.net/eric/promptkit/internal/prompt"
@@ -71,6 +72,11 @@ type preparationState struct {
start time.Time start time.Time
} }
type preparedOperation struct {
run *domain.PreparedRun
validation validate.PreparedValidation
}
func NewRunner( func NewRunner(
promptDefs promptdef.Repository, promptDefs promptdef.Repository,
profiles profile.Repository, profiles profile.Repository,
@@ -131,41 +137,62 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
return nil, err return nil, err
} }
if r.admitter != nil { release, err := r.admitRun(ctx, state.effectiveModel.BackendID)
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)
if err != nil { if err != nil {
return nil, err return nil, err
} }
defer release()
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{ operation, err := r.completePreparation(ctx, req, state)
Prompt: domain.RenderedPrompt{SessionID: prepared.SessionID, Messages: prepared.Messages}, if err != nil {
Target: prepared.EffectiveModelParams, return nil, err
TargetPresence: prepared.TargetPresence, }
StructuredOutput: prepared.StructuredOutput, 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 err != nil {
if errors.Is(err, llm.ErrInvalidRequest) { if errors.Is(err, llm.ErrInvalidRequest) {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err) return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
} }
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err) return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
} }
usage := genResp.Usage
outputArtifact := buildOutputArtifact(genResp.Content, prepared.OutputContract.Format) 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 { if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err) 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, PreviousOutput: genResp.Content,
ValidationErrors: validationResult.Errors, ValidationErrors: validationResult.Errors,
SessionID: prepared.SessionID, SessionID: prepared.SessionID,
Target: prepared.EffectiveModelParams, Target: executionTarget,
TargetPresence: prepared.TargetPresence,
StructuredOutput: prepared.StructuredOutput, StructuredOutput: prepared.StructuredOutput,
Attempt: attemptsUsed, Attempt: attemptsUsed,
MaxAttempts: prepared.OutputContract.RepairAttempts, MaxAttempts: prepared.OutputContract.RepairAttempts,
@@ -193,9 +221,10 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
} }
genResp = repairResp genResp = repairResp
usage = addTokenUsage(usage, repairResp.Usage)
outputArtifact = buildOutputArtifact(genResp.Content, prepared.OutputContract.Format) 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 { if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err) 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() end := time.Now().UTC()
executionTarget.APIKey = ""
return &domain.RunResult{ return &domain.RunResult{
RunID: runID, RunID: runID,
@@ -218,21 +248,35 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
SelectedBackendID: prepared.SelectedBackendID, SelectedBackendID: prepared.SelectedBackendID,
ModelName: prepared.EffectiveModelParams.Model, ModelName: prepared.EffectiveModelParams.Model,
Endpoint: prepared.EffectiveModelParams.Endpoint, Endpoint: prepared.EffectiveModelParams.Endpoint,
EffectiveModelParams: prepared.EffectiveModelParams, EffectiveModelParams: executionTarget,
InputHashes: prepared.InputHashes, InputHashes: prepared.InputHashes,
Usage: genResp.Usage, Usage: usage,
StartTime: start, StartTime: start,
EndTime: end, EndTime: end,
Duration: end.Sub(start), Duration: end.Sub(start),
}, nil }, 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) { func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.PreparedRun, error) {
state, err := r.resolvePreparation(ctx, req, time.Now().UTC()) state, err := r.resolvePreparation(ctx, req, time.Now().UTC())
if err != nil { if err != nil {
return nil, err 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( func (r *Runner) resolvePreparation(
@@ -248,13 +292,15 @@ func (r *Runner) resolvePreparation(
return nil, fmt.Errorf("%w: session_id: %v", ErrInvalidRequest, err) 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 { 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 { 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) 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) 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 { if err != nil {
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err) return nil, err
} }
var selectedBackend *domain.Backend effectiveModel, targetPresence := resolveExecutionTarget(selection.backend, selection.profile, req.Execution)
if backendID := strings.TrimSpace(execProfile.BackendID); backendID != "" { effectiveModel.APIKey = req.APIKey
execProfile.BackendID = backendID effectiveModel, err = normalizeResolvedExecutionTarget(effectiveModel)
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)
if err != nil { if err != nil {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err) 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 { if err := validateAPIKey(effectiveModel.APIKeyEnv, effectiveModel.APIKey, effectiveModel.APIKeyRequired); err != nil {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err) return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
} }
effectiveContract := resolveOutputContract(def, req.Validation)
return &preparationState{ return &preparationState{
definition: def, definition: def,
directSessionID: directSessionID, directSessionID: directSessionID,
promptDefinitionHash: promptDefinitionHash, promptDefinitionHash: promptDefinitionHash,
selectedProfileID: selectedProfileID, selectedProfileID: selection.id,
effectiveModel: effectiveModel, effectiveModel: effectiveModel,
targetPresence: targetPresence, targetPresence: targetPresence,
effectiveContract: effectiveContract, effectiveContract: effectiveContract,
@@ -315,16 +342,76 @@ func (r *Runner) completePreparation(
ctx context.Context, ctx context.Context,
req domain.RunRequest, req domain.RunRequest,
state *preparationState, state *preparationState,
) (*domain.PreparedRun, error) { ) (*preparedOperation, error) {
structuredOutput, err := r.resolveStructuredOutput( validationPlan, err := r.prepareValidation(ctx, state.effectiveContract)
ctx, if err != nil {
return nil, err
}
structuredOutput, err := r.structuredOutputFromValidationPlan(
state.definition, state.definition,
state.effectiveContract, state.effectiveContract,
validationPlan,
) )
if err != nil { if err != nil {
return nil, err 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)) resolvedInputs := make(map[string]*domain.Artifact, len(req.Inputs))
inputHashes := make(map[string]string, len(req.Inputs)) inputHashes := make(map[string]string, len(req.Inputs))
for name, ref := range req.Inputs { for name, ref := range req.Inputs {
@@ -354,13 +441,15 @@ func (r *Runner) completePreparation(
} }
end := time.Now().UTC() end := time.Now().UTC()
effectiveModel := state.effectiveModel
effectiveModel.APIKey = ""
return &domain.PreparedRun{ return &domain.PreparedRun{
PromptID: state.definition.ID, PromptID: state.definition.ID,
PromptVersion: state.definition.Version, PromptVersion: state.definition.Version,
PromptHash: state.promptDefinitionHash, PromptHash: state.promptDefinitionHash,
SelectedProfileID: state.selectedProfileID, SelectedProfileID: state.selectedProfileID,
SelectedBackendID: state.effectiveModel.BackendID, SelectedBackendID: state.effectiveModel.BackendID,
EffectiveModelParams: state.effectiveModel, EffectiveModelParams: effectiveModel,
TargetPresence: state.targetPresence, TargetPresence: state.targetPresence,
OutputContract: state.effectiveContract, OutputContract: state.effectiveContract,
StructuredOutput: structuredOutput, StructuredOutput: structuredOutput,
@@ -374,29 +463,29 @@ func (r *Runner) completePreparation(
}, nil }, nil
} }
func (r *Runner) resolveStructuredOutput(ctx context.Context, def *domain.PromptDefinition, contract domain.OutputContract) (*domain.StructuredOutputSpec, error) { func (r *Runner) admitRun(ctx context.Context, backendID string) (func(), error) {
if contract.ValidationMode != domain.ValidationJSONSchema { if r.admitter == nil {
return nil, nil return func() {}, nil
} }
release, err := r.admitter.Admit(ctx, backendID)
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)
if err != nil { 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{ return &domain.StructuredOutputSpec{
Type: domain.StructuredOutputJSONSchema, Type: domain.StructuredOutputJSONSchema,
JSONSchema: &domain.StructuredOutputJSONSpec{ JSONSchema: &domain.StructuredOutputJSONSpec{
Name: deriveStructuredSchemaName(def.ID, def.Version), Name: deriveStructuredSchemaName(def.ID, def.Version),
Strict: true, Strict: true,
Schema: schemaDoc, Schema: schemaDocument,
}, },
}, nil }
} }
func deriveStructuredSchemaName(promptID string, promptVersion string) string { func deriveStructuredSchemaName(promptID string, promptVersion string) string {
@@ -425,25 +514,6 @@ func deriveStructuredSchemaName(promptID string, promptVersion string) string {
return name 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 { func (r *Runner) shouldAttemptRepair(contract domain.OutputContract, validationResult domain.ValidationResult) bool {
if r.repairer == nil { if r.repairer == nil {
return false return false
@@ -499,7 +569,7 @@ func mergeExecutionTarget(base domain.ExecutionTarget, override domain.Execution
return out 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 out := base
var presence domain.ExecutionTargetPresence var presence domain.ExecutionTargetPresence
if override.Endpoint != "" { if override.Endpoint != "" {
@@ -509,30 +579,18 @@ func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.E
out.Model = override.Model out.Model = override.Model
} }
if override.Temperature != nil { 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 out.Temperature = *override.Temperature
presence.Temperature = true presence.Temperature = true
} }
if override.MaxTokens != nil { 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 out.MaxTokens = *override.MaxTokens
presence.MaxTokens = true presence.MaxTokens = true
} }
if override.TopP != nil { 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 out.TopP = *override.TopP
presence.TopP = true presence.TopP = true
} }
if override.TimeoutSeconds != nil { 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 out.TimeoutSeconds = *override.TimeoutSeconds
presence.TimeoutSeconds = true presence.TimeoutSeconds = true
} }
@@ -548,34 +606,30 @@ func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.E
if len(override.ExtraParams) > 0 { if len(override.ExtraParams) > 0 {
out.ExtraParams = copyExtraParams(override.ExtraParams) 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 := defaults.ExecutionTargetDefault()
out = mergeExecutionTarget(out, backendToTarget(backendValue)) out = mergeExecutionTarget(out, backendToTarget(backendValue))
out = mergeExecutionTarget(out, executionProfileToTarget(profileValue)) out = mergeExecutionTarget(out, executionProfileToTarget(profileValue))
var presence domain.ExecutionTargetPresence var presence domain.ExecutionTargetPresence
if override != nil { if override != nil {
var err error out, presence = mergeExecutionTargetOverride(out, *override)
out, presence, err = mergeExecutionTargetOverride(out, *override)
if err != nil {
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, err
}
} }
return out, presence, nil return out, presence
} }
func validateAPIKey(apiKeyEnv string, apiKey string, apiKeyRequired bool) error { func validateAPIKey(apiKeyEnv string, apiKey string, apiKeyRequired bool) error {
if strings.TrimSpace(apiKey) != "" { if strings.TrimSpace(apiKey) != "" {
return nil return nil
} }
if !apiKeyRequired {
return nil
}
envName := strings.TrimSpace(apiKeyEnv) envName := strings.TrimSpace(apiKeyEnv)
if envName == "" { if envName == "" {
if apiKeyRequired { return ErrAPIKeyRequired
return ErrAPIKeyRequired
}
return nil
} }
if strings.TrimSpace(os.Getenv(envName)) == "" { if strings.TrimSpace(os.Getenv(envName)) == "" {
return fmt.Errorf("%w: api key environment variable %q is not set", ErrAPIKeyEnvMissing, envName) return fmt.Errorf("%w: api key environment variable %q is not set", ErrAPIKeyEnvMissing, envName)
@@ -630,18 +684,21 @@ func copyExtraParams(src map[string]any) map[string]any {
return cp 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 contract := def.Validation
if contract.Format == "" { if contract.Format == "" {
contract.Format = def.OutputFormat contract.Format = def.OutputFormat
} }
if override != nil { if override != nil {
contract = *override contract = *override
if contract.Format == "" {
contract.Format = domain.FormatText
}
} }
if contract.Format == "" { if err := domain.ValidateOutputContract(contract); err != nil {
contract.Format = domain.FormatText return domain.OutputContract{}, err
} }
return contract return contract, nil
} }
func hashRenderedPrompt(p domain.RenderedPrompt) string { func hashRenderedPrompt(p domain.RenderedPrompt) string {

View File

@@ -6,9 +6,11 @@ import (
"encoding/hex" "encoding/hex"
"errors" "errors"
"fmt" "fmt"
"math"
"path/filepath" "path/filepath"
"reflect" "reflect"
"regexp" "regexp"
"strconv"
"strings" "strings"
"sync" "sync"
"testing" "testing"
@@ -189,16 +191,34 @@ func (f *fakeValidator) Validate(ctx context.Context, artifact *domain.Artifact,
return f.result, nil return f.result, nil
} }
func (f *fakeValidator) LoadSchemaDocument(ctx context.Context, schemaPath string) (any, error) { func (f *fakeValidator) PrepareValidation(_ context.Context, contract domain.OutputContract) (validate.PreparedValidation, error) {
f.schemaLoads++ var schemaDocument any
f.schemaLoadPath = schemaPath if contract.ValidationMode == domain.ValidationJSONSchema {
if f.schemaErr != nil { f.schemaLoads++
return nil, f.schemaErr 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 &fakePreparedValidator{validator: f, contract: contract, schemaDocument: schemaDocument}, nil
return f.schemaDoc, nil }
}
return map[string]any{"type": "object"}, 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 { type fakeRepairer struct {
@@ -208,6 +228,30 @@ type fakeRepairer struct {
reqs []RepairRequest 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 { type fakeRunAdmitter struct {
backendIDs []string backendIDs []string
err error err error
@@ -502,6 +546,35 @@ func TestRunnerDirectSessionResolution(t *testing.T) {
t.Fatalf("invalid direct session invoked generation %d times", llmClient.calls) 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) { func TestRunnerPrepareUsesPromptDefaultProfileWhenNoExplicitProfileID(t *testing.T) {
@@ -704,16 +777,23 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
} }
func TestRunnerPrepareInvalidRequestNumericOverridesFail(t *testing.T) { func TestRunnerPrepareInvalidRequestNumericOverridesFail(t *testing.T) {
tests := []struct { type testCase struct {
name string name string
override *domain.ExecutionTargetOverride override *domain.ExecutionTargetOverride
}{ }
tests := []testCase{
{name: "temperature below range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(-0.1)}}, {name: "temperature below range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(-0.1)}},
{name: "temperature above range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(2.1)}}, {name: "temperature above range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(2.1)}},
{name: "max tokens below range", override: &domain.ExecutionTargetOverride{MaxTokens: intPtr(-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 below range", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(-0.1)}},
{name: "top p above range", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(1.1)}}, {name: "top p above range", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(1.1)}},
{name: "timeout below range", override: &domain.ExecutionTargetOverride{TimeoutSeconds: intPtr(-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 { for _, tc := range tests {
@@ -1357,14 +1437,14 @@ func TestRunnerAdmissionUsesResolvedBackendIdentity(t *testing.T) {
func TestRunnerAdmissionFailureSkipsCompletionCollaborators(t *testing.T) { func TestRunnerAdmissionFailureSkipsCompletionCollaborators(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
admissionError error admissionError error
wantBackendContext bool wantCapacityType bool
}{ }{
{ {
name: "capacity exhausted", name: "capacity exhausted",
admissionError: capacity.ErrCapacityExceeded, admissionError: capacity.ErrCapacityExceeded,
wantBackendContext: true, wantCapacityType: true,
}, },
{ {
name: "context canceled", name: "context canceled",
@@ -1414,8 +1494,16 @@ func TestRunnerAdmissionFailureSkipsCompletionCollaborators(t *testing.T) {
if errors.Is(err, ErrInvalidRequest) || errors.Is(err, ErrLLMGenerate) { if errors.Is(err, ErrInvalidRequest) || errors.Is(err, ErrLLMGenerate) {
t.Fatalf("admission error was recategorized: %v", err) t.Fatalf("admission error was recategorized: %v", err)
} }
if tc.wantBackendContext && !strings.Contains(err.Error(), "custom") { var capacityErr *CapacityError
t.Fatalf("capacity error lacks backend context: %v", err) 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"}) { if !reflect.DeepEqual(admitter.backendIDs, []string{"custom"}) {
t.Fatalf("admitted backend IDs=%#v, want custom", admitter.backendIDs) 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) { func TestRunnerRunAPIKeyEnvResolvesFromEnvironment(t *testing.T) {
t.Setenv("PROMPTKIT_TEST_API_KEY", "secret") t.Setenv("PROMPTKIT_TEST_API_KEY", "secret")
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
@@ -1729,22 +1787,26 @@ func TestRunnerRunAPIKeyEnvResolvesFromEnvironment(t *testing.T) {
} }
} }
func TestRunnerRunAPIKeyEnvMissingEnvironmentValueFailsClearly(t *testing.T) { func TestRunnerRunOptionalAPIKeyEnvMissingEnvironmentValueReachesLLM(t *testing.T) {
const environmentName = "PROMPTKIT_MISSING_KEY"
t.Setenv(environmentName, "")
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{
"exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "PROMPTKIT_MISSING_KEY"}, "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: environmentName},
}} }}
runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, nil) llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}
runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil)
_, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) result, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()})
if !errors.Is(err, ErrInvalidRequest) { if err != nil || result == nil {
t.Fatalf("expected ErrInvalidRequest, got %v", err) t.Fatalf("optional credential run = (%+v, %v), want success", result, err)
} }
if !errors.Is(err, ErrAPIKeyEnvMissing) { if llmClient.calls != 1 {
t.Fatalf("expected ErrAPIKeyEnvMissing, got %v", err) t.Fatalf("LLM calls = %d, want 1", llmClient.calls)
} }
if !strings.Contains(err.Error(), "PROMPTKIT_MISSING_KEY") { if llmClient.lastReq.Target.APIKeyEnv != environmentName {
t.Fatalf("expected missing env name in error, got %v", err) t.Fatalf("LLM api_key_env = %q, want %q", llmClient.lastReq.Target.APIKeyEnv, environmentName)
} }
} }
@@ -1757,7 +1819,7 @@ func TestRunnerRunDirectAPIKeyBypassesMissingEnvAndReachesLLM(t *testing.T) {
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}
runner := NewRunner(promptRepo, execRepo, nil, defaultArtifactReader(), defaultRenderer(), llmClient, nil, nil) 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", PromptID: "p",
ProfileID: "exec", ProfileID: "exec",
APIKey: directKey, APIKey: directKey,
@@ -1769,6 +1831,9 @@ func TestRunnerRunDirectAPIKeyBypassesMissingEnvAndReachesLLM(t *testing.T) {
if llmClient.lastReq.Target.APIKey != directKey { if llmClient.lastReq.Target.APIKey != directKey {
t.Fatalf("expected direct API key to reach LLM request") 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" { 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) t.Fatalf("expected api_key_env name to remain on target, got %q", llmClient.lastReq.Target.APIKeyEnv)
} }
@@ -2033,50 +2098,255 @@ func TestRunnerRunValidationStillWorks(t *testing.T) {
} }
} }
func TestRunnerRunStructuredRepairRemainsBoundedAndUsesEffectiveModelSettings(t *testing.T) { func TestRunnerRepairStateMachine(t *testing.T) {
repairer := &fakeRepairer{responses: []*domain.GenerateResponse{{Content: `{"broken":`}, {Content: `{"still":`}}} failed := func(mode domain.ValidationMode, diagnostic string) domain.ValidationResult {
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: `{"initial":`}} 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( tests := []struct {
&fakePromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 1)}, name string
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ mode domain.ValidationMode
"exec": {ID: "exec", BackendID: "custom", Model: "profile-model", TimeoutSeconds: 55}, budget int
}}, fakeBackendResolver{backends: map[string]domain.Backend{ validationResults []domain.ValidationResult
"custom": {ID: "custom", Endpoint: "http://backend/v1"}, 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(), for _, tc := range tests {
defaultRenderer(), t.Run(tc.name, func(t *testing.T) {
llmClient, definition := promptDef(domain.FormatJSON, tc.mode, tc.budget)
validate.NewStandardValidator("."), plan := &recordingPreparedValidation{results: tc.validationResults}
repairer, nil) 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{ result, err := runner.Run(context.Background(), domain.RunRequest{
PromptID: "p", PromptID: "p",
ProfileID: "exec", ProfileID: "exec",
Inputs: singleInputRef(), SessionID: " repair-session ",
Execution: &domain.ExecutionTargetOverride{Endpoint: "http://override/v1", Model: "override-model", TimeoutSeconds: intPtr(22)}, APIKey: "direct-secret",
}) Inputs: singleInputRef(),
if err != nil { Execution: tc.execution,
t.Fatalf("expected no error, got %v", err) })
} if err != nil {
if repairer.calls != 1 || res.Validation.RepairAttempts != 1 { t.Fatalf("run: %v", err)
t.Fatalf("expected one bounded repair, calls=%d attempts=%d", repairer.calls, res.Validation.RepairAttempts) }
} if len(client.requests) != tc.wantRepairs+1 || len(repairer.reqs) != tc.wantRepairs {
if len(repairer.reqs) != 1 { t.Fatalf(
t.Fatalf("expected one repair request, got %d", len(repairer.reqs)) "generation/repair calls = (%d, %d), want (%d, %d)",
} len(client.requests), len(repairer.reqs), tc.wantRepairs+1, tc.wantRepairs,
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 len(plan.artifacts) != tc.wantRepairs+1 {
if repairer.reqs[0].Target.TimeoutSeconds != 22 { t.Fatalf("validation calls = %d, want %d", len(plan.artifacts), tc.wantRepairs+1)
t.Fatalf("expected repair to use effective timeout, got %d", repairer.reqs[0].Target.TimeoutSeconds) }
}
if llmClient.lastReq.Target.BackendID != "custom" || initialRequest := client.requests[0]
repairer.reqs[0].Target.BackendID != "custom" || if initialRequest.TargetPresence != tc.wantPresence {
res.SelectedBackendID != "custom" { t.Fatalf("initial target presence = %+v, want %+v", initialRequest.TargetPresence, tc.wantPresence)
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) 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 +2441,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) { func TestExecutionProfileToTargetPopulatesAllFieldsAndCopiesExtraParams(t *testing.T) {
src := &domain.ExecutionProfile{ src := &domain.ExecutionProfile{
ID: "exec", ID: "exec",
@@ -2319,10 +2499,7 @@ func TestResolveExecutionTargetProfileValuesPopulateAllSupportedFields(t *testin
}, },
} }
target, presence, err := resolveExecutionTarget(nil, profileValue, nil) target, presence := resolveExecutionTarget(nil, profileValue, nil)
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if presence != (domain.ExecutionTargetPresence{}) { if presence != (domain.ExecutionTargetPresence{}) {
t.Fatalf("expected no request override presence, got %+v", presence) t.Fatalf("expected no request override presence, got %+v", presence)
} }
@@ -2374,10 +2551,7 @@ func TestResolveExecutionTargetRuntimeOverridesBeatProfileForAllOverrideableFiel
}, },
} }
target, presence, err := resolveExecutionTarget(nil, profileValue, override) target, presence := resolveExecutionTarget(nil, profileValue, override)
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if presence != (domain.ExecutionTargetPresence{Temperature: true, MaxTokens: true, TopP: true, TimeoutSeconds: true}) { if presence != (domain.ExecutionTargetPresence{Temperature: true, MaxTokens: true, TopP: true, TimeoutSeconds: true}) {
t.Fatalf("unexpected override presence: %+v", presence) t.Fatalf("unexpected override presence: %+v", presence)
} }
@@ -2424,12 +2598,9 @@ func TestResolveExecutionTargetReasoningOverrideStates(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
target, _, err := resolveExecutionTarget(nil, profileValue, &domain.ExecutionTargetOverride{ target, _ := resolveExecutionTarget(nil, profileValue, &domain.ExecutionTargetOverride{
ReasoningEffort: tt.override, ReasoningEffort: tt.override,
}) })
if err != nil {
t.Fatalf("resolve execution target: %v", err)
}
if target.ReasoningEffort != tt.want { if target.ReasoningEffort != tt.want {
t.Fatalf("reasoning effort = %q, want %q", target.ReasoningEffort, tt.want) t.Fatalf("reasoning effort = %q, want %q", target.ReasoningEffort, tt.want)
} }
@@ -2558,10 +2729,7 @@ func TestResolveExecutionTargetUsesBackendProfileAndRequestPrecedence(t *testing
ExtraParams: map[string]any{"request": true}, ExtraParams: map[string]any{"request": true},
} }
target, _, err := resolveExecutionTarget(backendValue, profileValue, override) target, _ := resolveExecutionTarget(backendValue, profileValue, override)
if err != nil {
t.Fatalf("resolve target: %v", err)
}
if target.BackendID != "custom" { if target.BackendID != "custom" {
t.Fatalf("endpoint override changed backend identity: %+v", target) t.Fatalf("endpoint override changed backend identity: %+v", target)
} }
@@ -2572,12 +2740,9 @@ func TestResolveExecutionTargetUsesBackendProfileAndRequestPrecedence(t *testing
t.Fatalf("expected whole-map request replacement, got %#v", target.ExtraParams) 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", ID: "exec", BackendID: "custom", Model: "profile-model",
}, nil) }, nil)
if err != nil {
t.Fatalf("resolve backend defaults: %v", err)
}
if target.Endpoint != backendValue.Endpoint || if target.Endpoint != backendValue.Endpoint ||
target.APIKeyEnv != backendValue.APIKeyEnv || target.APIKeyEnv != backendValue.APIKeyEnv ||
!reflect.DeepEqual(target.ExtraParams, backendValue.ExtraParams) { !reflect.DeepEqual(target.ExtraParams, backendValue.ExtraParams) {

View 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
}
}

View File

@@ -1,15 +1,18 @@
package validate package validate
import ( import (
"bytes"
"context" "context"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"io"
"io/fs" "io/fs"
"net/url" "net/url"
"os" "os"
"path" "path"
"path/filepath" "path/filepath"
"runtime"
"strings" "strings"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/domain"
@@ -19,6 +22,8 @@ import (
const jsonSchemaDraft2020 = "https://json-schema.org/draft/2020-12/schema" const jsonSchemaDraft2020 = "https://json-schema.org/draft/2020-12/schema"
const schemaReadChunkSize = 64 * 1024
// StandardValidator provides basic, JSON, and JSON Schema output validation. // StandardValidator provides basic, JSON, and JSON Schema output validation.
type StandardValidator struct { type StandardValidator struct {
schemaBaseDir string schemaBaseDir string
@@ -45,7 +50,121 @@ func (v *FSValidator) Validate(ctx context.Context, artifact *domain.Artifact, c
return validateArtifact(ctx, artifact, contract, v.validateJSONSchema) 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) { func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract, validateSchema schemaValidatorFunc) (domain.ValidationResult, error) {
select { select {
@@ -70,7 +189,11 @@ func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract d
res.IsValid = true res.IsValid = true
return res, nil return res, nil
case domain.ValidationBasic: 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.Status = domain.ValidationFailed
res.IsValid = false res.IsValid = false
res.Errors = []string{"output is empty"} res.Errors = []string{"output is empty"}
@@ -80,26 +203,32 @@ func validateArtifact(ctx context.Context, artifact *domain.Artifact, contract d
res.IsValid = true res.IsValid = true
return res, nil return res, nil
case domain.ValidationJSON: case domain.ValidationJSON:
_, jsonErr := parseJSON(artifact.Body) valid := json.Valid(artifact.Body)
if jsonErr != nil { if err := ctx.Err(); err != nil {
return domain.ValidationResult{}, err
}
if !valid {
res.Status = domain.ValidationFailed res.Status = domain.ValidationFailed
res.IsValid = false res.IsValid = false
res.Errors = []string{fmt.Sprintf("invalid JSON: %v", jsonErr)} res.Errors = []string{"invalid JSON"}
return res, nil return res, nil
} }
res.Status = domain.ValidationPassed res.Status = domain.ValidationPassed
res.IsValid = true res.IsValid = true
return res, nil return res, nil
case domain.ValidationJSONSchema: case domain.ValidationJSONSchema:
instance, jsonErr := parseJSON(artifact.Body) instance, jsonErr := decodeJSONValue(ctx, artifact.Body)
if jsonErr != nil { if jsonErr != nil {
if contextErr := ctx.Err(); contextErr != nil {
return domain.ValidationResult{}, contextErr
}
res.Status = domain.ValidationFailed res.Status = domain.ValidationFailed
res.IsValid = false res.IsValid = false
res.Errors = []string{fmt.Sprintf("invalid JSON: %v", jsonErr)} res.Errors = []string{fmt.Sprintf("invalid JSON: %v", jsonErr)}
return res, nil return res, nil
} }
validationErrors, err := validateSchema(instance, contract.SchemaPath) validationErrors, err := validateSchema(ctx, instance, contract.SchemaPath)
if err != nil { if err != nil {
return domain.ValidationResult{}, err 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) { func (v *StandardValidator) validateJSONSchema(ctx context.Context, instance any, schemaPath string) ([]string, error) {
resolvedSchemaPath, err := v.resolveSchemaPath(schemaPath) resolvedSchemaPath, err := v.resolveSchemaPath(ctx, schemaPath)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -128,20 +257,21 @@ func (v *StandardValidator) validateJSONSchema(instance any, schemaPath string)
if err != nil { if err != nil {
return nil, err return nil, err
} }
compiler := newSchemaCompiler(standardSchemaLoader{root: schemaRoot}) if err := ctx.Err(); err != nil {
schema, err := compiler.Compile(resolvedSchemaPath) return nil, err
}
compiler := newSchemaCompiler(standardSchemaLoader{ctx: ctx, root: schemaRoot})
resourceURL := fileSchemaResourceURL(resolvedSchemaPath)
schema, err := compileJSONSchema(ctx, resourceURL.String(), compiler.Compile)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", resolvedSchemaPath, err) return nil, fmt.Errorf("failed to compile JSON schema %q: %w", resolvedSchemaPath, err)
} }
if err := schema.Validate(instance); err != nil { return executeJSONSchema(ctx, schema, instance)
return []string{fmt.Sprintf("json schema validation failed: %v", err)}, nil
}
return nil, nil
} }
func (v *FSValidator) validateJSONSchema(instance any, schemaPath string) ([]string, error) { func (v *FSValidator) validateJSONSchema(ctx context.Context, instance any, schemaPath string) ([]string, error) {
schemaName, schemaDoc, err := v.loadSchemaDocument(schemaPath) schemaName, schemaDoc, err := v.loadSchemaDocument(ctx, schemaPath)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -150,71 +280,60 @@ func (v *FSValidator) validateJSONSchema(instance any, schemaPath string) ([]str
if err := validateSchemaDialect(schemaDoc); err != nil { if err := validateSchemaDialect(schemaDoc); err != nil {
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", schemaName, err) return nil, fmt.Errorf("failed to compile JSON schema %q: %w", schemaName, err)
} }
compiler := newSchemaCompiler(fsSchemaLoader{fsys: v.fsys, root: filecatalog.CleanFSRoot(v.root)}) compiler := newSchemaCompiler(fsSchemaLoader{ctx: ctx, fsys: v.fsys, root: filecatalog.CleanFSRoot(v.root)})
if err := compiler.AddResource(resourceURL, schemaDoc); err != nil { 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) 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 { if err != nil {
return nil, fmt.Errorf("failed to compile JSON schema %q: %w", schemaName, err) return nil, fmt.Errorf("failed to compile JSON schema %q: %w", schemaName, err)
} }
if err := schema.Validate(instance); err != nil { return executeJSONSchema(ctx, schema, instance)
return []string{fmt.Sprintf("json schema validation failed: %v", err)}, nil
}
return nil, nil
} }
func parseJSON(body []byte) (any, error) { func decodeJSONValue(ctx context.Context, body []byte) (any, error) {
var v any if err := ctx.Err(); err != nil {
if err := json.Unmarshal(body, &v); err != nil {
return nil, err 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) { var value any
select { decodeErr := decoder.Decode(&value)
case <-ctx.Done(): if err := ctx.Err(); err != nil {
return nil, ctx.Err()
default:
}
resolved, err := v.resolveSchemaPath(schemaPath)
if err != nil {
return nil, err return nil, err
} }
if decodeErr != nil {
raw, err := os.ReadFile(resolved) return nil, decodeErr
if err != nil {
return nil, fmt.Errorf("failed to read schema file %q: %w", resolved, err)
} }
var doc any var trailing any
if err := json.Unmarshal(raw, &doc); err != nil { trailingErr := decoder.Decode(&trailing)
return nil, fmt.Errorf("failed to decode JSON schema %q: %w", resolved, err) if err := ctx.Err(); err != nil {
}
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 {
return nil, err 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) == "" { if strings.TrimSpace(schemaPath) == "" {
return "", errors.New("schema path is required for json_schema validation") 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 { if err != nil {
return "", err return "", err
} }
if err := ctx.Err(); err != nil {
return "", err
}
resolved, err := containedFilesystemPath(root, schemaPath) resolved, err := containedFilesystemPath(root, schemaPath)
if err != nil { if err != nil {
return "", err return "", err
} }
if err := ctx.Err(); err != nil {
return "", err
}
if _, err := os.Stat(resolved); err != nil { 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) return "", fmt.Errorf("failed to access schema file %q: %w", resolved, err)
} }
if err := ctx.Err(); err != nil {
return "", err
}
return resolved, nil return resolved, nil
} }
@@ -250,28 +381,39 @@ func (v *StandardValidator) schemaRoot() (string, error) {
return resolved, nil return resolved, nil
} }
func (v *FSValidator) loadSchemaDocument(schemaPath string) (string, any, error) { func (v *FSValidator) loadSchemaDocument(ctx context.Context, schemaPath string) (string, any, error) {
resolved, err := v.resolveSchemaPath(schemaPath) resolved, err := v.resolveSchemaPath(ctx, schemaPath)
if err != nil { if err != nil {
return "", nil, err 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 { if err != nil {
return "", nil, fmt.Errorf("failed to read schema file %q: %w", resolved, err) return "", nil, fmt.Errorf("failed to read schema file %q: %w", resolved, err)
} }
var doc any doc, err := decodeJSONValue(ctx, raw)
if err := json.Unmarshal(raw, &doc); err != nil { if err != nil {
return "", nil, fmt.Errorf("failed to decode JSON schema %q: %w", resolved, err) return "", nil, fmt.Errorf("failed to decode JSON schema %q: %w", resolved, err)
} }
if err := validateSchemaDialect(doc); err != nil { 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) 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 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) == "" { if strings.TrimSpace(schemaPath) == "" {
return "", errors.New("schema path is required for json_schema validation") 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) cleanRoot := filecatalog.CleanFSRoot(v.root)
rootInfo, err := fs.Stat(v.fsys, cleanRoot) rootInfo, err := fs.Stat(v.fsys, cleanRoot)
if err != nil { if err != nil {
if contextErr := ctx.Err(); contextErr != nil {
return "", contextErr
}
return "", fmt.Errorf("failed to access schema source %q: %w", cleanRoot, err) return "", fmt.Errorf("failed to access schema source %q: %w", cleanRoot, err)
} }
if err := ctx.Err(); err != nil {
return "", err
}
var resolved string var resolved string
if rootInfo.IsDir() { if rootInfo.IsDir() {
@@ -297,15 +445,21 @@ func (v *FSValidator) resolveSchemaPath(schemaPath string) (string, error) {
if err != nil { if err != nil {
return "", err 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)) return "", fmt.Errorf("schema path %q does not match schema file %q", cleanSchemaPath, path.Base(cleanRoot))
} }
resolved = cleanRoot resolved = cleanRoot
} }
if _, err := fs.Stat(v.fsys, resolved); err != nil { 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) return "", fmt.Errorf("failed to access schema file %q: %w", resolved, err)
} }
if err := ctx.Err(); err != nil {
return "", err
}
return resolved, nil return resolved, nil
} }
@@ -321,8 +475,19 @@ func cleanSchemaFSPath(schemaPath string) (string, error) {
return cleaned, nil return cleaned, nil
} }
func fsSchemaResourceURL(schemaName string) string { func fileSchemaResourceURL(schemaName string) *url.URL {
return "promptkit-schema:///" + strings.TrimPrefix(path.Clean(schemaName), "/") 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 { func newSchemaCompiler(loader jsonschema.URLLoader) *jsonschema.Compiler {
@@ -332,6 +497,35 @@ func newSchemaCompiler(loader jsonschema.URLLoader) *jsonschema.Compiler {
return 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 { func validateSchemaDialect(doc any) error {
object, ok := doc.(map[string]any) object, ok := doc.(map[string]any)
if !ok { if !ok {
@@ -352,19 +546,33 @@ func validateSchemaDialect(doc any) error {
} }
type standardSchemaLoader struct { type standardSchemaLoader struct {
ctx context.Context
root string root string
} }
func (l standardSchemaLoader) Load(resourceURL string) (any, error) { 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) fileName, err := (jsonschema.FileLoader{}).ToFile(resourceURL)
if err != nil { if err != nil {
return nil, fmt.Errorf("schema reference %q is not a contained file reference: %w", resourceURL, err) return nil, fmt.Errorf("schema reference %q is not a contained file reference: %w", resourceURL, err)
} }
resolved, err := containedFilesystemPath(l.root, fileName) resolved, err := containedFilesystemPath(l.root, fileName)
if err != nil { if err != nil {
if contextErr := l.ctx.Err(); contextErr != nil {
return nil, contextErr
}
return nil, err 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) { func containedFilesystemPath(root, name string) (string, error) {
@@ -390,38 +598,94 @@ func containedFilesystemPath(root, name string) (string, error) {
return candidate, nil return candidate, nil
} }
func loadJSONSchemaFile(name string) (any, error) { func loadJSONSchemaFile(ctx context.Context, name string) (any, error) {
raw, err := os.ReadFile(name) raw, err := readSchemaFile(ctx, func() (fs.File, error) {
return os.Open(name)
})
if err != nil { if err != nil {
return nil, err return nil, err
} }
var doc any doc, err := decodeJSONValue(ctx, raw)
if err := json.Unmarshal(raw, &doc); err != nil { if err != nil {
return nil, err return nil, err
} }
if err := validateSchemaDialect(doc); err != nil { 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 nil, err
} }
return doc, nil 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 { type fsSchemaLoader struct {
ctx context.Context
fsys fs.FS fsys fs.FS
root string root string
} }
func (l fsSchemaLoader) Load(resourceURL string) (any, error) { func (l fsSchemaLoader) Load(resourceURL string) (any, error) {
if err := l.ctx.Err(); err != nil {
return nil, err
}
parsed, err := url.Parse(resourceURL) parsed, err := url.Parse(resourceURL)
if err != nil { if err != nil {
return nil, fmt.Errorf("invalid schema reference %q: %w", resourceURL, err) 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) return nil, fmt.Errorf("schema reference %q is not allowed", resourceURL)
} }
name, err := url.PathUnescape(strings.TrimPrefix(parsed.Path, "/")) name := strings.TrimPrefix(parsed.Path, "/")
if err != nil {
return nil, fmt.Errorf("invalid schema reference %q: %w", resourceURL, err)
}
name = path.Clean(name) name = path.Clean(name)
if l.root == "." { if l.root == "." {
if strings.HasPrefix(name, "../") || name == ".." { 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) rootInfo, err := fs.Stat(l.fsys, l.root)
if err != nil { 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 return nil, err
} }
if !rootInfo.IsDir() && name != l.root { if !rootInfo.IsDir() && name != l.root {
return nil, fmt.Errorf("schema reference %q is outside the configured schema file", resourceURL) 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 { if err != nil {
return nil, err return nil, err
} }
var doc any doc, err := decodeJSONValue(l.ctx, raw)
if err := json.Unmarshal(raw, &doc); err != nil { if err != nil {
return nil, err return nil, err
} }
if err := validateSchemaDialect(doc); err != nil { 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 nil, err
} }
return doc, nil return doc, nil

View File

@@ -3,8 +3,11 @@ package validate
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"net/url"
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"runtime"
"strconv" "strconv"
"strings" "strings"
"testing" "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) { func TestStandardValidatorJSONSchemaSuccess(t *testing.T) {
tmp := t.TempDir() tmp := t.TempDir()
schemaPath := filepath.Join(tmp, "schema.json") 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) { func TestStandardValidatorJSONSchemaNestedSchemaPathSuccess(t *testing.T) {
tmp := t.TempDir() tmp := t.TempDir()
nestedDir := filepath.Join(tmp, "dnd") 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) { func TestFSValidatorJSONSchemaSuccess(t *testing.T) {
v := NewFSValidator(fstest.MapFS{ v := NewFSValidator(fstest.MapFS{
"schemas/events.schema.json": &fstest.MapFile{Data: []byte(`{ "schemas/events.schema.json": &fstest.MapFile{Data: []byte(`{
@@ -295,17 +356,128 @@ func TestFSValidatorJSONSchemaSuccess(t *testing.T) {
} }
} }
func TestFSValidatorJSONSchemaRegistrationError(t *testing.T) { func TestFSValidatorPreparedSchemaSurvivesSourceMutation(t *testing.T) {
v := NewFSValidator(fstest.MapFS{ rootSchema := []byte(`{
"schemas/%zz.json": &fstest.MapFile{Data: []byte(`{"type":"object"}`)}, "$schema": "https://json-schema.org/draft/2020-12/schema",
}, "schemas") "title": "original root",
"type": "object",
res, err := v.Validate(context.Background(), &domain.Artifact{Body: []byte(`{}`)}, domain.OutputContract{ "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, ValidationMode: domain.ValidationJSONSchema,
SchemaPath: "%zz.json", SchemaPath: "schema.json",
}) })
if err == nil || !strings.Contains(err.Error(), "failed to register JSON schema") { if err != nil {
t.Fatalf("expected schema registration error, got result=%#v error=%v", res, err) 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) { func TestStandardValidatorJSONSchemaReferenceBoundaries(t *testing.T) {
root := t.TempDir() root := t.TempDir()
if err := os.WriteFile(filepath.Join(root, "child.json"), []byte(`{ 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) { func TestJSONSchemaDialectIsDraft2020(t *testing.T) {
tests := []struct { tests := []struct {
name string 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)
}
}

View File

@@ -2,15 +2,33 @@ package validate
import ( import (
"context" "context"
"gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/domain"
) )
// Validator validates the generated artifact based on the output contract. // 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 { type Validator interface {
Validate(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract) (domain.ValidationResult, error) Validate(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract) (domain.ValidationResult, error)
} }
// SchemaDocumentLoader loads JSON schema documents using validator path semantics. // PreparedValidation validates artifacts against one frozen output contract.
type SchemaDocumentLoader interface { // Its cancellation boundary is synchronous: Validate does not detach schema
LoadSchemaDocument(ctx context.Context, schemaPath string) (any, error) // 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
View File

@@ -2,9 +2,16 @@ package promptkit
import ( import (
"encoding/json" "encoding/json"
"fmt"
"math"
"time" "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 // MarshalJSON implements json.Marshaler for PreparedRun. It uses RFC 3339
// timestamps, integer duration_ms, and omits zero timing values. // timestamps, integer duration_ms, and omits zero timing values.
func (r PreparedRun) MarshalJSON() ([]byte, error) { func (r PreparedRun) MarshalJSON() ([]byte, error) {
@@ -21,38 +28,11 @@ func (r PreparedRun) MarshalJSON() ([]byte, error) {
durationMS = &r.DurationMS durationMS = &r.DurationMS
} }
return json.Marshal(struct { return json.Marshal(preparedRunJSON{
PromptID string `json:"prompt_id"` preparedRunJSONFields: preparedRunJSONFields(r),
PromptVersion string `json:"prompt_version,omitempty"` StartTime: startTime,
PromptHash string `json:"prompt_hash,omitempty"` EndTime: endTime,
SelectedProfileID string `json:"selected_profile_id"` DurationMS: durationMS,
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,
}) })
} }
@@ -74,89 +54,57 @@ func (r RunResult) MarshalJSON() ([]byte, error) {
} }
return json.Marshal(runResultJSON{ return json.Marshal(runResultJSON{
RunID: r.RunID, runResultJSONFields: runResultJSONFields(r),
Artifact: r.Artifact, StartTime: startTime,
RawOutput: r.RawOutput, EndTime: endTime,
Validation: r.Validation, DurationMS: durationMS,
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,
}) })
} }
// UnmarshalJSON implements json.Unmarshaler for RunResult. It decodes // 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 { func (r *RunResult) UnmarshalJSON(data []byte) error {
var wire runResultJSON var wire runResultJSON
if err := json.Unmarshal(data, &wire); err != nil { if err := json.Unmarshal(data, &wire); err != nil {
return err return err
} }
*r = RunResult{ result := RunResult(wire.runResultJSONFields)
RunID: wire.RunID, if wire.DurationMS != nil {
Artifact: wire.Artifact, if *wire.DurationMS < minDurationMilliseconds || *wire.DurationMS > maxDurationMilliseconds {
RawOutput: wire.RawOutput, return fmt.Errorf(
Validation: wire.Validation, "decode RunResult duration_ms: %d cannot be represented as time.Duration",
PromptID: wire.PromptID, *wire.DurationMS,
PromptVersion: wire.PromptVersion, )
PromptHash: wire.PromptHash, }
SessionID: wire.SessionID, result.Duration = time.Duration(*wire.DurationMS) * time.Millisecond
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,
} }
if wire.StartTime != nil { if wire.StartTime != nil {
r.StartTime = *wire.StartTime result.StartTime = *wire.StartTime
} }
if wire.EndTime != nil { if wire.EndTime != nil {
r.EndTime = *wire.EndTime result.EndTime = *wire.EndTime
} }
*r = result
return nil return nil
} }
type runResultJSON struct { type preparedRunJSONFields PreparedRun
RunID string `json:"run_id"`
Artifact Artifact `json:"artifact"` type preparedRunJSON struct {
RawOutput string `json:"raw_output"` preparedRunJSONFields
Validation ValidationResult `json:"validation"` StartTime *time.Time `json:"start_time,omitempty"`
PromptID string `json:"prompt_id"` EndTime *time.Time `json:"end_time,omitempty"`
PromptVersion string `json:"prompt_version,omitempty"` DurationMS *int64 `json:"duration_ms,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"`
} }
func valueOrZero(value *int64) int64 { type runResultJSONFields RunResult
if value == nil {
return 0 type runResultJSON struct {
} runResultJSONFields
return *value 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
View 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)
}
}
}

View 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
}

View 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
View 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
}

View File

@@ -0,0 +1,789 @@
package promptkit_test
import (
"context"
"encoding/json"
"errors"
"fmt"
"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.WithProfiles(promptkit.Profile{
ID: "profile",
Endpoint: "http://example.test/v1",
Model: "model",
APIKeyRequired: true,
}),
promptkit.WithLLMClient(client),
)
if err != nil {
t.Fatalf("construct credential engine: %v", err)
}
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{
PromptID: "prepared",
Execution: &promptkit.ExecutionTargetOverride{APIKeyEnv: environmentName},
})
if err != nil {
t.Fatalf("prepare credential execution: %v", err)
}
t.Setenv(environmentName, "")
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 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)
}
}
}

View File

@@ -0,0 +1,243 @@
package promptkit_test
import (
"context"
"errors"
"io/fs"
"strings"
"sync"
"testing"
"testing/fstest"
"gitea.maximumdirect.net/eric/promptkit"
)
func TestProfileInheritanceBuiltInAliasWorkflow(t *testing.T) {
engine, err := promptkit.NewEngine(promptkit.Config{
PromptDir: frameworkPromptDir,
SchemaDir: frameworkSchemaDir,
}, promptkit.WithProfiles(promptkit.Profile{
ID: "weather-light",
BaseProfileID: "deepseek-4-flash",
ReasoningEffort: "high",
TimeoutSeconds: 120,
}))
if err != nil {
t.Fatalf("construct alias engine: %v", err)
}
base, err := engine.InspectProfile(context.Background(), "deepseek-4-flash")
if err != nil {
t.Fatalf("inspect base: %v", err)
}
child, err := engine.InspectProfile(context.Background(), "weather-light")
if err != nil {
t.Fatalf("inspect alias: %v", err)
}
if child.ProfileID != "weather-light" ||
child.EffectiveModelParams.BackendID != base.EffectiveModelParams.BackendID ||
child.EffectiveModelParams.Model != base.EffectiveModelParams.Model ||
child.EffectiveModelParams.ReasoningEffort != "high" ||
child.EffectiveModelParams.TimeoutSeconds != 120 {
t.Fatalf("alias inspection = %+v, base = %+v", child, base)
}
prepared, err := engine.Prepare(context.Background(), promptkit.RunRequest{
PromptID: frameworkMarkdownSummaryPromptID,
ProfileID: "weather-light",
Inputs: map[string]promptkit.ArtifactRef{
"transcript": promptkit.Inline("Rin opens the gate."),
"glossary": promptkit.Inline("gate: A guarded passage."),
},
})
if err != nil {
t.Fatalf("prepare alias: %v", err)
}
if prepared.SelectedProfileID != "weather-light" {
t.Fatalf("SelectedProfileID = %q", prepared.SelectedProfileID)
}
timeout := 15
reasoning := "low"
overridden, err := engine.Prepare(context.Background(), promptkit.RunRequest{
PromptID: frameworkMarkdownSummaryPromptID,
ProfileID: "weather-light",
Execution: &promptkit.ExecutionTargetOverride{
TimeoutSeconds: &timeout,
ReasoningEffort: &reasoning,
},
Inputs: map[string]promptkit.ArtifactRef{
"transcript": promptkit.Inline("Rin opens the gate."),
"glossary": promptkit.Inline("gate: A guarded passage."),
},
})
if err != nil {
t.Fatalf("prepare override: %v", err)
}
if overridden.EffectiveModelParams.TimeoutSeconds != timeout ||
overridden.EffectiveModelParams.ReasoningEffort != reasoning {
t.Fatalf("runtime override target = %+v", overridden.EffectiveModelParams)
}
}
func TestProfileInheritanceYAMLAliasOfBuiltIn(t *testing.T) {
engine, err := promptkit.NewEngine(promptkit.Config{},
promptkit.WithPromptFS(fstest.MapFS{}, "."),
promptkit.WithProfileFS(fstest.MapFS{
"alias.yaml": &fstest.MapFile{Data: []byte("id: yaml-alias\nbase_profile: deepseek-4-flash\n")},
}, "."),
)
if err != nil {
t.Fatalf("construct YAML alias engine: %v", err)
}
inspection, err := engine.InspectProfile(context.Background(), "yaml-alias")
if err != nil {
t.Fatalf("inspect YAML alias: %v", err)
}
if inspection.ProfileID != "yaml-alias" || inspection.EffectiveModelParams.Model == "" {
t.Fatalf("YAML alias inspection = %+v", inspection)
}
}
func TestProfileInheritanceRetainsRequiredCredentialBehavior(t *testing.T) {
engine, err := promptkit.NewEngine(promptkit.Config{},
promptkit.WithPromptFS(fstest.MapFS{}, "."),
promptkit.WithBackend(promptkit.Backend{
ID: "credential-backend",
Endpoint: "https://credential.example/v1",
APIKeyEnv: "OPTIONAL_BACKEND_KEY",
}),
promptkit.WithProfiles(
promptkit.Profile{
ID: "credential-base",
BackendID: "credential-backend",
Model: "model",
APIKeyRequired: true,
},
promptkit.Profile{ID: "credential-child", BaseProfileID: "credential-base"},
),
)
if err != nil {
t.Fatalf("construct credential inheritance engine: %v", err)
}
inspection, err := engine.InspectProfile(context.Background(), "credential-child")
if err != nil {
t.Fatalf("inspect credential child: %v", err)
}
if !inspection.APIKeyRequired || inspection.EffectiveModelParams.APIKeyEnv != "" {
t.Fatalf("credential inspection = %+v", inspection)
}
}
func TestProfileInheritancePreservesPublicErrorIdentities(t *testing.T) {
newEngine := func(t *testing.T, profiles fs.FS) *promptkit.Engine {
t.Helper()
options := []promptkit.Option{promptkit.WithPromptFS(fstest.MapFS{}, ".")}
if profiles != nil {
options = append(options, promptkit.WithProfileFS(profiles, "."))
}
engine, err := promptkit.NewEngine(promptkit.Config{}, options...)
if err != nil {
t.Fatalf("construct engine: %v", err)
}
return engine
}
tests := []struct {
name string
profiles fs.FS
profile string
contains []string
want error
wantNot error
}{
{
name: "missing selected profile",
profile: "missing",
want: promptkit.ErrProfileNotFound,
wantNot: promptkit.ErrProfileLoad,
},
{
name: "missing base",
profiles: fstest.MapFS{
"child.yaml": &fstest.MapFile{Data: []byte("id: child\nbase_profile: missing\n")},
},
profile: "child",
contains: []string{"child", "missing"},
want: promptkit.ErrProfileLoad,
wantNot: promptkit.ErrProfileNotFound,
},
{
name: "cycle",
profiles: fstest.MapFS{
"a.yaml": &fstest.MapFile{Data: []byte("id: a\nbase_profile: b\n")},
"b.yaml": &fstest.MapFile{Data: []byte("id: b\nbase_profile: a\n")},
},
profile: "a",
contains: []string{"a", "b"},
want: promptkit.ErrProfileLoad,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result, err := newEngine(t, tc.profiles).InspectProfile(context.Background(), tc.profile)
if result != nil || !errors.Is(err, tc.want) || (tc.wantNot != nil && errors.Is(err, tc.wantNot)) {
t.Fatalf("inspection=(%+v, %v), want %v without %v", result, err, tc.want, tc.wantNot)
}
for _, fragment := range tc.contains {
if !strings.Contains(err.Error(), fragment) {
t.Fatalf("error = %v, want %q", err, fragment)
}
}
})
}
}
func TestProfileInheritanceFreezesPreparedExecution(t *testing.T) {
profiles := &mutableInheritanceProfileFS{files: fstest.MapFS{
"child.yaml": &fstest.MapFile{Data: []byte("id: child\nbase_profile: base\n")},
"base.yaml": &fstest.MapFile{Data: []byte("id: base\nendpoint: https://base.example/v1\nmodel: first-model\n")},
}}
client := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "ok"}}
engine, err := promptkit.NewEngine(promptkit.Config{},
promptkit.WithPromptFS(contractPromptFS("prepared", "child", "content"), "."),
promptkit.WithProfileFS(profiles, "."),
promptkit.WithLLMClient(client),
)
if err != nil {
t.Fatalf("construct engine: %v", err)
}
prepared, err := engine.PrepareExecution(context.Background(), promptkit.RunRequest{PromptID: "prepared"})
if err != nil {
t.Fatalf("prepare execution: %v", err)
}
profiles.set("base.yaml", "id: base\nendpoint: https://base.example/v1\nmodel: second-model\n")
result, err := engine.RunPrepared(context.Background(), prepared)
if err != nil || result == nil || len(client.requests) != 1 || client.requests[0].Target.Model != "first-model" {
t.Fatalf("prepared execution=(%+v, %v), requests=%+v", result, err, client.requests)
}
inspection, err := engine.InspectProfile(context.Background(), "child")
if err != nil || inspection.EffectiveModelParams.Model != "second-model" {
t.Fatalf("fresh inspection=(%+v, %v)", inspection, err)
}
}
type mutableInheritanceProfileFS struct {
mu sync.RWMutex
files fstest.MapFS
}
func (f *mutableInheritanceProfileFS) Open(name string) (fs.File, error) {
f.mu.RLock()
defer f.mu.RUnlock()
return f.files.Open(name)
}
func (f *mutableInheritanceProfileFS) set(name, content string) {
f.mu.Lock()
defer f.mu.Unlock()
f.files[name] = &fstest.MapFile{Data: []byte(content)}
}

Some files were not shown because too many files have changed in this diff Show More