Centralize execution setting and session validation
This commit is contained in:
@@ -70,8 +70,8 @@ func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleCli
|
||||
}
|
||||
|
||||
func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) {
|
||||
if req.Target.TimeoutSeconds < 0 {
|
||||
return nil, fmt.Errorf("%w: timeout_seconds must be greater than or equal to 0", ErrInvalidRequest)
|
||||
if err := domain.ValidateExecutionTargetSettings(req.Target); err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
endpoint := strings.TrimSpace(req.Target.Endpoint)
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"math"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -1128,7 +1129,7 @@ func TestOpenAICompatibleClientCancellationReturnsRequestFailure(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleClientNegativeTimeoutRejected(t *testing.T) {
|
||||
func TestOpenAICompatibleClientRejectsInvalidExecutionSettings(t *testing.T) {
|
||||
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
|
||||
BaseURL: "http://example.com/v1",
|
||||
Model: "m",
|
||||
@@ -1137,15 +1138,34 @@ func TestOpenAICompatibleClientNegativeTimeoutRejected(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err = client.Generate(context.Background(), domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||
Target: domain.ExecutionTarget{TimeoutSeconds: -1},
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid request error")
|
||||
type testCase struct {
|
||||
name string
|
||||
target domain.ExecutionTarget
|
||||
}
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
tests := []testCase{
|
||||
{name: "non-finite temperature", target: domain.ExecutionTarget{Temperature: math.NaN()}},
|
||||
{name: "negative max tokens", target: domain.ExecutionTarget{MaxTokens: -1}},
|
||||
{name: "non-finite top p", target: domain.ExecutionTarget{TopP: math.Inf(1)}},
|
||||
{name: "negative timeout", target: domain.ExecutionTarget{TimeoutSeconds: -1}},
|
||||
}
|
||||
if strconv.IntSize == 64 {
|
||||
durationLimit := int64(math.MaxInt64 / int64(time.Second))
|
||||
tests = append(tests, testCase{name: "unrepresentable timeout", target: domain.ExecutionTarget{TimeoutSeconds: int(durationLimit) + 1}})
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := client.Generate(context.Background(), domain.GenerateRequest{
|
||||
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
|
||||
Target: tt.target,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid request error")
|
||||
}
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user