Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 16 additions & 4 deletions model/gemini/gemini.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ import (
"trpc.group/trpc-go/trpc-agent-go/model"
imodel "trpc.group/trpc-go/trpc-agent-go/model/internal/model"
"trpc.group/trpc-go/trpc-agent-go/model/internal/modeltailoring"
"trpc.group/trpc-go/trpc-agent-go/platform"
"trpc.group/trpc-go/trpc-agent-go/tool"
)

Expand Down Expand Up @@ -269,7 +270,7 @@ func (m *Model) handleNonStreamingResponse(
if err != nil {
errorResponse := &model.Response{
Error: &model.ResponseError{
Message: err.Error(),
Message: redactErrorMessage(err),
Type: model.ErrorTypeAPIError,
},
Timestamp: time.Now(),
Expand All @@ -295,7 +296,7 @@ func (m *Model) handleNonStreamingResponse(
if err != nil {
errorResponse := &model.Response{
Error: &model.ResponseError{
Message: err.Error(),
Message: redactErrorMessage(err),
Type: model.ErrorTypeAPIError,
},
Timestamp: time.Now(),
Expand Down Expand Up @@ -375,7 +376,7 @@ func (m *Model) handleStreamingResponse(
}
sendTerminalResponse(&model.Response{
Error: &model.ResponseError{
Message: err.Error(),
Message: redactErrorMessage(err),
Type: model.ErrorTypeAPIError,
},
Timestamp: time.Now(),
Expand Down Expand Up @@ -411,7 +412,7 @@ func (m *Model) handleStreamingResponse(
if retryErr != nil {
sendTerminalResponse(&model.Response{
Error: &model.ResponseError{
Message: retryErr.Error(),
Message: redactErrorMessage(retryErr),
Type: model.ErrorTypeAPIError,
},
Timestamp: time.Now(),
Expand Down Expand Up @@ -447,6 +448,17 @@ func (m *Model) handleStreamingResponse(
sendTerminalResponse(finalResponse)
}

func redactErrorMessage(err error) string {
if err == nil {
return ""
}
redactor, redactorErr := platform.NewRedactor()
if redactorErr != nil {
return "redacted error detail unavailable"
}
return redactor.Redact(err.Error())
}

// convertContentBlock builds a single assistant message from Gemini Candidate.
func (m *Model) convertContentBlock(candidates []*genai.Candidate) (model.Message, string) {
var (
Expand Down
36 changes: 30 additions & 6 deletions model/gemini/gemini_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1483,7 +1483,7 @@ func TestModel_GenerateContentError(t *testing.T) {
},
},
}
err := errors.New("error")
err := errors.New("error: Authorization: Bearer raw-token api_key=sk-testsecret token=raw-token secret: raw-secret password=raw-password Cookie: session=raw-cookie")
ctrl := gomock.NewController(t)
defer ctrl.Finish()

Expand Down Expand Up @@ -1521,7 +1521,19 @@ func TestModel_GenerateContentError(t *testing.T) {
m := &Model{
client: mockClient,
}
_, _ = m.GenerateContent(tt.args.ctx, tt.args.request)
ch, _ := m.GenerateContent(tt.args.ctx, tt.args.request)
if tt.name != "error" {
return
}
resp := <-ch
require.NotNil(t, resp.Error)
assert.Contains(t, resp.Error.Message, "error")
assert.NotContains(t, resp.Error.Message, "raw-token")
assert.NotContains(t, resp.Error.Message, "sk-testsecret")
assert.NotContains(t, resp.Error.Message, "raw-secret")
assert.NotContains(t, resp.Error.Message, "raw-password")
assert.NotContains(t, resp.Error.Message, "raw-cookie")
assert.Contains(t, resp.Error.Message, "****")
})
}
}
Expand Down Expand Up @@ -1760,7 +1772,7 @@ func TestModel_GenerateContentStreamingError(t *testing.T) {

t.Run("immediate_stream_error", func(t *testing.T) {
// Test when the stream immediately returns an error on the first chunk
streamErr := errors.New("stream connection failed")
streamErr := errors.New("stream connection failed: Authorization: Bearer raw-token api_key=sk-testsecret token=raw-token secret: raw-secret password=raw-password Cookie: session=raw-cookie")
ctrl := gomock.NewController(t)
defer ctrl.Finish()

Expand All @@ -1781,7 +1793,13 @@ func TestModel_GenerateContentStreamingError(t *testing.T) {
resp := <-respChan
assert.NotNil(t, resp)
assert.NotNil(t, resp.Error)
assert.Equal(t, "stream connection failed", resp.Error.Message)
assert.Contains(t, resp.Error.Message, "stream connection failed")
assert.NotContains(t, resp.Error.Message, "raw-token")
assert.NotContains(t, resp.Error.Message, "sk-testsecret")
assert.NotContains(t, resp.Error.Message, "raw-secret")
assert.NotContains(t, resp.Error.Message, "raw-password")
assert.NotContains(t, resp.Error.Message, "raw-cookie")
assert.Contains(t, resp.Error.Message, "****")
assert.Equal(t, model.ErrorTypeAPIError, resp.Error.Type)
assert.True(t, resp.Done)

Expand Down Expand Up @@ -2240,7 +2258,7 @@ func TestModel_NonStreaming_MalformedFunctionCallRetryError(t *testing.T) {
{FinishReason: genai.FinishReason("MALFORMED_FUNCTION_CALL")},
},
}
retryErr := errors.New("network error on retry")
retryErr := errors.New("network error on retry: Authorization: Bearer raw-token api_key=sk-testsecret token=raw-token secret: raw-secret password=raw-password Cookie: session=raw-cookie")

ctrl := gomock.NewController(t)
defer ctrl.Finish()
Expand All @@ -2264,7 +2282,13 @@ func TestModel_NonStreaming_MalformedFunctionCallRetryError(t *testing.T) {
require.Len(t, responses, 1)
require.True(t, responses[0].Done)
require.NotNil(t, responses[0].Error)
require.Equal(t, "network error on retry", responses[0].Error.Message)
require.Contains(t, responses[0].Error.Message, "network error on retry")
require.NotContains(t, responses[0].Error.Message, "raw-token")
require.NotContains(t, responses[0].Error.Message, "sk-testsecret")
require.NotContains(t, responses[0].Error.Message, "raw-secret")
require.NotContains(t, responses[0].Error.Message, "raw-password")
require.NotContains(t, responses[0].Error.Message, "raw-cookie")
require.Contains(t, responses[0].Error.Message, "****")
}

// TestModel_Streaming_MalformedFunctionCallRetry verifies that when the
Expand Down
18 changes: 15 additions & 3 deletions model/huggingface/huggingface.go
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ import (
"trpc.group/trpc-go/trpc-agent-go/model"
imodel "trpc.group/trpc-go/trpc-agent-go/model/internal/model"
"trpc.group/trpc-go/trpc-agent-go/model/internal/modeltailoring"
"trpc.group/trpc-go/trpc-agent-go/platform"
)

// Model implements the model.Model interface for HuggingFace API.
Expand Down Expand Up @@ -215,7 +216,7 @@ func (m *Model) handleNonStreamingRequest(
if err != nil {
responseChan <- &model.Response{
Error: &model.ResponseError{
Message: fmt.Sprintf("failed to make request: %v", err),
Message: redactErrorMessage(fmt.Errorf("failed to make request: %w", err)),
},
}
return
Expand Down Expand Up @@ -248,7 +249,7 @@ func (m *Model) handleStreamingRequest(
streamErr = err
terminalResponse = &model.Response{
Error: &model.ResponseError{
Message: fmt.Sprintf("failed to make streaming request: %v", err),
Message: redactErrorMessage(fmt.Errorf("failed to make streaming request: %w", err)),
},
}
} else {
Expand All @@ -266,7 +267,7 @@ func (m *Model) handleStreamingRequest(
streamErr = err
terminalResponse = &model.Response{
Error: &model.ResponseError{
Message: fmt.Sprintf("error reading stream: %v", err),
Message: redactErrorMessage(fmt.Errorf("error reading stream: %w", err)),
},
}
break
Expand Down Expand Up @@ -309,6 +310,17 @@ func (m *Model) handleStreamingRequest(
}
}

func redactErrorMessage(err error) string {
if err == nil {
return ""
}
redactor, redactorErr := platform.NewRedactor()
if redactorErr != nil {
return "redacted error detail unavailable"
}
return redactor.Redact(err.Error())
}

// makeRequest makes a non-streaming HTTP request to the HuggingFace API.
func (m *Model) makeRequest(ctx context.Context, hfRequest *ChatCompletionRequest) (*ChatCompletionResponse, error) {
// Marshal request to JSON.
Expand Down
24 changes: 24 additions & 0 deletions model/huggingface/huggingface_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -351,6 +351,30 @@ func TestModel_GenerateContent_NonStreaming(t *testing.T) {
}
}

func TestModel_GenerateContent_NonStreamingRedactsRequestError(t *testing.T) {
m, err := New(
"mistralai/Mistral-7B-Instruct-v0.2",
WithAPIKey("test-api-key"),
WithBaseURL("http://127.0.0.1:1/path?api_key=sk-testsecret&token=raw-token&secret=raw-secret&password=raw-password&cookie=raw-cookie"),
)
require.NoError(t, err)

responseChan, err := m.GenerateContent(context.Background(), &model.Request{
Messages: []model.Message{{Role: model.RoleUser, Content: "Hello"}},
})
require.NoError(t, err)

resp := <-responseChan
require.NotNil(t, resp.Error)
assert.Contains(t, resp.Error.Message, "failed to make request")
assert.NotContains(t, resp.Error.Message, "raw-token")
assert.NotContains(t, resp.Error.Message, "sk-testsecret")
assert.NotContains(t, resp.Error.Message, "raw-secret")
assert.NotContains(t, resp.Error.Message, "raw-password")
assert.NotContains(t, resp.Error.Message, "raw-cookie")
assert.Contains(t, resp.Error.Message, "****")
}

func TestModel_GenerateContent_Streaming(t *testing.T) {
tests := []struct {
name string
Expand Down
22 changes: 19 additions & 3 deletions model/openai/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ import (
"trpc.group/trpc-go/trpc-agent-go/model"
imodel "trpc.group/trpc-go/trpc-agent-go/model/internal/model"
"trpc.group/trpc-go/trpc-agent-go/model/internal/modeltailoring"
"trpc.group/trpc-go/trpc-agent-go/platform"
"trpc.group/trpc-go/trpc-agent-go/tool"
)

Expand Down Expand Up @@ -2376,7 +2377,7 @@ func (m *Model) emitStreamingFinalResponse(
// Send error response.
emit(&model.Response{
Error: &model.ResponseError{
Message: stream.Err().Error(),
Message: redactErrorMessage(stream.Err()),
Type: model.ErrorTypeStreamError,
},
Timestamp: time.Now(),
Expand Down Expand Up @@ -2519,7 +2520,7 @@ func (m *Model) handleNonStreamingResponseWithEmitter(
}
emit(&model.Response{
Error: &model.ResponseError{
Message: err.Error(),
Message: redactErrorMessage(err),
Type: model.ErrorTypeAPIError,
},
Timestamp: time.Now(),
Expand Down Expand Up @@ -3041,7 +3042,7 @@ func extractEmbeddedErrorResponse(cc *openai.ChatCompletion) *model.Response {
}
log.Debugf("OpenAI-compatible API returned HTTP 200 with embedded error (type=%s)", errBody.Type)
respErr := &model.ResponseError{
Message: errMsg,
Message: redactErrorText(errMsg),
Type: model.ErrorTypeAPIError,
}
if code := normalizeEmbeddedErrorCode(errBody.Code); code != "" {
Expand All @@ -3057,6 +3058,21 @@ func extractEmbeddedErrorResponse(cc *openai.ChatCompletion) *model.Response {
}
}

func redactErrorMessage(err error) string {
if err == nil {
return ""
}
return redactErrorText(err.Error())
}

func redactErrorText(message string) string {
redactor, err := platform.NewRedactor()
if err != nil {
return "redacted error detail unavailable"
}
return redactor.Redact(message)
}

// normalizeEmbeddedErrorString extracts a JSON string value.
// Returns "" for absent, null, or non-string values.
func normalizeEmbeddedErrorString(raw json.RawMessage) string {
Expand Down
46 changes: 31 additions & 15 deletions model/openai/openai_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9613,13 +9613,15 @@ func TestExtractEmbeddedErrorResponse(t *testing.T) {
}

tests := []struct {
name string
completion *openaigo.ChatCompletion
wantNil bool
wantMessage string
wantType string
wantCode string
wantParam string
name string
completion *openaigo.ChatCompletion
wantNil bool
wantMessage string
wantType string
wantCode string
wantParam string
wantRedacted bool
wantNoSubstrings []string
}{
{
name: "nil completion",
Expand Down Expand Up @@ -9670,16 +9672,18 @@ func TestExtractEmbeddedErrorResponse(t *testing.T) {
name: "empty completion with embedded error containing ret_code",
completion: mustCompletion(t, `{
"error": {
"message": "API key exceeded rate limit",
"message": "API key exceeded rate limit: Authorization: Bearer raw-token api_key=sk-testsecret token=raw-token secret: raw-secret password=raw-password Cookie: session=raw-cookie",
"type": "requestAuthError",
"code": "rate_limited",
"ret_code": -2000
}
}`),
wantNil: false,
wantMessage: "API key exceeded rate limit",
wantType: model.ErrorTypeAPIError,
wantCode: "rate_limited",
wantNil: false,
wantMessage: "API key exceeded rate limit",
wantType: model.ErrorTypeAPIError,
wantCode: "rate_limited",
wantRedacted: true,
wantNoSubstrings: []string{"raw-token", "sk-testsecret", "raw-secret", "raw-password", "raw-cookie"},
},
{
name: "empty completion with error message only, type defaults to api_error",
Expand Down Expand Up @@ -9759,7 +9763,13 @@ func TestExtractEmbeddedErrorResponse(t *testing.T) {
}
require.NotNil(t, resp, "expected non-nil error response")
require.NotNil(t, resp.Error, "expected Error field")
assert.Equal(t, tt.wantMessage, resp.Error.Message)
assert.Contains(t, resp.Error.Message, tt.wantMessage)
for _, forbidden := range tt.wantNoSubstrings {
assert.NotContains(t, resp.Error.Message, forbidden)
}
if tt.wantRedacted {
assert.Contains(t, resp.Error.Message, "****")
}
assert.Equal(t, tt.wantType, resp.Error.Type)
assert.True(t, resp.Done, "error response must be Done")
if tt.wantCode != "" {
Expand All @@ -9784,7 +9794,7 @@ func TestModel_GenerateContent_EmbeddedErrorHTTP200(t *testing.T) {
// Simulate a provider returning HTTP 200 with an error body.
fmt.Fprint(w, `{
"error": {
"message": "When using tool_choice, 'tools' must be set.",
"message": "When using tool_choice, 'tools' must be set. Authorization: Bearer raw-token api_key=sk-testsecret token=raw-token secret: raw-secret password=raw-password Cookie: session=raw-cookie",
"type": "invalid_request_error",
"ret_code": -1000
}
Expand Down Expand Up @@ -9815,7 +9825,13 @@ func TestModel_GenerateContent_EmbeddedErrorHTTP200(t *testing.T) {
require.Len(t, responses, 1, "expected exactly one response")
resp := responses[0]
require.NotNil(t, resp.Error, "response must carry an error")
assert.Equal(t, "When using tool_choice, 'tools' must be set.", resp.Error.Message)
assert.Contains(t, resp.Error.Message, "When using tool_choice, 'tools' must be set.")
assert.NotContains(t, resp.Error.Message, "raw-token")
assert.NotContains(t, resp.Error.Message, "sk-testsecret")
assert.NotContains(t, resp.Error.Message, "raw-secret")
assert.NotContains(t, resp.Error.Message, "raw-password")
assert.NotContains(t, resp.Error.Message, "raw-cookie")
assert.Contains(t, resp.Error.Message, "****")
assert.Equal(t, model.ErrorTypeAPIError, resp.Error.Type)
assert.True(t, resp.Done)
assert.Empty(t, resp.Choices, "error response should have no choices")
Expand Down
Loading