From b38f7b4dc3f52d3b33864cf13e017b150f647d32 Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Sat, 4 Jul 2026 13:26:53 +0000 Subject: [PATCH] Serialize runtime extra parameters outbound --- internal/llm/openai_compatible_client.go | 86 ++++++++++-- internal/llm/openai_compatible_client_test.go | 124 +++++++++++++++++- 2 files changed, 198 insertions(+), 12 deletions(-) diff --git a/internal/llm/openai_compatible_client.go b/internal/llm/openai_compatible_client.go index 705e899..23ff064 100644 --- a/internal/llm/openai_compatible_client.go +++ b/internal/llm/openai_compatible_client.go @@ -90,7 +90,12 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) } - payload, err := json.Marshal(wireReq) + wirePayload, err := openAIChatRequestPayload(wireReq) + if err != nil { + return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) + } + + payload, err := json.Marshal(wirePayload) if err != nil { return nil, fmt.Errorf("%w: failed to encode request: %v", ErrRequestFailed, err) } @@ -194,6 +199,12 @@ func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultMod if strings.TrimSpace(req.Target.ServiceTier) != "" { wireReq.ServiceTier = req.Target.ServiceTier } + if strings.TrimSpace(req.Target.ReasoningEffort) != "" { + wireReq.ReasoningEffort = req.Target.ReasoningEffort + } + if len(req.Target.ExtraParams) > 0 { + wireReq.ExtraParams = req.Target.ExtraParams + } if req.StructuredOutput != nil { responseFormat, err := toOpenAIResponseFormat(req.StructuredOutput) if err != nil { @@ -206,14 +217,71 @@ func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultMod } type openAIChatRequest struct { - Model string `json:"model"` - SessionID string `json:"session_id,omitempty"` - Messages []openAIChatRequestMessage `json:"messages"` - Temperature *float64 `json:"temperature,omitempty"` - MaxTokens *int `json:"max_tokens,omitempty"` - TopP *float64 `json:"top_p,omitempty"` - ServiceTier string `json:"service_tier,omitempty"` - ResponseFormat *openAIResponseFormat `json:"response_format,omitempty"` + Model string `json:"model"` + SessionID string `json:"session_id,omitempty"` + Messages []openAIChatRequestMessage `json:"messages"` + Temperature *float64 `json:"temperature,omitempty"` + MaxTokens *int `json:"max_tokens,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + ServiceTier string `json:"service_tier,omitempty"` + ReasoningEffort string `json:"reasoning_effort,omitempty"` + ResponseFormat *openAIResponseFormat `json:"response_format,omitempty"` + ExtraParams map[string]any `json:"-"` +} + +func openAIChatRequestPayload(req openAIChatRequest) (map[string]any, error) { + out := map[string]any{ + "model": req.Model, + "messages": req.Messages, + } + if req.SessionID != "" { + out["session_id"] = req.SessionID + } + if req.Temperature != nil { + out["temperature"] = *req.Temperature + } + if req.MaxTokens != nil { + out["max_tokens"] = *req.MaxTokens + } + if req.TopP != nil { + out["top_p"] = *req.TopP + } + if req.ServiceTier != "" { + out["service_tier"] = req.ServiceTier + } + if req.ReasoningEffort != "" { + out["reasoning_effort"] = req.ReasoningEffort + } + if req.ResponseFormat != nil { + out["response_format"] = req.ResponseFormat + } + + for key, value := range req.ExtraParams { + if key == "" { + return nil, errors.New("extra_params key must not be empty") + } + if _, reserved := reservedOpenAIChatRequestFields[key]; reserved { + return nil, fmt.Errorf("extra_params key %q collides with reserved request field", key) + } + if _, err := json.Marshal(value); err != nil { + return nil, fmt.Errorf("extra_params.%s must be JSON-serializable: %w", key, err) + } + out[key] = value + } + + return out, nil +} + +var reservedOpenAIChatRequestFields = map[string]struct{}{ + "model": {}, + "session_id": {}, + "messages": {}, + "temperature": {}, + "max_tokens": {}, + "top_p": {}, + "service_tier": {}, + "reasoning_effort": {}, + "response_format": {}, } type openAIChatRequestMessage struct { diff --git a/internal/llm/openai_compatible_client_test.go b/internal/llm/openai_compatible_client_test.go index b685312..0a5ddd2 100644 --- a/internal/llm/openai_compatible_client_test.go +++ b/internal/llm/openai_compatible_client_test.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "math" "net/http" "net/http/httptest" "strings" @@ -417,7 +418,7 @@ func TestOpenAICompatibleClientOmitsResponseFormatWhenNoStructuredOutput(t *test } } -func TestOpenAICompatibleClientOmitsReasoningEffortAndExtraParams(t *testing.T) { +func TestOpenAICompatibleClientSerializesReasoningEffortAndExtraParams(t *testing.T) { var observedBody map[string]any ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { defer r.Body.Close() @@ -439,18 +440,135 @@ func TestOpenAICompatibleClientOmitsReasoningEffortAndExtraParams(t *testing.T) Model: "model", ReasoningEffort: "high", ExtraParams: map[string]any{ - "provider_option": "on", + "string_value": "on", + "number_value": 42, + "boolean_value": true, + "object_value": map[string]any{"nested": "value", "count": 2}, + "array_value": []any{"first", 3, false}, }, }, }) if err != nil { t.Fatalf("expected no error, got %v", err) } + if observedBody["reasoning_effort"] != "high" { + t.Fatalf("expected reasoning_effort high, got %#v", observedBody["reasoning_effort"]) + } + if observedBody["string_value"] != "on" { + t.Fatalf("unexpected string extra param: %#v", observedBody["string_value"]) + } + if observedBody["number_value"] != float64(42) { + t.Fatalf("unexpected number extra param: %#v", observedBody["number_value"]) + } + if observedBody["boolean_value"] != true { + t.Fatalf("unexpected boolean extra param: %#v", observedBody["boolean_value"]) + } + objectValue, ok := observedBody["object_value"].(map[string]any) + if !ok || objectValue["nested"] != "value" || objectValue["count"] != float64(2) { + t.Fatalf("unexpected object extra param: %#v", observedBody["object_value"]) + } + if _, exists := observedBody["extra_params"]; exists { + t.Fatalf("expected extra_params wrapper omitted, got %#v", observedBody["extra_params"]) + } + arrayValue, ok := observedBody["array_value"].([]any) + if !ok || len(arrayValue) != 3 || arrayValue[0] != "first" || arrayValue[1] != float64(3) || arrayValue[2] != false { + t.Fatalf("unexpected array extra param: %#v", observedBody["array_value"]) + } +} + +func TestOpenAICompatibleClientOmitsReasoningEffortWhenUnset(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) + } 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"]) + t.Fatalf("expected extra_params wrapper omitted, got %#v", observedBody["extra_params"]) + } +} + +func TestOpenAICompatibleClientRejectsInvalidExtraParamsBeforeProviderCall(t *testing.T) { + tests := []struct { + name string + extraParams map[string]any + want string + }{ + {name: "empty key", extraParams: map[string]any{"": "empty"}, want: "key must not be empty"}, + {name: "unserializable value", extraParams: map[string]any{"bad": math.Inf(1)}, want: "JSON-serializable"}, + } + for _, key := range []string{ + "model", + "session_id", + "messages", + "temperature", + "max_tokens", + "top_p", + "service_tier", + "reasoning_effort", + "response_format", + } { + tests = append(tests, struct { + name string + extraParams map[string]any + want string + }{ + name: "reserved key " + key, + extraParams: map[string]any{key: "collision"}, + want: "reserved request field", + }) + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + called := false + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + called = true + _, _ = 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", ExtraParams: tc.extraParams}, + }) + if err == nil { + t.Fatal("expected invalid request error") + } + if !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("expected ErrInvalidRequest, got %v", err) + } + if !strings.Contains(err.Error(), tc.want) { + t.Fatalf("expected error to contain %q, got %v", tc.want, err) + } + if called { + t.Fatal("provider should not be called for invalid extra_params") + } + }) } }