Refine execution target mapping helpers and coverage across usecase, HTTP, and LLM

This commit is contained in:
2026-05-26 13:07:35 +00:00
parent 75fa0a030a
commit 79901fbb86
5 changed files with 428 additions and 63 deletions

View File

@@ -63,18 +63,7 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
var model *domain.ExecutionTarget
if req.Model != nil {
model = &domain.ExecutionTarget{
Endpoint: req.Model.Endpoint,
Model: req.Model.Model,
Temperature: req.Model.Temperature,
MaxTokens: req.Model.MaxTokens,
TopP: req.Model.TopP,
TimeoutSeconds: req.Model.TimeoutSeconds,
ServiceTier: req.Model.ServiceTier,
ReasoningEffort: req.Model.ReasoningEffort,
APIKeyEnv: req.Model.APIKeyEnv,
ExtraParams: req.Model.ExtraParams,
}
model = executionTargetFromModelOverrideDTO(req.Model)
}
res, err := h.runner.Run(r.Context(), domain.RunRequest{
@@ -110,19 +99,8 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
SelectedProfileID: res.SelectedProfileID,
ModelName: res.ModelName,
Endpoint: res.Endpoint,
ModelParams: modelParamsDTO{
Endpoint: res.EffectiveModelParams.Endpoint,
Model: res.EffectiveModelParams.Model,
Temperature: res.EffectiveModelParams.Temperature,
MaxTokens: res.EffectiveModelParams.MaxTokens,
TopP: res.EffectiveModelParams.TopP,
TimeoutSeconds: res.EffectiveModelParams.TimeoutSeconds,
ServiceTier: res.EffectiveModelParams.ServiceTier,
ReasoningEffort: res.EffectiveModelParams.ReasoningEffort,
APIKeyEnv: res.EffectiveModelParams.APIKeyEnv,
ExtraParams: res.EffectiveModelParams.ExtraParams,
},
InputHashes: res.InputHashes,
ModelParams: modelParamsDTOFromExecutionTarget(res.EffectiveModelParams),
InputHashes: res.InputHashes,
Usage: tokenUsageDTO{
PromptTokens: res.Usage.PromptTokens,
CompletionTokens: res.Usage.CompletionTokens,
@@ -143,6 +121,39 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
writeJSON(w, http.StatusOK, resp)
}
func executionTargetFromModelOverrideDTO(dto *modelOverrideRequestDTO) *domain.ExecutionTarget {
if dto == nil {
return nil
}
return &domain.ExecutionTarget{
Endpoint: dto.Endpoint,
Model: dto.Model,
Temperature: dto.Temperature,
MaxTokens: dto.MaxTokens,
TopP: dto.TopP,
TimeoutSeconds: dto.TimeoutSeconds,
ServiceTier: dto.ServiceTier,
ReasoningEffort: dto.ReasoningEffort,
APIKeyEnv: dto.APIKeyEnv,
ExtraParams: dto.ExtraParams,
}
}
func modelParamsDTOFromExecutionTarget(target domain.ExecutionTarget) modelParamsDTO {
return modelParamsDTO{
Endpoint: target.Endpoint,
Model: target.Model,
Temperature: target.Temperature,
MaxTokens: target.MaxTokens,
TopP: target.TopP,
TimeoutSeconds: target.TimeoutSeconds,
ServiceTier: target.ServiceTier,
ReasoningEffort: target.ReasoningEffort,
APIKeyEnv: target.APIKeyEnv,
ExtraParams: target.ExtraParams,
}
}
func mapValidation(v domain.ValidationResult) validationDTO {
return validationDTO{
Status: string(v.Status),

View File

@@ -8,6 +8,7 @@ import (
"fmt"
"net/http"
"net/http/httptest"
"reflect"
"strings"
"testing"
"time"
@@ -173,6 +174,136 @@ func TestHandlerPostRunsSuccessUsingPromptDefaultProfile(t *testing.T) {
}
}
func TestHandlerModelOverrideMapsAllSupportedExecutionFields(t *testing.T) {
r := &fakeRunner{result: &domain.RunResult{
Artifact: domain.Artifact{Body: []byte("ok")},
Validation: domain.ValidationResult{Status: domain.ValidationPassed, Mode: domain.ValidationBasic, IsValid: true},
EffectiveModelParams: domain.ExecutionTarget{Endpoint: "http://llm/v1", Model: "m1"},
}}
h := NewHandler(r)
reqBody := `{
"prompt_id": "prompt-1",
"inputs": {"transcript": {"type": "file", "uri": "./t.md"}},
"model": {
"endpoint": "http://override/v1",
"model": "override-model",
"temperature": 0.6,
"max_tokens": 250,
"top_p": 0.85,
"timeout_seconds": 33,
"service_tier": "flex",
"reasoning_effort": "medium",
"api_key_env": "SCRIPTORIUM_API_KEY",
"extra_params": {"provider_option":"on"}
}
}`
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString(reqBody))
w := httptest.NewRecorder()
h.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200, got %d body=%s", w.Code, w.Body.String())
}
if r.last.Execution == nil {
t.Fatalf("expected execution override in run request")
}
got := r.last.Execution
if got.Endpoint != "http://override/v1" ||
got.Model != "override-model" ||
got.Temperature != 0.6 ||
got.MaxTokens != 250 ||
got.TopP != 0.85 ||
got.TimeoutSeconds != 33 ||
got.ServiceTier != "flex" ||
got.ReasoningEffort != "medium" ||
got.APIKeyEnv != "SCRIPTORIUM_API_KEY" {
t.Fatalf("unexpected mapped execution target: %+v", got)
}
if !reflect.DeepEqual(got.ExtraParams, map[string]string{"provider_option": "on"}) {
t.Fatalf("unexpected mapped extra_params: %#v", got.ExtraParams)
}
}
func TestHandlerResponseMetadataModelParamsIncludesAllSupportedFields(t *testing.T) {
r := &fakeRunner{result: &domain.RunResult{
Artifact: domain.Artifact{
Name: "output",
ContentType: "text/plain",
Body: []byte("ok"),
Size: 2,
Hash: "abc",
},
Validation: domain.ValidationResult{Status: domain.ValidationPassed, Mode: domain.ValidationBasic, IsValid: true},
EffectiveModelParams: domain.ExecutionTarget{
Endpoint: "http://llm/v1",
Model: "gpt-test",
Temperature: 0.4,
MaxTokens: 321,
TopP: 0.7,
TimeoutSeconds: 45,
ServiceTier: "priority",
ReasoningEffort: "high",
APIKeyEnv: "SCRIPTORIUM_API_KEY",
ExtraParams: map[string]string{
"provider_option": "on",
},
},
}}
h := NewHandler(r)
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString(`{"prompt_id":"p","inputs":{"x":{"type":"file","uri":"a"}}}`))
w := httptest.NewRecorder()
h.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200, 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)
}
metadata := resp["metadata"].(map[string]any)
params := metadata["model_params"].(map[string]any)
if params["endpoint"] != "http://llm/v1" {
t.Fatalf("unexpected endpoint: %#v", params["endpoint"])
}
if params["model"] != "gpt-test" {
t.Fatalf("unexpected model: %#v", params["model"])
}
if params["temperature"] != 0.4 {
t.Fatalf("unexpected temperature: %#v", params["temperature"])
}
if params["max_tokens"] != float64(321) {
t.Fatalf("unexpected max_tokens: %#v", params["max_tokens"])
}
if params["top_p"] != 0.7 {
t.Fatalf("unexpected top_p: %#v", params["top_p"])
}
if params["timeout_seconds"] != float64(45) {
t.Fatalf("unexpected timeout_seconds: %#v", params["timeout_seconds"])
}
if params["service_tier"] != "priority" {
t.Fatalf("unexpected service_tier: %#v", params["service_tier"])
}
if params["reasoning_effort"] != "high" {
t.Fatalf("unexpected reasoning_effort: %#v", params["reasoning_effort"])
}
if params["api_key_env"] != "SCRIPTORIUM_API_KEY" {
t.Fatalf("unexpected api_key_env: %#v", params["api_key_env"])
}
extraParams, ok := params["extra_params"].(map[string]any)
if !ok {
t.Fatalf("expected extra_params object, got %#v", params["extra_params"])
}
if extraParams["provider_option"] != "on" {
t.Fatalf("unexpected extra_params.provider_option: %#v", extraParams["provider_option"])
}
}
func TestHandlerInvalidJSON(t *testing.T) {
h := NewHandler(&fakeRunner{})
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString("{"))

View File

@@ -75,14 +75,6 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
return nil, fmt.Errorf("%w: timeout_seconds must be greater than or equal to 0", ErrInvalidRequest)
}
model := strings.TrimSpace(req.Target.Model)
if model == "" {
model = strings.TrimSpace(c.defaultModel)
}
if model == "" {
return nil, fmt.Errorf("%w: model is required", ErrInvalidRequest)
}
endpoint := strings.TrimSpace(req.Target.Endpoint)
if endpoint == "" {
endpoint = c.baseURL
@@ -92,36 +84,9 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
}
endpoint = strings.TrimRight(endpoint, "/") + defaults.OpenAIChatCompletionsPath
wireReq := openAIChatRequest{
Model: model,
}
wireReq.Messages = make([]openAIChatMessage, 0, len(req.Prompt.Messages))
for _, msg := range req.Prompt.Messages {
wireReq.Messages = append(wireReq.Messages, openAIChatMessage{
Role: msg.Role,
Content: msg.Content,
})
}
if req.Target.Temperature != 0 {
wireReq.Temperature = &req.Target.Temperature
}
if req.Target.MaxTokens != 0 {
wireReq.MaxTokens = &req.Target.MaxTokens
}
if req.Target.TopP != 0 {
wireReq.TopP = &req.Target.TopP
}
if strings.TrimSpace(req.Target.ServiceTier) != "" {
wireReq.ServiceTier = req.Target.ServiceTier
}
if req.StructuredOutput != nil {
responseFormat, err := toOpenAIResponseFormat(req.StructuredOutput)
if err != nil {
return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
}
wireReq.ResponseFormat = responseFormat
wireReq, err := openAIChatRequestFromGenerateRequest(req, c.defaultModel)
if err != nil {
return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
}
payload, err := json.Marshal(wireReq)
@@ -190,6 +155,50 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
}, nil
}
func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultModel string) (openAIChatRequest, error) {
model := strings.TrimSpace(req.Target.Model)
if model == "" {
model = strings.TrimSpace(defaultModel)
}
if model == "" {
return openAIChatRequest{}, errors.New("model is required")
}
wireReq := openAIChatRequest{
Model: model,
}
wireReq.Messages = make([]openAIChatMessage, 0, len(req.Prompt.Messages))
for _, msg := range req.Prompt.Messages {
wireReq.Messages = append(wireReq.Messages, openAIChatMessage{
Role: msg.Role,
Content: msg.Content,
})
}
if req.Target.Temperature != 0 {
wireReq.Temperature = &req.Target.Temperature
}
if req.Target.MaxTokens != 0 {
wireReq.MaxTokens = &req.Target.MaxTokens
}
if req.Target.TopP != 0 {
wireReq.TopP = &req.Target.TopP
}
if strings.TrimSpace(req.Target.ServiceTier) != "" {
wireReq.ServiceTier = req.Target.ServiceTier
}
if req.StructuredOutput != nil {
responseFormat, err := toOpenAIResponseFormat(req.StructuredOutput)
if err != nil {
return openAIChatRequest{}, err
}
wireReq.ResponseFormat = responseFormat
}
return wireReq, nil
}
type openAIChatRequest struct {
Model string `json:"model"`
Messages []openAIChatMessage `json:"messages"`

View File

@@ -96,6 +96,15 @@ func TestOpenAICompatibleClientGenerateSuccess(t *testing.T) {
if got, ok := obs.Body["model"].(string); !ok || got != "gpt-test" {
t.Fatalf("unexpected model payload: %#v", obs.Body["model"])
}
if got, ok := obs.Body["temperature"].(float64); !ok || got != 0.4 {
t.Fatalf("unexpected temperature payload: %#v", obs.Body["temperature"])
}
if got, ok := obs.Body["max_tokens"].(float64); !ok || got != 123 {
t.Fatalf("unexpected max_tokens payload: %#v", obs.Body["max_tokens"])
}
if got, ok := obs.Body["top_p"].(float64); !ok || got != 0.7 {
t.Fatalf("unexpected top_p payload: %#v", obs.Body["top_p"])
}
if got, ok := obs.Body["service_tier"].(string); !ok || got != "priority" {
t.Fatalf("unexpected service_tier payload: %#v", obs.Body["service_tier"])
}
@@ -166,6 +175,43 @@ func TestOpenAICompatibleClientOmitsResponseFormatWhenNoStructuredOutput(t *test
}
}
func TestOpenAICompatibleClientOmitsReasoningEffortAndExtraParams(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",
ReasoningEffort: "high",
ExtraParams: map[string]string{
"provider_option": "on",
},
},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if _, exists := observedBody["reasoning_effort"]; exists {
t.Fatalf("expected reasoning_effort omitted, got %#v", observedBody["reasoning_effort"])
}
if _, exists := observedBody["extra_params"]; exists {
t.Fatalf("expected extra_params omitted, got %#v", observedBody["extra_params"])
}
}
func TestOpenAICompatibleClientNoAuthorizationHeaderWhenNoAPIKey(t *testing.T) {
hadAuth := false
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {

View File

@@ -1122,6 +1122,174 @@ func TestRunnerRunJSONSchemaRepairCarriesStructuredOutputSpec(t *testing.T) {
}
}
func TestExecutionProfileToTargetPopulatesAllFieldsAndCopiesExtraParams(t *testing.T) {
src := &domain.ExecutionProfile{
ID: "exec",
Endpoint: "http://profile/v1",
Model: "profile-model",
Temperature: 0.2,
MaxTokens: 123,
TopP: 0.75,
TimeoutSeconds: 90,
ServiceTier: "priority",
ReasoningEffort: "medium",
APIKeyEnv: "SCRIPTORIUM_API_KEY",
ExtraParams: map[string]string{
"provider_option": "on",
},
}
target := executionProfileToTarget(src)
if target.Endpoint != src.Endpoint ||
target.Model != src.Model ||
target.Temperature != src.Temperature ||
target.MaxTokens != src.MaxTokens ||
target.TopP != src.TopP ||
target.TimeoutSeconds != src.TimeoutSeconds ||
target.ServiceTier != src.ServiceTier ||
target.ReasoningEffort != src.ReasoningEffort ||
target.APIKeyEnv != src.APIKeyEnv {
t.Fatalf("expected all profile fields to populate target, got %+v", target)
}
if !reflect.DeepEqual(target.ExtraParams, src.ExtraParams) {
t.Fatalf("expected extra_params to match, got %#v", target.ExtraParams)
}
src.ExtraParams["provider_option"] = "changed"
if target.ExtraParams["provider_option"] != "on" {
t.Fatalf("expected extra_params copy to be independent, got %#v", target.ExtraParams)
}
}
func TestResolveExecutionTargetProfileValuesPopulateAllSupportedFields(t *testing.T) {
profileValue := &domain.ExecutionProfile{
ID: "exec",
Endpoint: "http://profile/v1",
Model: "profile-model",
Temperature: 0.3,
MaxTokens: 222,
TopP: 0.6,
TimeoutSeconds: 77,
ServiceTier: "priority",
ReasoningEffort: "low",
APIKeyEnv: "PROFILE_KEY",
ExtraParams: map[string]string{
"profile_option": "enabled",
},
}
target := resolveExecutionTarget(profileValue, nil)
if target.Endpoint != profileValue.Endpoint ||
target.Model != profileValue.Model ||
target.Temperature != profileValue.Temperature ||
target.MaxTokens != profileValue.MaxTokens ||
target.TopP != profileValue.TopP ||
target.TimeoutSeconds != profileValue.TimeoutSeconds ||
target.ServiceTier != profileValue.ServiceTier ||
target.ReasoningEffort != profileValue.ReasoningEffort ||
target.APIKeyEnv != profileValue.APIKeyEnv {
t.Fatalf("expected profile values to populate target, got %+v", target)
}
if !reflect.DeepEqual(target.ExtraParams, profileValue.ExtraParams) {
t.Fatalf("expected profile extra_params in target, got %#v", target.ExtraParams)
}
}
func TestResolveExecutionTargetRuntimeOverridesBeatProfileForAllOverrideableFields(t *testing.T) {
profileValue := &domain.ExecutionProfile{
ID: "exec",
Endpoint: "http://profile/v1",
Model: "profile-model",
Temperature: 0.2,
MaxTokens: 200,
TopP: 0.8,
TimeoutSeconds: 90,
ServiceTier: "priority",
ReasoningEffort: "medium",
APIKeyEnv: "PROFILE_KEY",
ExtraParams: map[string]string{
"profile_only": "yes",
},
}
override := &domain.ExecutionTarget{
Endpoint: "http://override/v1",
Model: "override-model",
Temperature: 0.9,
MaxTokens: 111,
TopP: 0.5,
TimeoutSeconds: 30,
ServiceTier: "flex",
ReasoningEffort: "high",
APIKeyEnv: "RUNTIME_KEY",
ExtraParams: map[string]string{
"runtime_only": "yes",
},
}
target := resolveExecutionTarget(profileValue, override)
if target.Endpoint != override.Endpoint ||
target.Model != override.Model ||
target.Temperature != override.Temperature ||
target.MaxTokens != override.MaxTokens ||
target.TopP != override.TopP ||
target.TimeoutSeconds != override.TimeoutSeconds ||
target.ServiceTier != override.ServiceTier ||
target.ReasoningEffort != override.ReasoningEffort ||
target.APIKeyEnv != override.APIKeyEnv {
t.Fatalf("expected runtime overrides to win for all fields, got %+v", target)
}
if !reflect.DeepEqual(target.ExtraParams, override.ExtraParams) {
t.Fatalf("expected runtime extra_params to replace profile extra_params, got %#v", target.ExtraParams)
}
}
func TestMergeExecutionTargetEmptyStringOverridesDoNotErase(t *testing.T) {
base := domain.ExecutionTarget{
Endpoint: "http://base/v1",
Model: "base-model",
ServiceTier: "priority",
ReasoningEffort: "medium",
APIKeyEnv: "BASE_KEY",
}
override := domain.ExecutionTarget{
Endpoint: "http://override/v1",
Model: "override-model",
ServiceTier: " ",
ReasoningEffort: " ",
APIKeyEnv: "",
}
merged := mergeExecutionTarget(base, override)
if merged.Endpoint != "http://override/v1" || merged.Model != "override-model" {
t.Fatalf("expected endpoint/model to override, got %+v", merged)
}
if merged.ServiceTier != "priority" {
t.Fatalf("expected empty service_tier override to be ignored, got %q", merged.ServiceTier)
}
if merged.ReasoningEffort != "medium" {
t.Fatalf("expected empty reasoning_effort override to be ignored, got %q", merged.ReasoningEffort)
}
if merged.APIKeyEnv != "BASE_KEY" {
t.Fatalf("expected empty api_key_env override to be ignored, got %q", merged.APIKeyEnv)
}
}
func TestMergeExecutionTargetEmptyExtraParamsDoesNotErase(t *testing.T) {
base := domain.ExecutionTarget{
ExtraParams: map[string]string{
"keep": "value",
},
}
override := domain.ExecutionTarget{
ExtraParams: map[string]string{},
}
merged := mergeExecutionTarget(base, override)
if !reflect.DeepEqual(merged.ExtraParams, base.ExtraParams) {
t.Fatalf("expected empty extra_params override not to erase base values, got %#v", merged.ExtraParams)
}
}
func TestBuildOutputArtifactDefaults(t *testing.T) {
tests := []struct {
name string