fix(copilot): harden provider configuration flow

This commit is contained in:
2026-07-13 14:57:28 +08:00
parent 19cee9af11
commit a712b91982
8 changed files with 153 additions and 36 deletions
+27 -4
View File
@@ -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")))
}
}
+17 -5
View File
@@ -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"),
))
}