fix(security): harden model client redirects

This commit is contained in:
wizardchen
2026-07-02 18:25:04 +08:00
committed by lyingbug
parent 663d13eb58
commit 1661827d6e
15 changed files with 226 additions and 20 deletions
+5 -2
View File
@@ -5,7 +5,6 @@ import (
"context"
"encoding/json"
"fmt"
"net/http"
"path/filepath"
"strings"
"time"
@@ -30,11 +29,15 @@ type OpenAIASR struct {
// NewOpenAIASR creates an OpenAI-compatible ASR instance.
func NewOpenAIASR(config *Config) (*OpenAIASR, error) {
if err := validateASRBaseURL(config.BaseURL); err != nil {
return nil, err
}
apiCfg := openai.DefaultConfig(config.APIKey)
if config.BaseURL != "" {
apiCfg.BaseURL = config.BaseURL
}
httpClient := &http.Client{Timeout: asrDefaultTimeout}
httpClient := newASRHTTPClient(asrDefaultTimeout)
// 注入用户自定义 HTTP header(类似 OpenAI Python SDK 的 extra_headers
if len(config.CustomHeaders) > 0 {
+25
View File
@@ -0,0 +1,25 @@
package asr
import (
"fmt"
"net/http"
"time"
secutils "github.com/Tencent/WeKnora/internal/utils"
)
func validateASRBaseURL(baseURL string) error {
if baseURL == "" {
return nil
}
if err := secutils.ValidateURLForSSRF(baseURL); err != nil {
return fmt.Errorf("base URL SSRF check failed: %w", err)
}
return nil
}
func newASRHTTPClient(timeout time.Duration) *http.Client {
cfg := secutils.DefaultSSRFSafeHTTPClientConfig()
cfg.Timeout = timeout
return secutils.NewSSRFSafeHTTPClient(cfg)
}
@@ -0,0 +1,26 @@
package asr
import (
"testing"
secutils "github.com/Tencent/WeKnora/internal/utils"
)
func withASRSSRFWhitelist(t *testing.T, raw string) {
t.Helper()
t.Setenv("SSRF_WHITELIST", raw)
secutils.ResetSSRFWhitelistForTest()
t.Cleanup(secutils.ResetSSRFWhitelistForTest)
}
func TestOpenAIASRRejectsInternalBaseURL(t *testing.T) {
withASRSSRFWhitelist(t, "")
_, err := NewOpenAIASR(&Config{
BaseURL: "http://169.254.169.254/latest/meta-data/",
ModelName: "asr-test",
})
if err == nil {
t.Fatalf("NewOpenAIASR returned nil error for blocked internal BaseURL")
}
}
+4 -1
View File
@@ -81,13 +81,16 @@ func NewAliyunReranker(config *RerankerConfig) (*AliyunReranker, error) {
if url := config.BaseURL; url != "" {
baseURL = url
}
if err := validateRerankBaseURL(baseURL); err != nil {
return nil, err
}
return &AliyunReranker{
modelName: config.ModelName,
modelID: config.ModelID,
apiKey: apiKey,
baseURL: baseURL,
client: &http.Client{},
client: newRerankHTTPClient(0),
}, nil
}
+4 -1
View File
@@ -54,13 +54,16 @@ func NewJinaReranker(config *RerankerConfig) (*JinaReranker, error) {
if url := config.BaseURL; url != "" {
baseURL = url
}
if err := validateRerankBaseURL(baseURL); err != nil {
return nil, err
}
return &JinaReranker{
modelName: config.ModelName,
modelID: config.ModelID,
apiKey: apiKey,
baseURL: baseURL,
client: &http.Client{},
client: newRerankHTTPClient(0),
}, nil
}
+5 -5
View File
@@ -27,6 +27,7 @@ type NvidiaReranker struct {
func (r *NvidiaReranker) SetCustomHeaders(headers map[string]string) {
r.customHeaders = headers
}
type NvidiaRerankDocument struct {
Text string `json:"text"`
}
@@ -57,17 +58,16 @@ func NewNvidiaReranker(config *RerankerConfig) (*NvidiaReranker, error) {
if url := config.BaseURL; url != "" {
baseURL = url
}
if err := validateRerankBaseURL(baseURL); err != nil {
return nil, err
}
return &NvidiaReranker{
modelName: config.ModelName,
modelID: config.ModelID,
apiKey: apiKey,
baseURL: baseURL,
client: &http.Client{
Transport: &http.Transport{
Proxy: http.ProxyFromEnvironment,
},
},
client: newRerankHTTPClient(0),
}, nil
}
+4 -1
View File
@@ -56,13 +56,16 @@ func NewOpenAIReranker(config *RerankerConfig) (*OpenAIReranker, error) {
if url := config.BaseURL; url != "" {
baseURL = url
}
if err := validateRerankBaseURL(baseURL); err != nil {
return nil, err
}
return &OpenAIReranker{
modelName: config.ModelName,
modelID: config.ModelID,
apiKey: apiKey,
baseURL: baseURL,
client: &http.Client{},
client: newRerankHTTPClient(0),
}, nil
}
+25
View File
@@ -0,0 +1,25 @@
package rerank
import (
"fmt"
"net/http"
"time"
secutils "github.com/Tencent/WeKnora/internal/utils"
)
func validateRerankBaseURL(baseURL string) error {
if baseURL == "" {
return nil
}
if err := secutils.ValidateURLForSSRF(baseURL); err != nil {
return fmt.Errorf("base URL SSRF check failed: %w", err)
}
return nil
}
func newRerankHTTPClient(timeout time.Duration) *http.Client {
cfg := secutils.DefaultSSRFSafeHTTPClientConfig()
cfg.Timeout = timeout
return secutils.NewSSRFSafeHTTPClient(cfg)
}
@@ -0,0 +1,52 @@
package rerank
import (
stderrors "errors"
"net/http"
"net/http/httptest"
"testing"
secutils "github.com/Tencent/WeKnora/internal/utils"
)
func withRerankSSRFWhitelist(t *testing.T, raw string) {
t.Helper()
t.Setenv("SSRF_WHITELIST", raw)
secutils.ResetSSRFWhitelistForTest()
t.Cleanup(secutils.ResetSSRFWhitelistForTest)
}
func TestOpenAIRerankerRejectsInternalBaseURL(t *testing.T) {
withRerankSSRFWhitelist(t, "")
_, err := NewOpenAIReranker(&RerankerConfig{
BaseURL: "http://169.254.169.254/latest/meta-data/",
ModelName: "rerank-test",
})
if err == nil {
t.Fatalf("NewOpenAIReranker returned nil error for blocked internal BaseURL")
}
}
func TestOpenAIRerankerBlocksRedirectToInternalURL(t *testing.T) {
withRerankSSRFWhitelist(t, "127.0.0.1")
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound)
}))
defer server.Close()
reranker, err := NewOpenAIReranker(&RerankerConfig{
BaseURL: server.URL,
ModelName: "rerank-test",
APIKey: "sk-test",
})
if err != nil {
t.Fatalf("NewOpenAIReranker: %v", err)
}
_, err = reranker.Rerank(t.Context(), "query", []string{"doc"})
if !stderrors.Is(err, secutils.ErrSSRFRedirectBlocked) {
t.Fatalf("Rerank error = %v, want ErrSSRFRedirectBlocked", err)
}
}
+6 -2
View File
@@ -35,6 +35,10 @@ func NewWeKnoraCloudReranker(config *RerankerConfig) (*WeKnoraCloudReranker, err
if config.AppSecret == "" {
return nil, fmt.Errorf("WeKnoraCloud reranker: AppSecret is required")
}
baseURL := strings.TrimRight(config.BaseURL, "/")
if err := validateRerankBaseURL(baseURL); err != nil {
return nil, err
}
remoteModelName := ""
if config.ExtraConfig != nil {
remoteModelName = strings.TrimSpace(config.ExtraConfig["remote_model_name"])
@@ -45,8 +49,8 @@ func NewWeKnoraCloudReranker(config *RerankerConfig) (*WeKnoraCloudReranker, err
modelID: config.ModelID,
appID: config.AppID,
apiKey: config.AppSecret,
baseURL: strings.TrimRight(config.BaseURL, "/"),
client: &http.Client{Timeout: 60 * time.Second},
baseURL: baseURL,
client: newRerankHTTPClient(60 * time.Second),
}, nil
}
+4 -1
View File
@@ -65,13 +65,16 @@ func NewZhipuReranker(config *RerankerConfig) (*ZhipuReranker, error) {
if url := config.BaseURL; url != "" {
baseURL = url
}
if err := validateRerankBaseURL(baseURL); err != nil {
return nil, err
}
return &ZhipuReranker{
modelName: config.ModelName,
modelID: config.ModelID,
apiKey: apiKey,
baseURL: baseURL,
client: &http.Client{},
client: newRerankHTTPClient(0),
}, nil
}
+6 -2
View File
@@ -49,6 +49,10 @@ type RemoteAPIVLM struct {
// NewRemoteAPIVLM creates a remote-API backed VLM instance.
func NewRemoteAPIVLM(config *Config) (*RemoteAPIVLM, error) {
if err := validateVLMBaseURL(config.BaseURL); err != nil {
return nil, err
}
providerName := provider.ProviderName(config.Provider)
if providerName == "" {
providerName = provider.DetectProvider(config.BaseURL)
@@ -73,7 +77,7 @@ func NewRemoteAPIVLM(config *Config) (*RemoteAPIVLM, error) {
apiCfg.BaseURL = config.BaseURL
}
}
httpClient := &http.Client{Timeout: vlmHTTPTimeout()}
httpClient := newVLMHTTPClient(vlmHTTPTimeout())
// 注入用户自定义 HTTP header(类似 OpenAI Python SDK 的 extra_headers
if len(config.CustomHeaders) > 0 {
@@ -105,7 +109,7 @@ func NewRemoteAPIVLM(config *Config) (*RemoteAPIVLM, error) {
// Predict sends an image with a text prompt to the OpenAI-compatible API.
func (v *RemoteAPIVLM) Predict(ctx context.Context, imgBytesList [][]byte, prompt string) (string, error) {
var parts []openai.ChatMessagePart
// Add text prompt first
parts = append(parts, openai.ChatMessagePart{
Type: openai.ChatMessagePartTypeText,
+25
View File
@@ -0,0 +1,25 @@
package vlm
import (
"fmt"
"net/http"
"time"
secutils "github.com/Tencent/WeKnora/internal/utils"
)
func validateVLMBaseURL(baseURL string) error {
if baseURL == "" {
return nil
}
if err := secutils.ValidateURLForSSRF(baseURL); err != nil {
return fmt.Errorf("base URL SSRF check failed: %w", err)
}
return nil
}
func newVLMHTTPClient(timeout time.Duration) *http.Client {
cfg := secutils.DefaultSSRFSafeHTTPClientConfig()
cfg.Timeout = timeout
return secutils.NewSSRFSafeHTTPClient(cfg)
}
@@ -0,0 +1,26 @@
package vlm
import (
"testing"
secutils "github.com/Tencent/WeKnora/internal/utils"
)
func withVLMSSRFWhitelist(t *testing.T, raw string) {
t.Helper()
t.Setenv("SSRF_WHITELIST", raw)
secutils.ResetSSRFWhitelistForTest()
t.Cleanup(secutils.ResetSSRFWhitelistForTest)
}
func TestRemoteAPIVLMRejectsInternalBaseURL(t *testing.T) {
withVLMSSRFWhitelist(t, "")
_, err := NewRemoteAPIVLM(&Config{
BaseURL: "http://169.254.169.254/latest/meta-data/",
ModelName: "vlm-test",
})
if err == nil {
t.Fatalf("NewRemoteAPIVLM returned nil error for blocked internal BaseURL")
}
}
+9 -5
View File
@@ -36,6 +36,10 @@ func NewWeKnoraCloudVLM(config *Config) (*WeKnoraCloudVLM, error) {
if config.AppSecret == "" {
return nil, fmt.Errorf("WeKnoraCloud VLM: AppSecret is required")
}
baseURL := strings.TrimRight(config.BaseURL, "/")
if err := validateVLMBaseURL(baseURL); err != nil {
return nil, err
}
remoteModelName := ""
if config.Extra != nil {
if v, ok := config.Extra["remote_model_name"]; ok {
@@ -50,15 +54,15 @@ func NewWeKnoraCloudVLM(config *Config) (*WeKnoraCloudVLM, error) {
modelID: config.ModelID,
appID: config.AppID,
apiKey: config.AppSecret,
baseURL: strings.TrimRight(config.BaseURL, "/"),
client: &http.Client{Timeout: vlmHTTPTimeout()},
baseURL: baseURL,
client: newVLMHTTPClient(vlmHTTPTimeout()),
}, nil
}
type weKnoraCloudVLMContentPart struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
ImageURL *weKnoraCloudVLMImageURL `json:"image_url,omitempty"`
Type string `json:"type"`
Text string `json:"text,omitempty"`
ImageURL *weKnoraCloudVLMImageURL `json:"image_url,omitempty"`
}
type weKnoraCloudVLMImageURL struct {