From 82c304f9bf5da733019739a35d37f09738307e6f Mon Sep 17 00:00:00 2001 From: linseasea Date: Thu, 13 Aug 2026 18:44:26 +0800 Subject: [PATCH 1/2] fix(relay): mask upstream credentials in all relayed error paths (stream, non-stream, headers, task/mj/alpha-search) --- common/str.go | 6 + relay/alpha_search_handler.go | 12 +- relay/channel/openai/relay_realtime.go | 3 + relay/helper/stream_scanner.go | 148 ++++++++++++++++++ relay/mjproxy_handler.go | 20 ++- relaykit/relayconvert/kitutil/mask.go | 90 +++++++++++ .../relayconvert/kitutil/mask_formats_test.go | 50 ++++++ relaykit/types/error.go | 17 +- service/error.go | 8 +- service/http.go | 18 +++ 10 files changed, 362 insertions(+), 10 deletions(-) create mode 100644 relaykit/relayconvert/kitutil/mask_formats_test.go diff --git a/common/str.go b/common/str.go index b412c7dc11..c0074eb75f 100644 --- a/common/str.go +++ b/common/str.go @@ -141,3 +141,9 @@ func MaskEmail(email string) string { func MaskSensitiveInfo(str string) string { return kitutil.MaskSensitiveInfo(str) } + +// MaskSensitiveKeys masks only credential-shaped values (no URL/domain/IP +// rewriting), for enum-like fields where full masking would mangle values. +func MaskSensitiveKeys(str string) string { + return kitutil.MaskSensitiveKeys(str) +} diff --git a/relay/alpha_search_handler.go b/relay/alpha_search_handler.go index 246eebe864..9bd77fce16 100644 --- a/relay/alpha_search_handler.go +++ b/relay/alpha_search_handler.go @@ -95,8 +95,18 @@ func AlphaSearchHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError if contentType := httpResp.Header.Get("Content-Type"); contentType != "" { c.Writer.Header().Set("Content-Type", contentType) } + // F-65: read + mask the upstream body before relaying. /v1/alpha/search + // used to io.Copy the raw body to the client, bypassing the F-13/F-20 + // masking system; an upstream that echoes the channel Authorization value + // in a 200 body would leak the key. (F-36 already caps body size here via + // DoApiRequest's LimitReader.) + responseBody, err := io.ReadAll(httpResp.Body) + if err != nil { + return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) + } + responseBody = []byte(common.MaskSensitiveInfo(string(responseBody))) c.Writer.WriteHeader(httpResp.StatusCode) - if _, err := io.Copy(c.Writer, httpResp.Body); err != nil { + if _, err := c.Writer.Write(responseBody); err != nil { return types.NewError(err, types.ErrorCodeDoRequestFailed, types.ErrOptionWithSkipRetry()) } diff --git a/relay/channel/openai/relay_realtime.go b/relay/channel/openai/relay_realtime.go index bea97780df..7771464d19 100644 --- a/relay/channel/openai/relay_realtime.go +++ b/relay/channel/openai/relay_realtime.go @@ -187,6 +187,9 @@ func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types. localUsage.OutputTokenDetails.AudioTokens += audioToken } + // F-13 layer A: upstream realtime frames are relayed verbatim; + // mask error frames that may echo the channel key before forwarding. + message = []byte(helper.MaskStreamErrorData(string(message))) err = helper.WssString(c, clientConn, string(message)) if err != nil { errChan <- fmt.Errorf("error writing to client: %v", err) diff --git a/relay/helper/stream_scanner.go b/relay/helper/stream_scanner.go index f88bef5206..bd920310c1 100644 --- a/relay/helper/stream_scanner.go +++ b/relay/helper/stream_scanner.go @@ -3,6 +3,7 @@ package helper import ( "bufio" "context" + "encoding/json" "fmt" "io" "net/http" @@ -74,6 +75,152 @@ func ExtendWriteDeadline(c *gin.Context) { _ = http.NewResponseController(c.Writer).SetWriteDeadline(time.Now().Add(streamWriteTimeout)) } +// MaskStreamErrorData masks sensitive values inside SSE error chunks before +// they reach a downstream handler. Upstreams may echo the gateway's upstream +// Authorization header (i.e. the channel API key) in streamed error messages; +// the non-stream error path already masks, but streamed chunks used to be +// relayed verbatim (F-13 layer A). +// +// F-20: also handles chunks that carry the key outside a top-level "error" key +// (e.g. {"message":"invalid key: sk-..."}) and masks non-message fields inside +// the error object (code/type/param), which the first F-13 masker skipped. +func MaskStreamErrorData(data string) string { + if !json.Valid([]byte(data)) { + return data + } + var probe map[string]json.RawMessage + if err := json.Unmarshal([]byte(data), &probe); err != nil { + return data + } + changed := false + maskField := func(raw json.RawMessage, masker func(string) string) (json.RawMessage, bool) { + var s string + if err := json.Unmarshal(raw, &s); err != nil || s == "" { + return raw, false + } + masked := masker(s) + if masked == s { + return raw, false + } + b, err := json.Marshal(masked) + if err != nil { + return raw, false + } + return b, true + } + // Top-level fields commonly used by error-ish chunks (F-20b: no "error" key). + for _, key := range []string{"message", "code", "type"} { + if raw, ok := probe[key]; ok { + masker := common.MaskSensitiveInfo + if key != "message" { + // Enum-like fields must not be URL/domain/IP-mangled. + masker = common.MaskSensitiveKeys + } + if b, ok2 := maskField(raw, masker); ok2 { + probe[key] = b + changed = true + } + } + } + // "error" key: object or string. + if raw, ok := probe["error"]; ok && len(raw) > 0 { + switch raw[0] { + case '{': + var errObj map[string]json.RawMessage + if err := json.Unmarshal(raw, &errObj); err == nil { + for _, key := range []string{"message", "code", "type", "param"} { + if v, ok := errObj[key]; ok { + masker := common.MaskSensitiveInfo + if key != "message" { + masker = common.MaskSensitiveKeys + } + if b, ok2 := maskField(v, masker); ok2 { + errObj[key] = b + changed = true + } + } + } + if b, err := json.Marshal(errObj); err == nil { + probe["error"] = b + } + } + case '"': + if b, ok2 := maskField(raw, common.MaskSensitiveKeys); ok2 { + probe["error"] = b + changed = true + } + } + } + if !changed { + // F-20 residual: upstreams may also echo the channel key inside + // non-error chunks (e.g. a mid-stream "usage" frame carrying unknown + // metadata). Recursively mask string values of the chunk, but only + // re-marshal when something actually changed so legitimate chunks + // keep their original byte layout (JSON key order preserved). + if strings.Contains(data, "usage") { + var walk func(raw json.RawMessage) (json.RawMessage, bool) + walk = func(raw json.RawMessage) (json.RawMessage, bool) { + var obj map[string]json.RawMessage + if err := json.Unmarshal(raw, &obj); err == nil { + objChanged := false + for k, v := range obj { + nv, ch := walk(v) + if ch { + obj[k] = nv + objChanged = true + } + } + if objChanged { + if b, err := json.Marshal(obj); err == nil { + return b, true + } + } + return raw, false + } + var arr []json.RawMessage + if err := json.Unmarshal(raw, &arr); err == nil { + arrChanged := false + for i, v := range arr { + nv, ch := walk(v) + if ch { + arr[i] = nv + arrChanged = true + } + } + if arrChanged { + if b, err := json.Marshal(arr); err == nil { + return b, true + } + } + return raw, false + } + if b, ok2 := maskField(raw, common.MaskSensitiveKeys); ok2 { + return b, true + } + return raw, false + } + anyChanged := false + for key, raw := range probe { + nv, ch := walk(raw) + if ch { + probe[key] = nv + anyChanged = true + } + } + if anyChanged { + if b, err := json.Marshal(probe); err == nil { + return string(b) + } + } + } + return data + } + if b, err := json.Marshal(probe); err == nil { + return string(b) + } + return data +} + func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, dataHandler func(data string, sr *StreamResult)) { if resp == nil || dataHandler == nil { @@ -262,6 +409,7 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon if data == "" { continue } + data = MaskStreamErrorData(data) if !strings.HasPrefix(data, "[DONE]") { info.SetFirstResponseTime() info.ReceivedResponseCount++ diff --git a/relay/mjproxy_handler.go b/relay/mjproxy_handler.go index daa0585a07..93a7fa7d33 100644 --- a/relay/mjproxy_handler.go +++ b/relay/mjproxy_handler.go @@ -77,7 +77,8 @@ func RelayMidjourneyImage(c *gin.Context) { if resp.StatusCode != http.StatusOK { responseBody, _ := io.ReadAll(resp.Body) c.JSON(resp.StatusCode, gin.H{ - "error": string(responseBody), + // F-42: mask upstream error body before echoing it to the client. + "error": common.MaskSensitiveInfo(string(responseBody)), }) return } @@ -264,7 +265,9 @@ func RelaySwapFace(c *gin.Context, info *relaycommon.RelayInfo) *dto.MidjourneyR MjId: midjResponse.Result, Prompt: "InsightFace", PromptEn: "", - Description: midjResponse.Description, + // F-64: mask upstream-controlled text before persistence; task fetch + // replays these fields to the client verbatim. + Description: common.MaskSensitiveInfo(midjResponse.Description), State: "", SubmitTime: info.StartTime.UnixNano() / int64(time.Millisecond), StartTime: time.Now().UnixNano() / int64(time.Millisecond), @@ -321,6 +324,9 @@ func RelayMidjourneyTaskImageSeed(c *gin.Context) *dto.MidjourneyResponse { if err != nil { return service.MidjourneyErrorWrapper(constant.MjRequestError, "unmarshal_response_body_failed") } + // F-64: the marshaled upstream response is copied to the client verbatim; + // apply the same key-masking used by the OpenAI/Claude error paths. + respBody = []byte(common.MaskSensitiveInfo(string(respBody))) service.IOCopyBytesGracefully(c, nil, respBody) return nil } @@ -577,7 +583,8 @@ func RelayMidjourneySubmit(c *gin.Context, relayInfo *relaycommon.RelayInfo) *dt MjId: midjResponse.Result, Prompt: midjRequest.Prompt, PromptEn: "", - Description: midjResponse.Description, + // F-64: mask upstream-controlled text before persistence. + Description: common.MaskSensitiveInfo(midjResponse.Description), State: "", SubmitTime: time.Now().UnixNano() / int64(time.Millisecond), StartTime: 0, @@ -601,7 +608,9 @@ func RelayMidjourneySubmit(c *gin.Context, relayInfo *relaycommon.RelayInfo) *dt } if midjResponse.Code != 1 && midjResponse.Code != 21 && midjResponse.Code != 22 { //非1-提交成功,21-任务已存在和22-排队中,则记录错误原因 - midjourneyTask.FailReason = midjResponse.Description + // F-64: upstream error bodies may echo the channel key (mj-api-secret); + // mask before persistence so /task/:id/fetch cannot replay it. + midjourneyTask.FailReason = common.MaskSensitiveInfo(midjResponse.Description) consumeQuota = false } @@ -646,6 +655,9 @@ func RelayMidjourneySubmit(c *gin.Context, relayInfo *relaycommon.RelayInfo) *dt responseBody = []byte(newBody) } //resp.Body = io.NopCloser(bytes.NewBuffer(responseBody)) + // F-64: raw upstream body is relayed verbatim; mask it before writing so an + // upstream that echoes mj-api-secret (or any channel key) cannot leak it. + responseBody = []byte(common.MaskSensitiveInfo(string(responseBody))) bodyReader := io.NopCloser(bytes.NewBuffer(responseBody)) //for k, v := range resp.Header { diff --git a/relaykit/relayconvert/kitutil/mask.go b/relaykit/relayconvert/kitutil/mask.go index d193a56e84..b482bd4284 100644 --- a/relaykit/relayconvert/kitutil/mask.go +++ b/relaykit/relayconvert/kitutil/mask.go @@ -12,6 +12,41 @@ var ( maskIPPattern = regexp.MustCompile(`\b(?:\d{1,3}\.){3}\d{1,3}\b`) // maskApiKeyPattern matches patterns like 'api_key:xxx' or "api_key:xxx" to mask the API key value maskApiKeyPattern = regexp.MustCompile(`(['"]?)api_key:([^\s'"]+)(['"]?)`) + + // maskBareKeyPattern masks bare upstream API key formats that providers may + // echo in error messages (sk-..., gsk_..., hf_..., xai-..., AIza...). + maskBareKeyPattern = regexp.MustCompile(`(?i)(sk-[A-Za-z0-9_\-]{8,}|gsk_[A-Za-z0-9_\-]{8,}|hf_[A-Za-z0-9_\-]{8,}|xai-[A-Za-z0-9_\-]{8,}|AIza[A-Za-z0-9_\-]{8,})`) + + // maskMoreKeyPrefixPattern covers additional provider key prefixes that + // F-13's bare pattern missed (Perplexity, NVIDIA NIM, Replicate, personal + // access tokens, GitHub tokens, GitLab tokens). + maskMoreKeyPrefixPattern = regexp.MustCompile(`(?i)(pplx-[A-Za-z0-9_\-]{8,}|nvapi-[A-Za-z0-9_\-]{8,}|r8_[A-Za-z0-9_\-]{8,}|pat_[A-Za-z0-9_\-]{8,}|rt_[A-Za-z0-9_\-]{8,}|ghp_[A-Za-z0-9_\-]{8,}|gho_[A-Za-z0-9_\-]{8,}|github_pat_[A-Za-z0-9_\-]{8,}|glpat-[A-Za-z0-9_\-]{8,})`) + + // maskJwtPattern masks JWT-shaped credentials (header.payload.signature). + maskJwtPattern = regexp.MustCompile(`(?i)eyJ[A-Za-z0-9_\-]+\.[A-Za-z0-9_\-]+\.[A-Za-z0-9_\-]+`) + + // maskKeyValuePattern masks prefix-less long values adjacent to credential + // field names (api_key:/token:/secret:/authorization:/password:/credential:), + // which covers providers that issue prefix-less keys (Cohere, Cloudflare...). + maskKeyValuePattern = regexp.MustCompile(`(?i)(['"]?(?:api[_-]?key|key|token|secret|authorization|password|credential)['"]?\s*[:=]\s*['"]?)([A-Za-z0-9_\-\.]{12,})(['"]?)`) + + // maskBearerPattern masks "Bearer " tokens. + maskBearerPattern = regexp.MustCompile(`(?i)(bearer\s+)([A-Za-z0-9_\-\.]{12,})`) + + // maskInvalidKeyPattern masks prefix-less values after "invalid/bad/wrong/ + // unauthorized/incorrect key " phrasing (no colon). + maskInvalidKeyPattern = regexp.MustCompile(`(?i)\b(invalid|bad|wrong|unauthorized|incorrect|missing)\s+key\s+([A-Za-z0-9_\-\.]{12,})`) + + // maskPemBlockPattern masks PEM private key blocks, including the + // JSON-escaped form (\\n) used when a service-account key JSON is echoed + // back inside a string. Covers RSA/EC/Ed25519/OPENSSH/ENCRYPTED keys. + maskPemBlockPattern = regexp.MustCompile(`(?s)-----BEGIN [A-Z ]*PRIVATE KEY-----.*?-----END [A-Z ]*PRIVATE KEY-----`) + + // maskServiceAccountFieldPattern masks prefix-less long values adjacent to + // service-account credential fields (private_key, client_secret, client_id) + // and handles JSON-escaped quotes (\"field\":\"value\") inside nested + // strings — the base key-value pattern misses both. + maskServiceAccountFieldPattern = regexp.MustCompile(`(?i)((?:\\?["'])?(?:private[_-]?key|client[_-]?secret|client[_-]?id)(?:\\?["'])?\s*[:=]\s*(?:\\?["'])?)([A-Za-z0-9_\-\.]{12,})(?:\\?["'])?`) ) // maskHostTail returns the tail parts of a domain/host that should be preserved. @@ -130,5 +165,60 @@ func MaskSensitiveInfo(str string) string { // Mask API keys (e.g., "api_key:AIzaSyAAAaUooTUni8AdaOkSRMda30n_Q4vrV70" -> "api_key:***") str = maskApiKeyPattern.ReplaceAllString(str, "${1}api_key:***${3}") + // Mask bare key formats (defense in depth; upstreams may echo the gateway's + // Authorization value verbatim in error messages — F-13 layer B). + str = maskBareKeyPattern.ReplaceAllStringFunc(str, func(k string) string { + if len(k) > 6 { + return k[:4] + "***" + } + return k + }) + + // F-43: additional prefixed formats, JWTs, and key-word-adjacent values. + str = maskMoreKeyPrefixPattern.ReplaceAllStringFunc(str, func(k string) string { + if len(k) > 6 { + return k[:4] + "***" + } + return k + }) + str = maskJwtPattern.ReplaceAllString(str, "eyJ***") + str = maskKeyValuePattern.ReplaceAllString(str, "${1}***${3}") + str = maskBearerPattern.ReplaceAllString(str, "${1}***") + str = maskInvalidKeyPattern.ReplaceAllString(str, "${1} key ***") + str = maskServiceAccountFieldPattern.ReplaceAllString(str, "${1}***") + // F-67: PEM private key blocks (Vertex service-account keys and similar + // channel credentials) must never be echoed to clients even when the + // surrounding JSON field names were masked. + str = maskPemBlockPattern.ReplaceAllStringFunc(str, func(block string) string { + if len(block) > 30 { + return block[:20] + "***[REDACTED PEM]***" + } + return "***[REDACTED PEM]***" + }) + + return str +} + +// MaskSensitiveKeys masks only credential-shaped values (API keys, tokens, +// JWTs, "key:/token:" values) without URL/domain/IP rewriting, so enum-like +// fields (type/code/param in SSE events) are never mangled (F-20 regression). +func MaskSensitiveKeys(str string) string { + str = maskApiKeyPattern.ReplaceAllString(str, "${1}api_key:***${3}") + str = maskBareKeyPattern.ReplaceAllStringFunc(str, func(k string) string { + if len(k) > 6 { + return k[:4] + "***" + } + return k + }) + str = maskMoreKeyPrefixPattern.ReplaceAllStringFunc(str, func(k string) string { + if len(k) > 6 { + return k[:4] + "***" + } + return k + }) + str = maskJwtPattern.ReplaceAllString(str, "eyJ***") + str = maskKeyValuePattern.ReplaceAllString(str, "${1}***${3}") + str = maskBearerPattern.ReplaceAllString(str, "${1}***") + str = maskInvalidKeyPattern.ReplaceAllString(str, "${1} key ***") return str } diff --git a/relaykit/relayconvert/kitutil/mask_formats_test.go b/relaykit/relayconvert/kitutil/mask_formats_test.go new file mode 100644 index 0000000000..5159ea6f9e --- /dev/null +++ b/relaykit/relayconvert/kitutil/mask_formats_test.go @@ -0,0 +1,50 @@ +package kitutil + +import ( + "strings" + "testing" +) + +func TestMaskSensitiveInfoKeyFormats(t *testing.T) { + cases := []struct { + name string + in string + }{ + {"openai sk-", "invalid key sk-abc1234567890"}, + {"openrouter sk-or-v1-", "invalid key sk-or-v1-abcdefghijklmnop"}, + {"anthropic sk-ant-", "invalid key sk-ant-api03-abcdefghijklmnop"}, + {"groq gsk_", "invalid key gsk_abcdefghijklmnop"}, + {"hf hf_", "invalid key hf_abcdefghijklmnop"}, + {"xai xai-", "invalid key xai-abcdefghijklmnop"}, + {"gemini AIza", "invalid key AIzaSyABCDEFGHIJKLMNOPQRSTUVWXYZ012345"}, + {"perplexity pplx-", "invalid key pplx-abcdefghijklmnop"}, + {"nvidia nvapi-", "invalid key nvapi-abcdefghijklmnop"}, + {"replicate r8_", "invalid key r8_abcdefghijklmnop"}, + {"cohere no-prefix 40", "invalid key abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP"}, + {"cloudflare no-prefix 40", "invalid key 0123456789abcdef0123456789abcdef01234567"}, + {"jwt eyJ", "invalid token eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.signature1234567890"}, + {"token-colon no prefix", "invalid token: abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP"}, + {"key-colon no prefix", "invalid key: abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP"}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + out := MaskSensitiveInfo(tc.in) + original := tc.in + if out == original { + t.Logf("LEAK (unchanged): %s", tc.in) + } else { + t.Logf("masked: %s -> %s", tc.in, out) + } + // Heuristic leak check: the 12-char tail of the secret must not survive. + parts := strings.Fields(tc.in) + secret := parts[len(parts)-1] + if len(secret) > 12 { + tail := secret[len(secret)-12:] + if strings.Contains(out, tail) { + t.Errorf("secret tail %q still present in output: %q", tail, out) + } + } + }) + } +} diff --git a/relaykit/types/error.go b/relaykit/types/error.go index 387fdad769..a9c129e94f 100644 --- a/relaykit/types/error.go +++ b/relaykit/types/error.go @@ -202,7 +202,15 @@ func (e *NewAPIError) ToOpenAIError() OpenAIError { } } if e.errorCode != ErrorCodeCountTokenFailed { - result.Message = kitutil.MaskSensitiveInfo(result.Message) + // F-20a: upstreams may echo the channel key in non-message fields + // (code/type/param/metadata). Mask the serialized error as a whole so + // every field is covered, not just Message. + if raw, err := json.Marshal(result); err == nil { + var masked OpenAIError + if err := json.Unmarshal([]byte(kitutil.MaskSensitiveInfo(string(raw))), &masked); err == nil { + result = masked + } + } } if result.Message == "" { result.Message = string(e.errorType) @@ -231,7 +239,12 @@ func (e *NewAPIError) ToClaudeError() ClaudeError { } } if e.errorCode != ErrorCodeCountTokenFailed { - result.Message = kitutil.MaskSensitiveInfo(result.Message) + if raw, err := json.Marshal(result); err == nil { + var masked ClaudeError + if err := json.Unmarshal([]byte(kitutil.MaskSensitiveInfo(string(raw))), &masked); err == nil { + result = masked + } + } } if result.Message == "" { result.Message = string(e.errorType) diff --git a/service/error.go b/service/error.go index f14f1bbab6..d43052ba1a 100644 --- a/service/error.go +++ b/service/error.go @@ -200,11 +200,13 @@ func TaskErrorWrapperLocal(err error, code string, statusCode int) *taskdto.Task func TaskErrorWrapper(err error, code string, statusCode int) *taskdto.TaskError { text := err.Error() lowerText := strings.ToLower(text) + // F-42: mask unconditionally. Task upstreams (suno/video/audio) may echo the + // channel API key in error bodies that do not contain "post"/"dial"/"http", + // which previously skipped masking entirely. if strings.Contains(lowerText, "post") || strings.Contains(lowerText, "dial") || strings.Contains(lowerText, "http") { common.SysLog(fmt.Sprintf("error: %s", text)) - //text = "请求上游地址失败" - text = common.MaskSensitiveInfo(text) } + text = common.MaskSensitiveInfo(text) //避免暴露内部错误 taskError := &taskdto.TaskError{ Code: code, @@ -223,7 +225,7 @@ func TaskErrorFromAPIError(apiErr *types.NewAPIError) *taskdto.TaskError { } return &taskdto.TaskError{ Code: string(apiErr.GetErrorCode()), - Message: apiErr.Err.Error(), + Message: common.MaskSensitiveInfo(apiErr.Err.Error()), StatusCode: apiErr.StatusCode, Error: apiErr.Err, } diff --git a/service/http.go b/service/http.go index d9818cfe15..2946b45cde 100644 --- a/service/http.go +++ b/service/http.go @@ -38,6 +38,24 @@ func ShouldCopyUpstreamHeader(c *gin.Context, k string, v []string) bool { } return false } + // F-66: never forward credential-bearing response headers from the upstream + // to the client. A misbehaving/malicious upstream may echo the channel key + // (or any credential it received) back in an Authorization / X-Api-Key / + // X-Goog-Api-Key response header; IOCopyBytesGracefully copies all + // non-blocked response headers verbatim. + switch strings.ToLower(k) { + case "authorization", + "proxy-authorization", + "x-api-key", + "x-goog-api-key", + "x-goog-auth", + "api-key", + "apikey", + "x-auth-token", + "x-access-token", + "x-amz-security-token": + return false + } return true } From 11ad5456bf90ca64770b1de46c5f918f2e25a8b1 Mon Sep 17 00:00:00 2001 From: linseasea Date: Mon, 17 Aug 2026 09:38:19 +0800 Subject: [PATCH 2/2] fix(relay): address review feedback on credential masking (per-field mask, word-boundary anchors, keys-only for 2xx, no content rewrite, nil-guard, explicit tests) --- relay/alpha_search_handler.go | 4 +- relay/helper/stream_scanner.go | 6 +++ relay/mjproxy_handler.go | 12 ++--- relaykit/relayconvert/kitutil/mask.go | 22 ++++++--- .../relayconvert/kitutil/mask_formats_test.go | 48 +++++++------------ relaykit/types/error.go | 24 +++++----- service/error.go | 2 +- 7 files changed, 61 insertions(+), 57 deletions(-) diff --git a/relay/alpha_search_handler.go b/relay/alpha_search_handler.go index 9bd77fce16..b04c81a9b6 100644 --- a/relay/alpha_search_handler.go +++ b/relay/alpha_search_handler.go @@ -104,7 +104,9 @@ func AlphaSearchHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError if err != nil { return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError) } - responseBody = []byte(common.MaskSensitiveInfo(string(responseBody))) + // F-65/F-6826: credential-only masking for successful bodies so resource + // URLs in search results are preserved while channel keys are masked. + responseBody = []byte(common.MaskSensitiveKeys(string(responseBody))) c.Writer.WriteHeader(httpResp.StatusCode) if _, err := c.Writer.Write(responseBody); err != nil { return types.NewError(err, types.ErrorCodeDoRequestFailed, types.ErrOptionWithSkipRetry()) diff --git a/relay/helper/stream_scanner.go b/relay/helper/stream_scanner.go index bd920310c1..934c4ca4ba 100644 --- a/relay/helper/stream_scanner.go +++ b/relay/helper/stream_scanner.go @@ -164,6 +164,12 @@ func MaskStreamErrorData(data string) string { if err := json.Unmarshal(raw, &obj); err == nil { objChanged := false for k, v := range obj { + // Never rewrite assistant content fields: masking is for + // error metadata, and re-marshalling a completion chunk + // can corrupt valid output (F-6826). + if k == "content" || k == "delta" || k == "text" || k == "input_text" { + continue + } nv, ch := walk(v) if ch { obj[k] = nv diff --git a/relay/mjproxy_handler.go b/relay/mjproxy_handler.go index 93a7fa7d33..ea568a2dee 100644 --- a/relay/mjproxy_handler.go +++ b/relay/mjproxy_handler.go @@ -267,7 +267,7 @@ func RelaySwapFace(c *gin.Context, info *relaycommon.RelayInfo) *dto.MidjourneyR PromptEn: "", // F-64: mask upstream-controlled text before persistence; task fetch // replays these fields to the client verbatim. - Description: common.MaskSensitiveInfo(midjResponse.Description), + Description: common.MaskSensitiveKeys(midjResponse.Description), State: "", SubmitTime: info.StartTime.UnixNano() / int64(time.Millisecond), StartTime: time.Now().UnixNano() / int64(time.Millisecond), @@ -325,8 +325,8 @@ func RelayMidjourneyTaskImageSeed(c *gin.Context) *dto.MidjourneyResponse { return service.MidjourneyErrorWrapper(constant.MjRequestError, "unmarshal_response_body_failed") } // F-64: the marshaled upstream response is copied to the client verbatim; - // apply the same key-masking used by the OpenAI/Claude error paths. - respBody = []byte(common.MaskSensitiveInfo(string(respBody))) + // use credential-only masking so image/resource URLs survive. + respBody = []byte(common.MaskSensitiveKeys(string(respBody))) service.IOCopyBytesGracefully(c, nil, respBody) return nil } @@ -584,7 +584,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayInfo *relaycommon.RelayInfo) *dt Prompt: midjRequest.Prompt, PromptEn: "", // F-64: mask upstream-controlled text before persistence. - Description: common.MaskSensitiveInfo(midjResponse.Description), + Description: common.MaskSensitiveKeys(midjResponse.Description), State: "", SubmitTime: time.Now().UnixNano() / int64(time.Millisecond), StartTime: 0, @@ -610,7 +610,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayInfo *relaycommon.RelayInfo) *dt //非1-提交成功,21-任务已存在和22-排队中,则记录错误原因 // F-64: upstream error bodies may echo the channel key (mj-api-secret); // mask before persistence so /task/:id/fetch cannot replay it. - midjourneyTask.FailReason = common.MaskSensitiveInfo(midjResponse.Description) + midjourneyTask.FailReason = common.MaskSensitiveKeys(midjResponse.Description) consumeQuota = false } @@ -657,7 +657,7 @@ func RelayMidjourneySubmit(c *gin.Context, relayInfo *relaycommon.RelayInfo) *dt //resp.Body = io.NopCloser(bytes.NewBuffer(responseBody)) // F-64: raw upstream body is relayed verbatim; mask it before writing so an // upstream that echoes mj-api-secret (or any channel key) cannot leak it. - responseBody = []byte(common.MaskSensitiveInfo(string(responseBody))) + responseBody = []byte(common.MaskSensitiveKeys(string(responseBody))) bodyReader := io.NopCloser(bytes.NewBuffer(responseBody)) //for k, v := range resp.Header { diff --git a/relaykit/relayconvert/kitutil/mask.go b/relaykit/relayconvert/kitutil/mask.go index b482bd4284..d8c6db0d0a 100644 --- a/relaykit/relayconvert/kitutil/mask.go +++ b/relaykit/relayconvert/kitutil/mask.go @@ -15,23 +15,23 @@ var ( // maskBareKeyPattern masks bare upstream API key formats that providers may // echo in error messages (sk-..., gsk_..., hf_..., xai-..., AIza...). - maskBareKeyPattern = regexp.MustCompile(`(?i)(sk-[A-Za-z0-9_\-]{8,}|gsk_[A-Za-z0-9_\-]{8,}|hf_[A-Za-z0-9_\-]{8,}|xai-[A-Za-z0-9_\-]{8,}|AIza[A-Za-z0-9_\-]{8,})`) + maskBareKeyPattern = regexp.MustCompile(`(?i)(\bsk-[A-Za-z0-9_\-]{8,}|\bgsk_[A-Za-z0-9_\-]{8,}|\bhf_[A-Za-z0-9_\-]{8,}|\bxai-[A-Za-z0-9_\-]{8,}|\bAIza[A-Za-z0-9_\-]{8,})`) // maskMoreKeyPrefixPattern covers additional provider key prefixes that // F-13's bare pattern missed (Perplexity, NVIDIA NIM, Replicate, personal // access tokens, GitHub tokens, GitLab tokens). - maskMoreKeyPrefixPattern = regexp.MustCompile(`(?i)(pplx-[A-Za-z0-9_\-]{8,}|nvapi-[A-Za-z0-9_\-]{8,}|r8_[A-Za-z0-9_\-]{8,}|pat_[A-Za-z0-9_\-]{8,}|rt_[A-Za-z0-9_\-]{8,}|ghp_[A-Za-z0-9_\-]{8,}|gho_[A-Za-z0-9_\-]{8,}|github_pat_[A-Za-z0-9_\-]{8,}|glpat-[A-Za-z0-9_\-]{8,})`) + maskMoreKeyPrefixPattern = regexp.MustCompile(`(?i)(\bpplx-[A-Za-z0-9_\-]{8,}|\bnvapi-[A-Za-z0-9_\-]{8,}|\br8_[A-Za-z0-9_\-]{8,}|\bpat_[A-Za-z0-9_\-]{8,}|\brt_[A-Za-z0-9_\-]{8,}|\bghp_[A-Za-z0-9_\-]{8,}|\bgho_[A-Za-z0-9_\-]{8,}|\bgithub_pat_[A-Za-z0-9_\-]{8,}|\bglpat-[A-Za-z0-9_\-]{8,})`) // maskJwtPattern masks JWT-shaped credentials (header.payload.signature). - maskJwtPattern = regexp.MustCompile(`(?i)eyJ[A-Za-z0-9_\-]+\.[A-Za-z0-9_\-]+\.[A-Za-z0-9_\-]+`) + maskJwtPattern = regexp.MustCompile(`(?i)\beyJ[A-Za-z0-9_\-]+\.[A-Za-z0-9_\-]+\.[A-Za-z0-9_\-]+`) // maskKeyValuePattern masks prefix-less long values adjacent to credential // field names (api_key:/token:/secret:/authorization:/password:/credential:), // which covers providers that issue prefix-less keys (Cohere, Cloudflare...). - maskKeyValuePattern = regexp.MustCompile(`(?i)(['"]?(?:api[_-]?key|key|token|secret|authorization|password|credential)['"]?\s*[:=]\s*['"]?)([A-Za-z0-9_\-\.]{12,})(['"]?)`) + maskKeyValuePattern = regexp.MustCompile(`(?i)((?:\\?["'])?\b(?:api[_-]?key|key|token|secret|authorization|password|credential)(?:\\?["'])?\s*[:=]\s*(?:\\?["'])?)([A-Za-z0-9_\-\.]{12,})(?:\\?["'])?`) // maskBearerPattern masks "Bearer " tokens. - maskBearerPattern = regexp.MustCompile(`(?i)(bearer\s+)([A-Za-z0-9_\-\.]{12,})`) + maskBearerPattern = regexp.MustCompile(`(?i)(\bbearer\s+)([A-Za-z0-9_\-\.]{12,})`) // maskInvalidKeyPattern masks prefix-less values after "invalid/bad/wrong/ // unauthorized/incorrect key " phrasing (no colon). @@ -46,7 +46,7 @@ var ( // service-account credential fields (private_key, client_secret, client_id) // and handles JSON-escaped quotes (\"field\":\"value\") inside nested // strings — the base key-value pattern misses both. - maskServiceAccountFieldPattern = regexp.MustCompile(`(?i)((?:\\?["'])?(?:private[_-]?key|client[_-]?secret|client[_-]?id)(?:\\?["'])?\s*[:=]\s*(?:\\?["'])?)([A-Za-z0-9_\-\.]{12,})(?:\\?["'])?`) + maskServiceAccountFieldPattern = regexp.MustCompile(`(?i)((?:\\?["'])?\b(?:private[_-]?key|client[_-]?secret|client[_-]?id)(?:\\?["'])?\s*[:=]\s*(?:\\?["'])?)([A-Za-z0-9_\-\.]{12,})(?:\\?["'])?`) ) // maskHostTail returns the tail parts of a domain/host that should be preserved. @@ -220,5 +220,15 @@ func MaskSensitiveKeys(str string) string { str = maskKeyValuePattern.ReplaceAllString(str, "${1}***${3}") str = maskBearerPattern.ReplaceAllString(str, "${1}***") str = maskInvalidKeyPattern.ReplaceAllString(str, "${1} key ***") + str = maskServiceAccountFieldPattern.ReplaceAllString(str, "${1}***") + // F-67: PEM private key blocks (Vertex service-account keys and similar + // channel credentials) must never be echoed to clients even when the + // surrounding JSON field names were masked. + str = maskPemBlockPattern.ReplaceAllStringFunc(str, func(block string) string { + if len(block) > 30 { + return block[:20] + "***[REDACTED PEM]***" + } + return "***[REDACTED PEM]***" + }) return str } diff --git a/relaykit/relayconvert/kitutil/mask_formats_test.go b/relaykit/relayconvert/kitutil/mask_formats_test.go index 5159ea6f9e..cbade1850b 100644 --- a/relaykit/relayconvert/kitutil/mask_formats_test.go +++ b/relaykit/relayconvert/kitutil/mask_formats_test.go @@ -1,7 +1,6 @@ package kitutil import ( - "strings" "testing" ) @@ -9,41 +8,30 @@ func TestMaskSensitiveInfoKeyFormats(t *testing.T) { cases := []struct { name string in string + want string }{ - {"openai sk-", "invalid key sk-abc1234567890"}, - {"openrouter sk-or-v1-", "invalid key sk-or-v1-abcdefghijklmnop"}, - {"anthropic sk-ant-", "invalid key sk-ant-api03-abcdefghijklmnop"}, - {"groq gsk_", "invalid key gsk_abcdefghijklmnop"}, - {"hf hf_", "invalid key hf_abcdefghijklmnop"}, - {"xai xai-", "invalid key xai-abcdefghijklmnop"}, - {"gemini AIza", "invalid key AIzaSyABCDEFGHIJKLMNOPQRSTUVWXYZ012345"}, - {"perplexity pplx-", "invalid key pplx-abcdefghijklmnop"}, - {"nvidia nvapi-", "invalid key nvapi-abcdefghijklmnop"}, - {"replicate r8_", "invalid key r8_abcdefghijklmnop"}, - {"cohere no-prefix 40", "invalid key abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP"}, - {"cloudflare no-prefix 40", "invalid key 0123456789abcdef0123456789abcdef01234567"}, - {"jwt eyJ", "invalid token eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.signature1234567890"}, - {"token-colon no prefix", "invalid token: abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP"}, - {"key-colon no prefix", "invalid key: abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP"}, + {"openai sk-", "invalid key sk-abc1234567890", "invalid key sk-a***"}, + {"openrouter sk-or-v1-", "invalid key sk-or-v1-abcdefghijklmnop", "invalid key sk-o***"}, + {"anthropic sk-ant-", "invalid key sk-ant-api03-abcdefghijklmnop", "invalid key sk-a***"}, + {"groq gsk_", "invalid key gsk_abcdefghijklmnop", "invalid key gsk_***"}, + {"hf hf_", "invalid key hf_abcdefghijklmnop", "invalid key hf_a***"}, + {"xai xai-", "invalid key xai-abcdefghijklmnop", "invalid key xai-***"}, + {"gemini AIza", "invalid key AIzaSyABCDEFGHIJKLMNOPQRSTUVWXYZ012345", "invalid key AIza***"}, + {"perplexity pplx-", "invalid key pplx-abcdefghijklmnop", "invalid key pplx***"}, + {"nvidia nvapi-", "invalid key nvapi-abcdefghijklmnop", "invalid key nvap***"}, + {"replicate r8_", "invalid key r8_abcdefghijklmnop", "invalid key r8_a***"}, + {"cohere no-prefix 40", "invalid key abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP", "invalid key ***"}, + {"cloudflare no-prefix 40", "invalid key 0123456789abcdef0123456789abcdef01234567", "invalid key ***"}, + {"jwt eyJ", "invalid token eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.signature1234567890", "invalid token eyJ***"}, + {"token-colon no prefix", "invalid token: abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP", "invalid token: ***"}, + {"key-colon no prefix", "invalid key: abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOP", "invalid key: ***"}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { out := MaskSensitiveInfo(tc.in) - original := tc.in - if out == original { - t.Logf("LEAK (unchanged): %s", tc.in) - } else { - t.Logf("masked: %s -> %s", tc.in, out) - } - // Heuristic leak check: the 12-char tail of the secret must not survive. - parts := strings.Fields(tc.in) - secret := parts[len(parts)-1] - if len(secret) > 12 { - tail := secret[len(secret)-12:] - if strings.Contains(out, tail) { - t.Errorf("secret tail %q still present in output: %q", tail, out) - } + if out != tc.want { + t.Errorf("mask mismatch\n in: %q\n got: %q\nwant: %q", tc.in, out, tc.want) } }) } diff --git a/relaykit/types/error.go b/relaykit/types/error.go index a9c129e94f..6441eafb43 100644 --- a/relaykit/types/error.go +++ b/relaykit/types/error.go @@ -203,13 +203,15 @@ func (e *NewAPIError) ToOpenAIError() OpenAIError { } if e.errorCode != ErrorCodeCountTokenFailed { // F-20a: upstreams may echo the channel key in non-message fields - // (code/type/param/metadata). Mask the serialized error as a whole so - // every field is covered, not just Message. - if raw, err := json.Marshal(result); err == nil { - var masked OpenAIError - if err := json.Unmarshal([]byte(kitutil.MaskSensitiveInfo(string(raw))), &masked); err == nil { - result = masked - } + // (code/type/param). Mask Message with URL/domain masking and the + // remaining string fields with credential-only masking so enum-like + // values are never mangled. String masking cannot fail; if the code + // field is not a string it cannot carry a credential and is kept. + result.Message = kitutil.MaskSensitiveInfo(result.Message) + result.Type = kitutil.MaskSensitiveKeys(result.Type) + result.Param = kitutil.MaskSensitiveKeys(result.Param) + if s, ok := result.Code.(string); ok { + result.Code = kitutil.MaskSensitiveKeys(s) } } if result.Message == "" { @@ -239,12 +241,8 @@ func (e *NewAPIError) ToClaudeError() ClaudeError { } } if e.errorCode != ErrorCodeCountTokenFailed { - if raw, err := json.Marshal(result); err == nil { - var masked ClaudeError - if err := json.Unmarshal([]byte(kitutil.MaskSensitiveInfo(string(raw))), &masked); err == nil { - result = masked - } - } + result.Message = kitutil.MaskSensitiveInfo(result.Message) + result.Type = kitutil.MaskSensitiveKeys(result.Type) } if result.Message == "" { result.Message = string(e.errorType) diff --git a/service/error.go b/service/error.go index d43052ba1a..8d3d9f02ac 100644 --- a/service/error.go +++ b/service/error.go @@ -225,7 +225,7 @@ func TaskErrorFromAPIError(apiErr *types.NewAPIError) *taskdto.TaskError { } return &taskdto.TaskError{ Code: string(apiErr.GetErrorCode()), - Message: common.MaskSensitiveInfo(apiErr.Err.Error()), + Message: common.MaskSensitiveInfo(apiErr.Error()), StatusCode: apiErr.StatusCode, Error: apiErr.Err, }