fix(copilot): harden provider configuration flow
This commit is contained in:
@@ -63,8 +63,8 @@ func TestNewOpenAIProvider_TrailingSlashTrimmed(t *testing.T) {
|
||||
|
||||
func TestChatRequest_MarshalJSON(t *testing.T) {
|
||||
req := ChatRequest{
|
||||
Model: "gpt-4",
|
||||
Messages: []ChatMessage{
|
||||
Model: "gpt-4",
|
||||
Messages: []ChatMessage{
|
||||
{Role: "system", Content: "You are a helpful assistant."},
|
||||
{Role: "user", Content: "Hello!"},
|
||||
},
|
||||
@@ -314,8 +314,9 @@ func TestOpenAIProvider_ChatCompletion_APIError(t *testing.T) {
|
||||
var apiErr *APIError
|
||||
require.True(t, errors.As(err, &apiErr), "error should wrap APIError")
|
||||
assert.Equal(t, http.StatusUnauthorized, apiErr.StatusCode)
|
||||
assert.Equal(t, "Invalid API key", apiErr.Message)
|
||||
assert.Equal(t, "provider authentication failed", apiErr.Message)
|
||||
assert.Equal(t, "invalid_request_error", apiErr.Type)
|
||||
assert.NotContains(t, err.Error(), "Invalid API key")
|
||||
}
|
||||
|
||||
func TestOpenAIProvider_ChatCompletion_RetryOn5xx(t *testing.T) {
|
||||
@@ -438,6 +439,28 @@ data: [DONE]
|
||||
assert.Equal(t, " world", chunks[2].Choices[0].Delta.Content)
|
||||
}
|
||||
|
||||
func TestOpenAIProvider_ParseSSEStreamSkipsMalformedProviderResponse(t *testing.T) {
|
||||
sseData := "data: secret-provider-response\n\ndata: [DONE]\n\n"
|
||||
p := NewOpenAIProvider(OpenAIProviderConfig{APIKey: "test"})
|
||||
|
||||
err := p.parseSSEStream(strings.NewReader(sseData), func(StreamChunk) error {
|
||||
t.Fatal("malformed provider chunks must not reach the callback")
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestParseAPIErrorDoesNotExposeProviderResponse(t *testing.T) {
|
||||
providerBody := []byte(`{"error":{"message":"authorization failed for sk-secret-value","type":"server_error","code":"provider_failure"},"debug":"complete response"}`)
|
||||
apiErr := parseAPIError(http.StatusInternalServerError, providerBody)
|
||||
|
||||
assert.Equal(t, "provider request failed", apiErr.Message)
|
||||
assert.Equal(t, "server_error", apiErr.Type)
|
||||
assert.Equal(t, "provider_failure", apiErr.Code)
|
||||
assert.NotContains(t, apiErr.Error(), "sk-secret-value")
|
||||
assert.NotContains(t, apiErr.Error(), "complete response")
|
||||
}
|
||||
|
||||
// --- Utility tests ---
|
||||
|
||||
func TestParseFloatEmbedding(t *testing.T) {
|
||||
@@ -481,4 +504,4 @@ func TestIsNonRetriableError(t *testing.T) {
|
||||
|
||||
// Non-APIError should not be treated as non-retriable
|
||||
assert.False(t, isNonRetriableError(fmt.Errorf("some error")))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -211,7 +211,7 @@ func (p *OpenAIProvider) doRequest(ctx context.Context, path string, body []byte
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
apiErr := parseAPIError(resp.StatusCode, respBody)
|
||||
applogger.L().Errorf("Provider API error (status %d, type=%s, code=%s)", resp.StatusCode, apiErr.Type, apiErr.Code)
|
||||
applogger.L().Errorf("Provider API error (status %d)", resp.StatusCode)
|
||||
return nil, apiErr
|
||||
}
|
||||
|
||||
@@ -251,7 +251,7 @@ func (p *OpenAIProvider) parseSSEStream(body io.Reader, onChunk func(StreamChunk
|
||||
|
||||
var chunk StreamChunk
|
||||
if err := json.Unmarshal([]byte(event.Data), &chunk); err != nil {
|
||||
applogger.L().Errorf("parseSSEStream: failed to unmarshal chunk: %v (data: %s)", err, event.Data)
|
||||
applogger.L().Errorf("parseSSEStream: failed to unmarshal provider chunk: %v", err)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -279,7 +279,7 @@ func (e *APIError) Error() string {
|
||||
func parseAPIError(statusCode int, body []byte) *APIError {
|
||||
apiErr := &APIError{
|
||||
StatusCode: statusCode,
|
||||
Message: string(body),
|
||||
Message: providerErrorMessage(statusCode),
|
||||
}
|
||||
|
||||
// Try to parse OpenAI error structure
|
||||
@@ -290,8 +290,7 @@ func parseAPIError(statusCode int, body []byte) *APIError {
|
||||
Code string `json:"code"`
|
||||
} `json:"error"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &errResp); err == nil && errResp.Error.Message != "" {
|
||||
apiErr.Message = errResp.Error.Message
|
||||
if err := json.Unmarshal(body, &errResp); err == nil {
|
||||
apiErr.Type = errResp.Error.Type
|
||||
apiErr.Code = errResp.Error.Code
|
||||
}
|
||||
@@ -299,6 +298,19 @@ func parseAPIError(statusCode int, body []byte) *APIError {
|
||||
return apiErr
|
||||
}
|
||||
|
||||
func providerErrorMessage(statusCode int) string {
|
||||
switch statusCode {
|
||||
case http.StatusUnauthorized, http.StatusForbidden:
|
||||
return "provider authentication failed"
|
||||
case http.StatusNotFound:
|
||||
return "provider endpoint or model was not found"
|
||||
case http.StatusTooManyRequests:
|
||||
return "provider rate limit exceeded"
|
||||
default:
|
||||
return "provider request failed"
|
||||
}
|
||||
}
|
||||
|
||||
// isNonRetriableError returns true for errors that should not be retried.
|
||||
func isNonRetriableError(err error) bool {
|
||||
if apiErr, ok := err.(*APIError); ok {
|
||||
|
||||
@@ -7,6 +7,8 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -626,22 +628,30 @@ func normalizeCopilotProviderError(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
message := strings.TrimSpace(err.Error())
|
||||
var providerErr *llm.APIError
|
||||
if errors.As(err, &providerErr) {
|
||||
switch providerErr.StatusCode {
|
||||
case http.StatusUnauthorized, http.StatusForbidden:
|
||||
return "provider authentication failed"
|
||||
case http.StatusNotFound:
|
||||
return "provider endpoint or model was not found"
|
||||
case http.StatusTooManyRequests:
|
||||
return "provider rate limit exceeded"
|
||||
default:
|
||||
return "provider request failed"
|
||||
}
|
||||
}
|
||||
|
||||
message := strings.ToLower(strings.TrimSpace(err.Error()))
|
||||
var networkErr net.Error
|
||||
switch {
|
||||
case errors.Is(err, context.DeadlineExceeded) || strings.Contains(strings.ToLower(message), "timeout"):
|
||||
case errors.Is(err, context.DeadlineExceeded) || errors.As(err, &networkErr) && networkErr.Timeout() || strings.Contains(message, "timeout"):
|
||||
return "provider request timed out"
|
||||
case strings.Contains(message, "401") || strings.Contains(message, "403"):
|
||||
return "provider authentication failed"
|
||||
case strings.Contains(message, "404"):
|
||||
return "provider endpoint or model was not found"
|
||||
case strings.Contains(message, "429"):
|
||||
return "provider rate limit exceeded"
|
||||
case strings.Contains(message, "connection refused") || strings.Contains(message, "no such host"):
|
||||
return "provider endpoint is unreachable"
|
||||
case strings.Contains(message, "unmarshal") || strings.Contains(message, "response format"):
|
||||
return "provider response format is incompatible"
|
||||
default:
|
||||
if len(message) > 240 {
|
||||
message = message[:240]
|
||||
}
|
||||
return message
|
||||
return "provider request failed"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
@@ -132,3 +133,17 @@ func TestCopilotConfigServiceTestsChatAndEmbeddingWithoutChangingSavedConfig(t *
|
||||
require.NoError(t, err)
|
||||
assert.False(t, payload.Configured, "testing candidate settings must not persist them")
|
||||
}
|
||||
|
||||
func TestNormalizeCopilotProviderErrorDoesNotExposeProviderResponse(t *testing.T) {
|
||||
providerErr := fmt.Errorf("chat completion: %w", &llm.APIError{
|
||||
StatusCode: http.StatusInternalServerError,
|
||||
Message: "complete provider response containing sk-secret-value",
|
||||
})
|
||||
|
||||
normalized := normalizeCopilotProviderError(providerErr)
|
||||
assert.Equal(t, "provider request failed", normalized)
|
||||
assert.NotContains(t, normalized, "sk-secret-value")
|
||||
assert.Equal(t, "provider response format is incompatible", normalizeCopilotProviderError(
|
||||
fmt.Errorf("unmarshal chat response: invalid character"),
|
||||
))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user