Implement runtime parameter completion fixes
All checks were successful
ci/woodpecker/tag/release Pipeline was successful
All checks were successful
ci/woodpecker/tag/release Pipeline was successful
This commit is contained in:
@@ -28,9 +28,9 @@ Serialized JSON fields:
|
|||||||
- `model` (required after fallback resolution)
|
- `model` (required after fallback resolution)
|
||||||
- `session_id` (only when the rendered prompt includes a non-empty session ID)
|
- `session_id` (only when the rendered prompt includes a non-empty session ID)
|
||||||
- `messages` (rendered prompt messages)
|
- `messages` (rendered prompt messages)
|
||||||
- `temperature` (only when non-zero)
|
- `temperature` (when non-zero, or when explicitly overridden to zero)
|
||||||
- `max_tokens` (only when non-zero)
|
- `max_tokens` (when non-zero, or when explicitly overridden to zero)
|
||||||
- `top_p` (only when non-zero)
|
- `top_p` (when non-zero, or when explicitly overridden to zero)
|
||||||
- `service_tier` (only when non-empty)
|
- `service_tier` (only when non-empty)
|
||||||
- `reasoning_effort` (only when non-empty)
|
- `reasoning_effort` (only when non-empty)
|
||||||
- `response_format` (only when structured output is provided)
|
- `response_format` (only when structured output is provided)
|
||||||
@@ -142,6 +142,7 @@ Base timeout comes from client configuration.
|
|||||||
Per-request override:
|
Per-request override:
|
||||||
|
|
||||||
- if `Target.TimeoutSeconds > 0`, use that value for request timeout
|
- if `Target.TimeoutSeconds > 0`, use that value for request timeout
|
||||||
|
- if `Target.TimeoutSeconds == 0` and the value came from an explicit request override, disable the HTTP client timeout
|
||||||
- if `Target.TimeoutSeconds < 0`, request is rejected (`ErrInvalidRequest`)
|
- if `Target.TimeoutSeconds < 0`, request is rejected (`ErrInvalidRequest`)
|
||||||
|
|
||||||
## Response Expectations
|
## Response Expectations
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
||||||
|
"gitea.maximumdirect.net/eric/scriptorium/internal/llm"
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/profile"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/profile"
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/promptdef"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/promptdef"
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/usecase"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/usecase"
|
||||||
@@ -32,6 +33,34 @@ func (f *fakeRunner) Run(ctx context.Context, req domain.RunRequest) (*domain.Ru
|
|||||||
return f.result, nil
|
return f.result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type handlerPromptRepo struct {
|
||||||
|
def *domain.PromptDefinition
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r handlerPromptRepo) GetPromptDefinition(ctx context.Context, id string, version string) (*domain.PromptDefinition, error) {
|
||||||
|
return r.def, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type handlerProfileRepo struct {
|
||||||
|
profile *domain.ExecutionProfile
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r handlerProfileRepo) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
||||||
|
return r.profile, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type handlerArtifactReader struct{}
|
||||||
|
|
||||||
|
func (handlerArtifactReader) Read(ctx context.Context, ref domain.ArtifactRef) (*domain.Artifact, error) {
|
||||||
|
return &domain.Artifact{Name: "input", Body: []byte("input"), Hash: "hash"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type handlerRenderer struct{}
|
||||||
|
|
||||||
|
func (handlerRenderer) Render(ctx context.Context, definition *domain.PromptDefinition, inputs map[string]*domain.Artifact, vars map[string]string) (*domain.RenderedPrompt, error) {
|
||||||
|
return &domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, nil
|
||||||
|
}
|
||||||
|
|
||||||
func TestHandlerPostRunsSuccessWithExplicitProfileID(t *testing.T) {
|
func TestHandlerPostRunsSuccessWithExplicitProfileID(t *testing.T) {
|
||||||
start := time.Now().UTC()
|
start := time.Now().UTC()
|
||||||
end := start.Add(2 * time.Second)
|
end := start.Add(2 * time.Second)
|
||||||
@@ -477,6 +506,54 @@ func TestHandlerMissingPromptID(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHandlerReservedExtraParamsThroughRunnerMapsToInvalidRequest(t *testing.T) {
|
||||||
|
llmClient, err := llm.NewOpenAICompatibleClient(llm.OpenAICompatibleConfig{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
runner := usecase.NewRunner(
|
||||||
|
handlerPromptRepo{def: &domain.PromptDefinition{
|
||||||
|
ID: "p",
|
||||||
|
Version: "1",
|
||||||
|
DefaultProfile: "exec",
|
||||||
|
Templates: []domain.PromptMessageTemplate{{Role: "user", Content: "hi"}},
|
||||||
|
OutputFormat: domain.FormatText,
|
||||||
|
Validation: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone},
|
||||||
|
}},
|
||||||
|
handlerProfileRepo{profile: &domain.ExecutionProfile{
|
||||||
|
ID: "exec",
|
||||||
|
Endpoint: "http://example.invalid/v1",
|
||||||
|
Model: "model",
|
||||||
|
}},
|
||||||
|
handlerArtifactReader{},
|
||||||
|
handlerRenderer{},
|
||||||
|
llmClient,
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
h := NewHandler(runner)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString(`{
|
||||||
|
"prompt_id":"p",
|
||||||
|
"inputs":{"x":{"type":"file","uri":"a"}},
|
||||||
|
"model":{"extra_params":{"model":"collision"}}
|
||||||
|
}`))
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
|
||||||
|
h.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
if w.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("expected 400, got %d body=%s", w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
var resp map[string]any
|
||||||
|
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("invalid JSON response: %v", err)
|
||||||
|
}
|
||||||
|
errBody := resp["error"].(map[string]any)
|
||||||
|
if errBody["code"] != "invalid_request" {
|
||||||
|
t.Fatalf("expected invalid_request code, got %#v", errBody["code"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestHandlerUsecaseErrorMapping(t *testing.T) {
|
func TestHandlerUsecaseErrorMapping(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
@@ -95,20 +95,21 @@ 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 `json:"prompt_id"`
|
||||||
PromptVersion string `json:"prompt_version,omitempty"`
|
PromptVersion string `json:"prompt_version,omitempty"`
|
||||||
PromptHash string `json:"prompt_hash,omitempty"`
|
PromptHash string `json:"prompt_hash,omitempty"`
|
||||||
SelectedProfileID string `json:"selected_profile_id"`
|
SelectedProfileID string `json:"selected_profile_id"`
|
||||||
EffectiveModelParams ExecutionTarget `json:"effective_model_params"`
|
EffectiveModelParams ExecutionTarget `json:"effective_model_params"`
|
||||||
OutputContract OutputContract `json:"output_contract"`
|
TargetPresence ExecutionTargetPresence `json:"-"`
|
||||||
StructuredOutput *StructuredOutputSpec `json:"structured_output,omitempty"`
|
OutputContract OutputContract `json:"output_contract"`
|
||||||
InputHashes map[string]string `json:"input_hashes,omitempty"`
|
StructuredOutput *StructuredOutputSpec `json:"structured_output,omitempty"`
|
||||||
SessionID string `json:"session_id,omitempty"`
|
InputHashes map[string]string `json:"input_hashes,omitempty"`
|
||||||
RenderedPromptHash string `json:"rendered_prompt_hash"`
|
SessionID string `json:"session_id,omitempty"`
|
||||||
Messages []RenderedMessage `json:"messages"`
|
RenderedPromptHash string `json:"rendered_prompt_hash"`
|
||||||
StartTime time.Time `json:"start_time,omitempty"`
|
Messages []RenderedMessage `json:"messages"`
|
||||||
EndTime time.Time `json:"end_time,omitempty"`
|
StartTime time.Time `json:"start_time,omitempty"`
|
||||||
DurationMS int64 `json:"duration_ms,omitempty"`
|
EndTime time.Time `json:"end_time,omitempty"`
|
||||||
|
DurationMS int64 `json:"duration_ms,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// ArtifactRef represents a reference to an input artifact.
|
// ArtifactRef represents a reference to an input artifact.
|
||||||
@@ -186,6 +187,15 @@ type ExecutionTargetOverride struct {
|
|||||||
ExtraParams map[string]any `json:"extra_params,omitempty"`
|
ExtraParams map[string]any `json:"extra_params,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ExecutionTargetPresence tracks which effective runtime fields came from an
|
||||||
|
// explicit request override even when the resolved value is a zero value.
|
||||||
|
type ExecutionTargetPresence struct {
|
||||||
|
Temperature bool
|
||||||
|
MaxTokens bool
|
||||||
|
TopP bool
|
||||||
|
TimeoutSeconds bool
|
||||||
|
}
|
||||||
|
|
||||||
// ExecutionTarget represents effective model runtime settings for a run.
|
// ExecutionTarget represents effective model runtime settings for a run.
|
||||||
type ExecutionTarget struct {
|
type ExecutionTarget struct {
|
||||||
Endpoint string `yaml:"endpoint" json:"endpoint"`
|
Endpoint string `yaml:"endpoint" json:"endpoint"`
|
||||||
@@ -225,6 +235,7 @@ type RenderedMessage struct {
|
|||||||
type GenerateRequest struct {
|
type GenerateRequest struct {
|
||||||
Prompt RenderedPrompt
|
Prompt RenderedPrompt
|
||||||
Target ExecutionTarget
|
Target ExecutionTarget
|
||||||
|
TargetPresence ExecutionTargetPresence
|
||||||
StructuredOutput *StructuredOutputSpec
|
StructuredOutput *StructuredOutputSpec
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -116,6 +116,8 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
|
|||||||
effectiveTimeout := c.timeout
|
effectiveTimeout := c.timeout
|
||||||
if req.Target.TimeoutSeconds > 0 {
|
if req.Target.TimeoutSeconds > 0 {
|
||||||
effectiveTimeout = time.Duration(req.Target.TimeoutSeconds) * time.Second
|
effectiveTimeout = time.Duration(req.Target.TimeoutSeconds) * time.Second
|
||||||
|
} else if req.TargetPresence.TimeoutSeconds {
|
||||||
|
effectiveTimeout = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
httpClient := c.httpClient
|
httpClient := c.httpClient
|
||||||
@@ -187,13 +189,13 @@ func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultMod
|
|||||||
wireReq.Messages = append(wireReq.Messages, openAIChatRequestMessageFromRenderedMessage(msg))
|
wireReq.Messages = append(wireReq.Messages, openAIChatRequestMessageFromRenderedMessage(msg))
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Target.Temperature != 0 {
|
if req.Target.Temperature != 0 || req.TargetPresence.Temperature {
|
||||||
wireReq.Temperature = &req.Target.Temperature
|
wireReq.Temperature = &req.Target.Temperature
|
||||||
}
|
}
|
||||||
if req.Target.MaxTokens != 0 {
|
if req.Target.MaxTokens != 0 || req.TargetPresence.MaxTokens {
|
||||||
wireReq.MaxTokens = &req.Target.MaxTokens
|
wireReq.MaxTokens = &req.Target.MaxTokens
|
||||||
}
|
}
|
||||||
if req.Target.TopP != 0 {
|
if req.Target.TopP != 0 || req.TargetPresence.TopP {
|
||||||
wireReq.TopP = &req.Target.TopP
|
wireReq.TopP = &req.Target.TopP
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(req.Target.ServiceTier) != "" {
|
if strings.TrimSpace(req.Target.ServiceTier) != "" {
|
||||||
|
|||||||
@@ -507,6 +507,127 @@ func TestOpenAICompatibleClientOmitsReasoningEffortWhenUnset(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatibleClientSerializesExplicitZeroNumericOverrides(t *testing.T) {
|
||||||
|
var observedBody map[string]any
|
||||||
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
defer r.Body.Close()
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil {
|
||||||
|
t.Fatalf("failed to decode request body: %v", err)
|
||||||
|
}
|
||||||
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
||||||
|
}))
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||||
|
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||||
|
Target: domain.ExecutionTarget{Model: "model"},
|
||||||
|
TargetPresence: domain.ExecutionTargetPresence{
|
||||||
|
Temperature: true,
|
||||||
|
MaxTokens: true,
|
||||||
|
TopP: true,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
if observedBody["temperature"] != float64(0) {
|
||||||
|
t.Fatalf("expected explicit zero temperature, got %#v", observedBody["temperature"])
|
||||||
|
}
|
||||||
|
if observedBody["max_tokens"] != float64(0) {
|
||||||
|
t.Fatalf("expected explicit zero max_tokens, got %#v", observedBody["max_tokens"])
|
||||||
|
}
|
||||||
|
if observedBody["top_p"] != float64(0) {
|
||||||
|
t.Fatalf("expected explicit zero top_p, got %#v", observedBody["top_p"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatibleClientOmitsImplicitZeroNumericFields(t *testing.T) {
|
||||||
|
var observedBody map[string]any
|
||||||
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
defer r.Body.Close()
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil {
|
||||||
|
t.Fatalf("failed to decode request body: %v", err)
|
||||||
|
}
|
||||||
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
||||||
|
}))
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||||
|
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||||
|
Target: domain.ExecutionTarget{Model: "model"},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
for _, field := range []string{"temperature", "max_tokens", "top_p"} {
|
||||||
|
if _, exists := observedBody[field]; exists {
|
||||||
|
t.Fatalf("expected implicit zero field %q to be omitted, got body %#v", field, observedBody)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatibleClientExplicitZeroTimeoutDisablesClientTimeout(t *testing.T) {
|
||||||
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
||||||
|
}))
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||||
|
BaseURL: ts.URL + "/v1",
|
||||||
|
Timeout: time.Nanosecond,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||||
|
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||||
|
Target: domain.ExecutionTarget{Model: "model", TimeoutSeconds: 0},
|
||||||
|
TargetPresence: domain.ExecutionTargetPresence{TimeoutSeconds: true},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected explicit zero timeout to disable client timeout, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatibleClientOmittedTimeoutUsesClientTimeout(t *testing.T) {
|
||||||
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
||||||
|
}))
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||||
|
BaseURL: ts.URL + "/v1",
|
||||||
|
Timeout: time.Nanosecond,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||||
|
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||||
|
Target: domain.ExecutionTarget{Model: "model", TimeoutSeconds: 0},
|
||||||
|
})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected omitted timeout to use client timeout")
|
||||||
|
}
|
||||||
|
if !errors.Is(err, ErrRequestFailed) {
|
||||||
|
t.Fatalf("expected ErrRequestFailed, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestOpenAICompatibleClientRejectsInvalidExtraParamsBeforeProviderCall(t *testing.T) {
|
func TestOpenAICompatibleClientRejectsInvalidExtraParamsBeforeProviderCall(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
@@ -92,9 +92,13 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
|
|||||||
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{
|
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{
|
||||||
Prompt: domain.RenderedPrompt{SessionID: prepared.SessionID, Messages: prepared.Messages},
|
Prompt: domain.RenderedPrompt{SessionID: prepared.SessionID, Messages: prepared.Messages},
|
||||||
Target: prepared.EffectiveModelParams,
|
Target: prepared.EffectiveModelParams,
|
||||||
|
TargetPresence: prepared.TargetPresence,
|
||||||
StructuredOutput: prepared.StructuredOutput,
|
StructuredOutput: prepared.StructuredOutput,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if errors.Is(err, llm.ErrInvalidRequest) {
|
||||||
|
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
|
||||||
|
}
|
||||||
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
|
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -187,7 +191,7 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr
|
|||||||
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
|
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
effectiveModel, err := resolveExecutionTarget(execProfile, req.Execution)
|
effectiveModel, targetPresence, err := resolveExecutionTarget(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)
|
||||||
}
|
}
|
||||||
@@ -233,6 +237,7 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr
|
|||||||
PromptHash: promptDefinitionHash,
|
PromptHash: promptDefinitionHash,
|
||||||
SelectedProfileID: selectedProfileID,
|
SelectedProfileID: selectedProfileID,
|
||||||
EffectiveModelParams: effectiveModel,
|
EffectiveModelParams: effectiveModel,
|
||||||
|
TargetPresence: targetPresence,
|
||||||
OutputContract: effectiveContract,
|
OutputContract: effectiveContract,
|
||||||
StructuredOutput: structuredOutput,
|
StructuredOutput: structuredOutput,
|
||||||
InputHashes: inputHashes,
|
InputHashes: inputHashes,
|
||||||
@@ -363,8 +368,9 @@ func mergeExecutionTarget(base domain.ExecutionTarget, override domain.Execution
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.ExecutionTargetOverride) (domain.ExecutionTarget, error) {
|
func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence, error) {
|
||||||
out := base
|
out := base
|
||||||
|
var presence domain.ExecutionTargetPresence
|
||||||
if override.Endpoint != "" {
|
if override.Endpoint != "" {
|
||||||
out.Endpoint = override.Endpoint
|
out.Endpoint = override.Endpoint
|
||||||
}
|
}
|
||||||
@@ -373,27 +379,31 @@ func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.E
|
|||||||
}
|
}
|
||||||
if override.Temperature != nil {
|
if override.Temperature != nil {
|
||||||
if *override.Temperature < 0 || *override.Temperature > 2 {
|
if *override.Temperature < 0 || *override.Temperature > 2 {
|
||||||
return domain.ExecutionTarget{}, errors.New("temperature must be between 0 and 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
|
||||||
}
|
}
|
||||||
if override.MaxTokens != nil {
|
if override.MaxTokens != nil {
|
||||||
if *override.MaxTokens < 0 {
|
if *override.MaxTokens < 0 {
|
||||||
return domain.ExecutionTarget{}, errors.New("max_tokens must be greater than or equal to 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
|
||||||
}
|
}
|
||||||
if override.TopP != nil {
|
if override.TopP != nil {
|
||||||
if *override.TopP < 0 || *override.TopP > 1 {
|
if *override.TopP < 0 || *override.TopP > 1 {
|
||||||
return domain.ExecutionTarget{}, errors.New("top_p must be between 0 and 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
|
||||||
}
|
}
|
||||||
if override.TimeoutSeconds != nil {
|
if override.TimeoutSeconds != nil {
|
||||||
if *override.TimeoutSeconds < 0 {
|
if *override.TimeoutSeconds < 0 {
|
||||||
return domain.ExecutionTarget{}, errors.New("timeout_seconds must be greater than or equal to 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
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(override.ServiceTier) != "" {
|
if strings.TrimSpace(override.ServiceTier) != "" {
|
||||||
out.ServiceTier = override.ServiceTier
|
out.ServiceTier = override.ServiceTier
|
||||||
@@ -407,20 +417,21 @@ 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, nil
|
return out, presence, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func resolveExecutionTarget(profileValue *domain.ExecutionProfile, override *domain.ExecutionTargetOverride) (domain.ExecutionTarget, error) {
|
func resolveExecutionTarget(profileValue *domain.ExecutionProfile, override *domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence, error) {
|
||||||
out := defaults.ExecutionTargetDefault()
|
out := defaults.ExecutionTargetDefault()
|
||||||
out = mergeExecutionTarget(out, executionProfileToTarget(profileValue))
|
out = mergeExecutionTarget(out, executionProfileToTarget(profileValue))
|
||||||
|
var presence domain.ExecutionTargetPresence
|
||||||
if override != nil {
|
if override != nil {
|
||||||
var err error
|
var err error
|
||||||
out, err = mergeExecutionTargetOverride(out, *override)
|
out, presence, err = mergeExecutionTargetOverride(out, *override)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return domain.ExecutionTarget{}, err
|
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return out, nil
|
return out, presence, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func validateAPIKeyEnv(apiKeyEnv string) error {
|
func validateAPIKeyEnv(apiKeyEnv string) error {
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/defaults"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/defaults"
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
||||||
|
"gitea.maximumdirect.net/eric/scriptorium/internal/llm"
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/profile"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/profile"
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/prompt"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/prompt"
|
||||||
"gitea.maximumdirect.net/eric/scriptorium/internal/promptdef"
|
"gitea.maximumdirect.net/eric/scriptorium/internal/promptdef"
|
||||||
@@ -302,6 +303,7 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
|||||||
wantMaxTokens int
|
wantMaxTokens int
|
||||||
wantTopP float64
|
wantTopP float64
|
||||||
wantTimeoutSecs int
|
wantTimeoutSecs int
|
||||||
|
wantPresence domain.ExecutionTargetPresence
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "omitted preserves profile values",
|
name: "omitted preserves profile values",
|
||||||
@@ -318,6 +320,7 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
|||||||
wantMaxTokens: 321,
|
wantMaxTokens: 321,
|
||||||
wantTopP: 0.8,
|
wantTopP: 0.8,
|
||||||
wantTimeoutSecs: 45,
|
wantTimeoutSecs: 45,
|
||||||
|
wantPresence: domain.ExecutionTargetPresence{Temperature: true},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "explicit zero max tokens",
|
name: "explicit zero max tokens",
|
||||||
@@ -326,6 +329,7 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
|||||||
wantMaxTokens: 0,
|
wantMaxTokens: 0,
|
||||||
wantTopP: 0.8,
|
wantTopP: 0.8,
|
||||||
wantTimeoutSecs: 45,
|
wantTimeoutSecs: 45,
|
||||||
|
wantPresence: domain.ExecutionTargetPresence{MaxTokens: true},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "explicit zero top p",
|
name: "explicit zero top p",
|
||||||
@@ -334,6 +338,7 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
|||||||
wantMaxTokens: 321,
|
wantMaxTokens: 321,
|
||||||
wantTopP: 0,
|
wantTopP: 0,
|
||||||
wantTimeoutSecs: 45,
|
wantTimeoutSecs: 45,
|
||||||
|
wantPresence: domain.ExecutionTargetPresence{TopP: true},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "explicit zero timeout",
|
name: "explicit zero timeout",
|
||||||
@@ -342,6 +347,7 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
|||||||
wantMaxTokens: 321,
|
wantMaxTokens: 321,
|
||||||
wantTopP: 0.8,
|
wantTopP: 0.8,
|
||||||
wantTimeoutSecs: 0,
|
wantTimeoutSecs: 0,
|
||||||
|
wantPresence: domain.ExecutionTargetPresence{TimeoutSeconds: true},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -382,6 +388,9 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
|||||||
got.TimeoutSeconds != tc.wantTimeoutSecs {
|
got.TimeoutSeconds != tc.wantTimeoutSecs {
|
||||||
t.Fatalf("unexpected effective numeric settings: %+v", got)
|
t.Fatalf("unexpected effective numeric settings: %+v", got)
|
||||||
}
|
}
|
||||||
|
if prepared.TargetPresence != tc.wantPresence {
|
||||||
|
t.Fatalf("unexpected target presence: got %+v want %+v", prepared.TargetPresence, tc.wantPresence)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -857,6 +866,9 @@ func TestRunnerRunSuccessful(t *testing.T) {
|
|||||||
if llmClient.lastReq.Target.TimeoutSeconds != 90 {
|
if llmClient.lastReq.Target.TimeoutSeconds != 90 {
|
||||||
t.Fatalf("expected timeout propagation, got %d", llmClient.lastReq.Target.TimeoutSeconds)
|
t.Fatalf("expected timeout propagation, got %d", llmClient.lastReq.Target.TimeoutSeconds)
|
||||||
}
|
}
|
||||||
|
if !llmClient.lastReq.TargetPresence.Temperature || !llmClient.lastReq.TargetPresence.TimeoutSeconds {
|
||||||
|
t.Fatalf("expected numeric override presence to be sent to llm, got %+v", llmClient.lastReq.TargetPresence)
|
||||||
|
}
|
||||||
if llmClient.lastReq.Prompt.SessionID != "session-123" {
|
if llmClient.lastReq.Prompt.SessionID != "session-123" {
|
||||||
t.Fatalf("expected session id to be sent to llm, got %q", llmClient.lastReq.Prompt.SessionID)
|
t.Fatalf("expected session id to be sent to llm, got %q", llmClient.lastReq.Prompt.SessionID)
|
||||||
}
|
}
|
||||||
@@ -1305,6 +1317,28 @@ func TestRunnerRunLLMFailure(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRunnerRunLLMInvalidRequestMapsToUsecaseInvalidRequest(t *testing.T) {
|
||||||
|
runner := NewRunner(
|
||||||
|
&fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)},
|
||||||
|
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||||
|
defaultArtifactReader(),
|
||||||
|
defaultRenderer(),
|
||||||
|
&fakeLLM{err: llm.ErrInvalidRequest},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
_, err := runner.Run(context.Background(), domain.RunRequest{
|
||||||
|
PromptID: "p",
|
||||||
|
ProfileID: "exec",
|
||||||
|
Inputs: singleInputRef(),
|
||||||
|
})
|
||||||
|
if !errors.Is(err, ErrInvalidRequest) {
|
||||||
|
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||||
|
}
|
||||||
|
if errors.Is(err, ErrLLMGenerate) {
|
||||||
|
t.Fatalf("did not expect ErrLLMGenerate, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestRunnerRunValidationStillWorks(t *testing.T) {
|
func TestRunnerRunValidationStillWorks(t *testing.T) {
|
||||||
validator := &fakeValidator{result: domain.ValidationResult{Status: domain.ValidationFailed, Mode: domain.ValidationBasic, Errors: []string{"bad"}, IsValid: false}}
|
validator := &fakeValidator{result: domain.ValidationResult{Status: domain.ValidationFailed, Mode: domain.ValidationBasic, Errors: []string{"bad"}, IsValid: false}}
|
||||||
runner := NewRunner(
|
runner := NewRunner(
|
||||||
@@ -1478,10 +1512,13 @@ func TestResolveExecutionTargetProfileValuesPopulateAllSupportedFields(t *testin
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
target, err := resolveExecutionTarget(profileValue, nil)
|
target, presence, err := resolveExecutionTarget(profileValue, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("expected no error, got %v", err)
|
t.Fatalf("expected no error, got %v", err)
|
||||||
}
|
}
|
||||||
|
if presence != (domain.ExecutionTargetPresence{}) {
|
||||||
|
t.Fatalf("expected no request override presence, got %+v", presence)
|
||||||
|
}
|
||||||
if target.Endpoint != profileValue.Endpoint ||
|
if target.Endpoint != profileValue.Endpoint ||
|
||||||
target.Model != profileValue.Model ||
|
target.Model != profileValue.Model ||
|
||||||
target.Temperature != profileValue.Temperature ||
|
target.Temperature != profileValue.Temperature ||
|
||||||
@@ -1529,10 +1566,13 @@ func TestResolveExecutionTargetRuntimeOverridesBeatProfileForAllOverrideableFiel
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
target, err := resolveExecutionTarget(profileValue, override)
|
target, presence, err := resolveExecutionTarget(profileValue, override)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("expected no error, got %v", err)
|
t.Fatalf("expected no error, got %v", err)
|
||||||
}
|
}
|
||||||
|
if presence != (domain.ExecutionTargetPresence{Temperature: true, MaxTokens: true, TopP: true, TimeoutSeconds: true}) {
|
||||||
|
t.Fatalf("unexpected override presence: %+v", presence)
|
||||||
|
}
|
||||||
if target.Endpoint != override.Endpoint ||
|
if target.Endpoint != override.Endpoint ||
|
||||||
target.Model != override.Model ||
|
target.Model != override.Model ||
|
||||||
target.Temperature != *override.Temperature ||
|
target.Temperature != *override.Temperature ||
|
||||||
|
|||||||
Reference in New Issue
Block a user