Centralize execution setting and session validation
This commit is contained in:
@@ -6,9 +6,11 @@ import (
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -502,6 +504,35 @@ func TestRunnerDirectSessionResolution(t *testing.T) {
|
||||
t.Fatalf("invalid direct session invoked generation %d times", llmClient.calls)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("malformed direct value fails before loading or generation", func(t *testing.T) {
|
||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "unexpected"}}
|
||||
runner := NewRunner(
|
||||
promptRepo,
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
nil, nil)
|
||||
|
||||
_, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
SessionID: "session" + string([]byte{0xff}),
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
if promptRepo.lastID != "" {
|
||||
t.Fatalf("invalid direct session loaded prompt %q", promptRepo.lastID)
|
||||
}
|
||||
if llmClient.calls != 0 {
|
||||
t.Fatalf("invalid direct session invoked generation %d times", llmClient.calls)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestRunnerPrepareUsesPromptDefaultProfileWhenNoExplicitProfileID(t *testing.T) {
|
||||
@@ -704,16 +735,23 @@ func TestRunnerPrepareRequestNumericOverridePresence(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRunnerPrepareInvalidRequestNumericOverridesFail(t *testing.T) {
|
||||
tests := []struct {
|
||||
type testCase struct {
|
||||
name string
|
||||
override *domain.ExecutionTargetOverride
|
||||
}{
|
||||
}
|
||||
tests := []testCase{
|
||||
{name: "temperature below range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(-0.1)}},
|
||||
{name: "temperature above range", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(2.1)}},
|
||||
{name: "max tokens below range", override: &domain.ExecutionTargetOverride{MaxTokens: intPtr(-1)}},
|
||||
{name: "top p below range", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(-0.1)}},
|
||||
{name: "top p above range", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(1.1)}},
|
||||
{name: "timeout below range", override: &domain.ExecutionTargetOverride{TimeoutSeconds: intPtr(-1)}},
|
||||
{name: "temperature is not finite", override: &domain.ExecutionTargetOverride{Temperature: float64Ptr(math.NaN())}},
|
||||
{name: "top p is not finite", override: &domain.ExecutionTargetOverride{TopP: float64Ptr(math.Inf(1))}},
|
||||
}
|
||||
if strconv.IntSize == 64 {
|
||||
durationLimit := int64(math.MaxInt64 / int64(time.Second))
|
||||
tests = append(tests, testCase{name: "timeout cannot be represented as a duration", override: &domain.ExecutionTargetOverride{TimeoutSeconds: intPtr(int(durationLimit) + 1)}})
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
@@ -2300,10 +2338,7 @@ func TestResolveExecutionTargetProfileValuesPopulateAllSupportedFields(t *testin
|
||||
},
|
||||
}
|
||||
|
||||
target, presence, err := resolveExecutionTarget(nil, profileValue, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
target, presence := resolveExecutionTarget(nil, profileValue, nil)
|
||||
if presence != (domain.ExecutionTargetPresence{}) {
|
||||
t.Fatalf("expected no request override presence, got %+v", presence)
|
||||
}
|
||||
@@ -2355,10 +2390,7 @@ func TestResolveExecutionTargetRuntimeOverridesBeatProfileForAllOverrideableFiel
|
||||
},
|
||||
}
|
||||
|
||||
target, presence, err := resolveExecutionTarget(nil, profileValue, override)
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
target, presence := resolveExecutionTarget(nil, profileValue, override)
|
||||
if presence != (domain.ExecutionTargetPresence{Temperature: true, MaxTokens: true, TopP: true, TimeoutSeconds: true}) {
|
||||
t.Fatalf("unexpected override presence: %+v", presence)
|
||||
}
|
||||
@@ -2405,12 +2437,9 @@ func TestResolveExecutionTargetReasoningOverrideStates(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
target, _, err := resolveExecutionTarget(nil, profileValue, &domain.ExecutionTargetOverride{
|
||||
target, _ := resolveExecutionTarget(nil, profileValue, &domain.ExecutionTargetOverride{
|
||||
ReasoningEffort: tt.override,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("resolve execution target: %v", err)
|
||||
}
|
||||
if target.ReasoningEffort != tt.want {
|
||||
t.Fatalf("reasoning effort = %q, want %q", target.ReasoningEffort, tt.want)
|
||||
}
|
||||
@@ -2539,10 +2568,7 @@ func TestResolveExecutionTargetUsesBackendProfileAndRequestPrecedence(t *testing
|
||||
ExtraParams: map[string]any{"request": true},
|
||||
}
|
||||
|
||||
target, _, err := resolveExecutionTarget(backendValue, profileValue, override)
|
||||
if err != nil {
|
||||
t.Fatalf("resolve target: %v", err)
|
||||
}
|
||||
target, _ := resolveExecutionTarget(backendValue, profileValue, override)
|
||||
if target.BackendID != "custom" {
|
||||
t.Fatalf("endpoint override changed backend identity: %+v", target)
|
||||
}
|
||||
@@ -2553,12 +2579,9 @@ func TestResolveExecutionTargetUsesBackendProfileAndRequestPrecedence(t *testing
|
||||
t.Fatalf("expected whole-map request replacement, got %#v", target.ExtraParams)
|
||||
}
|
||||
|
||||
target, _, err = resolveExecutionTarget(backendValue, &domain.ExecutionProfile{
|
||||
target, _ = resolveExecutionTarget(backendValue, &domain.ExecutionProfile{
|
||||
ID: "exec", BackendID: "custom", Model: "profile-model",
|
||||
}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("resolve backend defaults: %v", err)
|
||||
}
|
||||
if target.Endpoint != backendValue.Endpoint ||
|
||||
target.APIKeyEnv != backendValue.APIKeyEnv ||
|
||||
!reflect.DeepEqual(target.ExtraParams, backendValue.ExtraParams) {
|
||||
|
||||
Reference in New Issue
Block a user