From 79901fbb867eb9315464e94a30b5500a571cdcbd Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 26 May 2026 13:07:35 +0000 Subject: [PATCH] Refine execution target mapping helpers and coverage across usecase, HTTP, and LLM --- internal/adapter/http/handler.go | 61 ++++--- internal/adapter/http/handler_test.go | 131 ++++++++++++++ internal/llm/openai_compatible_client.go | 85 +++++---- internal/llm/openai_compatible_client_test.go | 46 +++++ internal/usecase/runner_test.go | 168 ++++++++++++++++++ 5 files changed, 428 insertions(+), 63 deletions(-) diff --git a/internal/adapter/http/handler.go b/internal/adapter/http/handler.go index 13a1e45..5ce9d8c 100644 --- a/internal/adapter/http/handler.go +++ b/internal/adapter/http/handler.go @@ -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), diff --git a/internal/adapter/http/handler_test.go b/internal/adapter/http/handler_test.go index 6f6471e..06ae4c5 100644 --- a/internal/adapter/http/handler_test.go +++ b/internal/adapter/http/handler_test.go @@ -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("{")) diff --git a/internal/llm/openai_compatible_client.go b/internal/llm/openai_compatible_client.go index 2da1be6..8e32b74 100644 --- a/internal/llm/openai_compatible_client.go +++ b/internal/llm/openai_compatible_client.go @@ -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"` diff --git a/internal/llm/openai_compatible_client_test.go b/internal/llm/openai_compatible_client_test.go index f59db46..702b08d 100644 --- a/internal/llm/openai_compatible_client_test.go +++ b/internal/llm/openai_compatible_client_test.go @@ -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) { diff --git a/internal/usecase/runner_test.go b/internal/usecase/runner_test.go index 48c8998..06df4ed 100644 --- a/internal/usecase/runner_test.go +++ b/internal/usecase/runner_test.go @@ -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