mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-09-01 14:53:07 +08:00
fix(security): harden model client redirects
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user