Validate and compose provider endpoints

This commit is contained in:
2026-08-11 23:38:45 +00:00
parent c281f721bc
commit 3a43550f70
18 changed files with 448 additions and 84 deletions

View File

@@ -51,9 +51,11 @@ type OpenAICompatibleClient struct {
}
func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) {
baseURL := strings.TrimSpace(cfg.BaseURL)
if baseURL != "" {
if _, err := url.ParseRequestURI(baseURL); err != nil {
baseURL := ""
if strings.TrimSpace(cfg.BaseURL) != "" {
var err error
baseURL, err = domain.NormalizeOpenAICompatibleBaseEndpoint(cfg.BaseURL)
if err != nil {
return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err)
}
}
@@ -75,7 +77,7 @@ func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleCli
}
return &OpenAICompatibleClient{
baseURL: strings.TrimRight(baseURL, "/"),
baseURL: baseURL,
defaultModel: cfg.Model,
httpClient: client,
}, nil
@@ -86,14 +88,18 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err)
}
endpoint := strings.TrimSpace(req.Target.Endpoint)
if endpoint == "" {
endpoint = c.baseURL
selectedEndpoint := req.Target.Endpoint
if strings.TrimSpace(selectedEndpoint) == "" {
selectedEndpoint = c.baseURL
}
if endpoint == "" {
return nil, fmt.Errorf("%w: endpoint is required", ErrInvalidRequest)
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(selectedEndpoint)
if err != nil {
return nil, fmt.Errorf("%w: invalid endpoint: %v", ErrInvalidRequest, err)
}
endpoint, err = url.JoinPath(endpoint, defaults.OpenAIChatCompletionsPath)
if err != nil {
return nil, fmt.Errorf("%w: invalid endpoint path: %v", ErrInvalidRequest, err)
}
endpoint = strings.TrimRight(endpoint, "/") + defaults.OpenAIChatCompletionsPath
wireReq, err := openAIChatRequestFromGenerateRequest(req, c.defaultModel)
if err != nil {

View File

@@ -64,14 +64,21 @@ func assertDeadlineNear(t *testing.T, deadline, before, after time.Time, duratio
}
func TestNewOpenAICompatibleClientRejectsInvalidBaseURL(t *testing.T) {
_, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "://invalid",
})
if err == nil {
t.Fatal("expected invalid configuration error")
}
if !errors.Is(err, ErrInvalidConfig) {
t.Fatalf("expected ErrInvalidConfig, got %v", err)
for _, endpoint := range []string{
"://invalid",
"/v1",
"https:///v1",
"ftp://provider.example/v1",
"https://user@provider.example/v1",
"https://provider.example/v1?mode=chat",
"https://provider.example/v1#chat",
} {
t.Run(endpoint, func(t *testing.T) {
_, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: endpoint})
if !errors.Is(err, ErrInvalidConfig) {
t.Fatalf("expected ErrInvalidConfig, got %v", err)
}
})
}
}
@@ -1312,6 +1319,94 @@ func TestOpenAICompatibleClientAllowsEmptyConfiguredBaseURL(t *testing.T) {
}
}
func TestOpenAICompatibleClientComposesCompletionURL(t *testing.T) {
tests := []struct {
name string
baseURL string
wantURL string
}{
{name: "HTTP host", baseURL: "http://provider.example", wantURL: "http://provider.example/chat/completions"},
{name: "HTTPS host", baseURL: "https://provider.example", wantURL: "https://provider.example/chat/completions"},
{name: "nested path", baseURL: "https://provider.example/api/openai/v1", wantURL: "https://provider.example/api/openai/v1/chat/completions"},
{name: "trailing slash", baseURL: "https://provider.example/v1/", wantURL: "https://provider.example/v1/chat/completions"},
{name: "repeated trailing slashes", baseURL: " https://provider.example/api/v1/// ", wantURL: "https://provider.example/api/v1/chat/completions"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
var selectedURL string
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: tc.baseURL,
Model: "m",
HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
selectedURL = req.URL.String()
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{"choices":[{"message":{"content":"ok"}}]}`)),
Request: req,
}, nil
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
})
if err != nil {
t.Fatalf("generate: %v", err)
}
if selectedURL != tc.wantURL {
t.Fatalf("selected URL = %q, want %q", selectedURL, tc.wantURL)
}
})
}
}
func TestOpenAICompatibleClientRejectsInvalidSelectedEndpointBeforeTransport(t *testing.T) {
invalidEndpoints := []string{
"/v1",
"https:///v1",
"ftp://provider.example/v1",
"https://user@provider.example/v1",
"https://provider.example/v1?mode=chat",
"https://provider.example/v1#chat",
"https://sensitive-endpoint.example/%zz",
}
for _, endpoint := range invalidEndpoints {
t.Run(endpoint, func(t *testing.T) {
transportCalls := 0
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "https://configured.example/v1",
Model: "m",
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
transportCalls++
return nil, errors.New("transport must not be called")
})},
})
if err != nil {
t.Fatalf("construct client: %v", err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{Endpoint: endpoint},
})
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
if strings.Contains(err.Error(), endpoint) {
t.Fatalf("error exposed selected endpoint %q: %v", endpoint, err)
}
if transportCalls != 0 {
t.Fatalf("transport calls = %d, want 0", transportCalls)
}
})
}
}
func TestOpenAICompatibleClientRequiresEndpointWhenUnsetEverywhere(t *testing.T) {
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "",