diff --git a/internal/models/asr/openai.go b/internal/models/asr/openai.go index 6244a5e58..66dee48e8 100644 --- a/internal/models/asr/openai.go +++ b/internal/models/asr/openai.go @@ -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 { diff --git a/internal/models/asr/transport.go b/internal/models/asr/transport.go new file mode 100644 index 000000000..638eca28b --- /dev/null +++ b/internal/models/asr/transport.go @@ -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) +} diff --git a/internal/models/asr/transport_security_test.go b/internal/models/asr/transport_security_test.go new file mode 100644 index 000000000..a2d97025a --- /dev/null +++ b/internal/models/asr/transport_security_test.go @@ -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") + } +} diff --git a/internal/models/rerank/aliyun_reranker.go b/internal/models/rerank/aliyun_reranker.go index ccd0c3075..4f8916603 100644 --- a/internal/models/rerank/aliyun_reranker.go +++ b/internal/models/rerank/aliyun_reranker.go @@ -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 } diff --git a/internal/models/rerank/jina_reranker.go b/internal/models/rerank/jina_reranker.go index 27f93af62..c318765c1 100644 --- a/internal/models/rerank/jina_reranker.go +++ b/internal/models/rerank/jina_reranker.go @@ -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 } diff --git a/internal/models/rerank/nvidia_reranker.go b/internal/models/rerank/nvidia_reranker.go index bdc8bd72f..b7ae2a4f3 100644 --- a/internal/models/rerank/nvidia_reranker.go +++ b/internal/models/rerank/nvidia_reranker.go @@ -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 } diff --git a/internal/models/rerank/remote_api.go b/internal/models/rerank/remote_api.go index ce78be20f..105223ae3 100644 --- a/internal/models/rerank/remote_api.go +++ b/internal/models/rerank/remote_api.go @@ -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 } diff --git a/internal/models/rerank/transport.go b/internal/models/rerank/transport.go new file mode 100644 index 000000000..88403d1d8 --- /dev/null +++ b/internal/models/rerank/transport.go @@ -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) +} diff --git a/internal/models/rerank/transport_security_test.go b/internal/models/rerank/transport_security_test.go new file mode 100644 index 000000000..4aa4d17fd --- /dev/null +++ b/internal/models/rerank/transport_security_test.go @@ -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) + } +} diff --git a/internal/models/rerank/weknoracloud.go b/internal/models/rerank/weknoracloud.go index ea985d759..393ae8c23 100644 --- a/internal/models/rerank/weknoracloud.go +++ b/internal/models/rerank/weknoracloud.go @@ -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 } diff --git a/internal/models/rerank/zhipu_reranker.go b/internal/models/rerank/zhipu_reranker.go index 2e2553f2c..c1c70468b 100644 --- a/internal/models/rerank/zhipu_reranker.go +++ b/internal/models/rerank/zhipu_reranker.go @@ -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 } diff --git a/internal/models/vlm/remote_api.go b/internal/models/vlm/remote_api.go index e82b83399..6b5112823 100644 --- a/internal/models/vlm/remote_api.go +++ b/internal/models/vlm/remote_api.go @@ -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, diff --git a/internal/models/vlm/transport.go b/internal/models/vlm/transport.go new file mode 100644 index 000000000..b9bb50af0 --- /dev/null +++ b/internal/models/vlm/transport.go @@ -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) +} diff --git a/internal/models/vlm/transport_security_test.go b/internal/models/vlm/transport_security_test.go new file mode 100644 index 000000000..d1e98fee5 --- /dev/null +++ b/internal/models/vlm/transport_security_test.go @@ -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") + } +} diff --git a/internal/models/vlm/weknoracloud.go b/internal/models/vlm/weknoracloud.go index 98428586d..2fd275b6d 100644 --- a/internal/models/vlm/weknoracloud.go +++ b/internal/models/vlm/weknoracloud.go @@ -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 {