fix(channel): improve proxy client compatibility and cache lifecycle (#6157)
* fix(channel): improve proxy client compatibility and cache lifecycle * test(controller): use non-fatal assertions for channel tests
This commit is contained in:
@@ -97,7 +97,6 @@ func RefreshCodexChannelCredential(ctx context.Context, channelID int, opts Code
|
||||
|
||||
if opts.ResetCaches {
|
||||
model.InitChannelCache()
|
||||
ResetProxyClientCache()
|
||||
}
|
||||
|
||||
return oauthKey, ch, nil
|
||||
|
||||
@@ -139,7 +139,6 @@ func runCodexCredentialAutoRefreshOnce() {
|
||||
}()
|
||||
model.InitChannelCache()
|
||||
}()
|
||||
ResetProxyClientCache()
|
||||
}
|
||||
|
||||
if common.DebugEnabled {
|
||||
|
||||
+207
-106
@@ -6,10 +6,12 @@ import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/QuantumNous/new-api/common"
|
||||
"github.com/QuantumNous/new-api/logger"
|
||||
"github.com/QuantumNous/new-api/setting/system_setting"
|
||||
|
||||
"golang.org/x/net/proxy"
|
||||
@@ -18,10 +20,24 @@ import (
|
||||
var (
|
||||
httpClient *http.Client
|
||||
ssrfProtectedHTTPClient *http.Client
|
||||
proxyClientLock sync.Mutex
|
||||
proxyClients = make(map[string]*http.Client)
|
||||
proxyClients = proxyHTTPClientCache{
|
||||
clients: make(map[string]*http.Client),
|
||||
aliases: make(map[string]string),
|
||||
}
|
||||
legacyProxyURLWarnings sync.Map
|
||||
)
|
||||
|
||||
type proxyHTTPClientCache struct {
|
||||
mutex sync.RWMutex
|
||||
clients map[string]*http.Client
|
||||
aliases map[string]string
|
||||
}
|
||||
|
||||
type proxyURLConfig struct {
|
||||
parsedURL *url.URL
|
||||
cacheKey string
|
||||
}
|
||||
|
||||
func checkRedirect(req *http.Request, via []*http.Request) error {
|
||||
urlStr := req.URL.String()
|
||||
if err := validateURLWithCurrentFetchSetting(urlStr, true); err != nil {
|
||||
@@ -53,30 +69,48 @@ func ValidateSSRFProtectedFetchURL(urlStr string) error {
|
||||
return validateURLWithCurrentFetchSetting(urlStr, true)
|
||||
}
|
||||
|
||||
func InitHttpClient() {
|
||||
transport := &http.Transport{
|
||||
MaxIdleConns: common.RelayMaxIdleConns,
|
||||
MaxIdleConnsPerHost: common.RelayMaxIdleConnsPerHost,
|
||||
IdleConnTimeout: time.Duration(common.RelayIdleConnTimeout) * time.Second,
|
||||
ForceAttemptHTTP2: true,
|
||||
Proxy: http.ProxyFromEnvironment, // Support HTTP_PROXY, HTTPS_PROXY, NO_PROXY env vars
|
||||
func newRelayHTTPTransport() *http.Transport {
|
||||
var transport *http.Transport
|
||||
if defaultTransport, ok := http.DefaultTransport.(*http.Transport); ok && defaultTransport != nil {
|
||||
transport = defaultTransport.Clone()
|
||||
} else {
|
||||
dialer := &net.Dialer{
|
||||
Timeout: 30 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
transport = &http.Transport{
|
||||
Proxy: http.ProxyFromEnvironment,
|
||||
DialContext: dialer.DialContext,
|
||||
ForceAttemptHTTP2: true,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: time.Second,
|
||||
}
|
||||
}
|
||||
transport.MaxIdleConns = common.RelayMaxIdleConns
|
||||
transport.MaxIdleConnsPerHost = common.RelayMaxIdleConnsPerHost
|
||||
transport.IdleConnTimeout = time.Duration(common.RelayIdleConnTimeout) * time.Second
|
||||
transport.ForceAttemptHTTP2 = true
|
||||
if common.TLSInsecureSkipVerify {
|
||||
transport.TLSClientConfig = common.InsecureTLSConfig
|
||||
}
|
||||
return transport
|
||||
}
|
||||
|
||||
if common.RelayTimeout == 0 {
|
||||
httpClient = &http.Client{
|
||||
Transport: transport,
|
||||
CheckRedirect: checkRedirect,
|
||||
}
|
||||
} else {
|
||||
httpClient = &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: time.Duration(common.RelayTimeout) * time.Second,
|
||||
CheckRedirect: checkRedirect,
|
||||
}
|
||||
func newRelayHTTPClient(transport *http.Transport) *http.Client {
|
||||
client := &http.Client{
|
||||
Transport: transport,
|
||||
CheckRedirect: checkRedirect,
|
||||
}
|
||||
if common.RelayTimeout != 0 {
|
||||
client.Timeout = time.Duration(common.RelayTimeout) * time.Second
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
func InitHttpClient() {
|
||||
transport := newRelayHTTPTransport()
|
||||
transport.Proxy = http.ProxyFromEnvironment
|
||||
httpClient = newRelayHTTPClient(transport)
|
||||
ssrfProtectedHTTPClient = newProtectedFetchHTTPClient()
|
||||
}
|
||||
|
||||
@@ -100,110 +134,177 @@ func GetSSRFProtectedHTTPClient() *http.Client {
|
||||
return ssrfProtectedHTTPClient
|
||||
}
|
||||
|
||||
// GetHttpClientWithProxy returns the default client or a proxy-enabled one when proxyURL is provided.
|
||||
func GetHttpClientWithProxy(proxyURL string) (*http.Client, error) {
|
||||
if proxyURL == "" {
|
||||
return GetHttpClient(), nil
|
||||
func newProxyURLConfig(parsedURL *url.URL) *proxyURLConfig {
|
||||
return &proxyURLConfig{
|
||||
parsedURL: parsedURL,
|
||||
cacheKey: parsedURL.String(),
|
||||
}
|
||||
return NewProxyHttpClient(proxyURL)
|
||||
}
|
||||
|
||||
// ResetProxyClientCache 清空代理客户端缓存,确保下次使用时重新初始化
|
||||
func ResetProxyClientCache() {
|
||||
proxyClientLock.Lock()
|
||||
defer proxyClientLock.Unlock()
|
||||
for _, client := range proxyClients {
|
||||
if transport, ok := client.Transport.(*http.Transport); ok && transport != nil {
|
||||
transport.CloseIdleConnections()
|
||||
func warnLegacyProxyURLOnce(config *proxyURLConfig) {
|
||||
if _, loaded := legacyProxyURLWarnings.LoadOrStore(config.cacheKey, struct{}{}); loaded {
|
||||
return
|
||||
}
|
||||
logger.LogWarn(
|
||||
context.Background(),
|
||||
fmt.Sprintf(
|
||||
"legacy proxy URL suffix ignored at runtime: scheme=%s host=%s; update the channel proxy setting",
|
||||
config.parsedURL.Scheme,
|
||||
config.parsedURL.Host,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
// NormalizeProxyURL validates a proxy URL using runtime-compatible rules and returns its canonical cache key.
|
||||
func NormalizeProxyURL(rawProxyURL string) (string, error) {
|
||||
parsedURL, legacySuffixStripped, err := common.ParseProxyURLRuntime(rawProxyURL)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if parsedURL == nil {
|
||||
return "", nil
|
||||
}
|
||||
config := newProxyURLConfig(parsedURL)
|
||||
if legacySuffixStripped {
|
||||
warnLegacyProxyURLOnce(config)
|
||||
}
|
||||
return config.cacheKey, nil
|
||||
}
|
||||
|
||||
// ValidateProxyURL validates a channel proxy URL without connecting to it.
|
||||
func ValidateProxyURL(rawProxyURL string) error {
|
||||
_, err := common.ParseProxyURLStrict(rawProxyURL)
|
||||
return err
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) get(rawCacheKey string) (*http.Client, bool) {
|
||||
cache.mutex.RLock()
|
||||
defer cache.mutex.RUnlock()
|
||||
cacheKey := rawCacheKey
|
||||
if canonicalKey, ok := cache.aliases[rawCacheKey]; ok {
|
||||
cacheKey = canonicalKey
|
||||
}
|
||||
client, ok := cache.clients[cacheKey]
|
||||
return client, ok
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) getOrCreate(rawCacheKey string, config *proxyURLConfig) (*http.Client, error) {
|
||||
cache.mutex.Lock()
|
||||
defer cache.mutex.Unlock()
|
||||
if client, ok := cache.clients[config.cacheKey]; ok {
|
||||
cache.aliases[rawCacheKey] = config.cacheKey
|
||||
return client, nil
|
||||
}
|
||||
|
||||
client, err := newProxyHTTPClient(config.parsedURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cache.clients[config.cacheKey] = client
|
||||
cache.aliases[rawCacheKey] = config.cacheKey
|
||||
return client, nil
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) remove(cacheKey string) *http.Client {
|
||||
cache.mutex.Lock()
|
||||
defer cache.mutex.Unlock()
|
||||
client := cache.clients[cacheKey]
|
||||
delete(cache.clients, cacheKey)
|
||||
for alias, canonicalKey := range cache.aliases {
|
||||
if canonicalKey == cacheKey {
|
||||
delete(cache.aliases, alias)
|
||||
}
|
||||
}
|
||||
proxyClients = make(map[string]*http.Client)
|
||||
return client
|
||||
}
|
||||
|
||||
// NewProxyHttpClient 创建支持代理的 HTTP 客户端
|
||||
func NewProxyHttpClient(proxyURL string) (*http.Client, error) {
|
||||
if proxyURL == "" {
|
||||
func (cache *proxyHTTPClientCache) reset() map[string]*http.Client {
|
||||
cache.mutex.Lock()
|
||||
defer cache.mutex.Unlock()
|
||||
oldClients := cache.clients
|
||||
cache.clients = make(map[string]*http.Client)
|
||||
cache.aliases = make(map[string]string)
|
||||
return oldClients
|
||||
}
|
||||
|
||||
func newProxyHTTPClient(proxyURL *url.URL) (*http.Client, error) {
|
||||
transport := newRelayHTTPTransport()
|
||||
|
||||
switch proxyURL.Scheme {
|
||||
case "http", "https":
|
||||
transport.Proxy = http.ProxyURL(proxyURL)
|
||||
|
||||
case "socks5", "socks5h":
|
||||
transport.Proxy = nil
|
||||
forwardDialer := &net.Dialer{
|
||||
Timeout: 30 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
dialer, err := proxy.FromURL(proxyURL, forwardDialer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
contextDialer, ok := dialer.(proxy.ContextDialer)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("SOCKS proxy dialer does not support context cancellation")
|
||||
}
|
||||
transport.DialContext = contextDialer.DialContext
|
||||
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported proxy scheme")
|
||||
}
|
||||
|
||||
return newRelayHTTPClient(transport), nil
|
||||
}
|
||||
|
||||
// GetHttpClientWithProxy returns the default client or a cached proxy-enabled client.
|
||||
func GetHttpClientWithProxy(rawProxyURL string) (*http.Client, error) {
|
||||
trimmedProxyURL := strings.TrimSpace(rawProxyURL)
|
||||
if trimmedProxyURL == "" {
|
||||
if client := GetHttpClient(); client != nil {
|
||||
return client, nil
|
||||
}
|
||||
return http.DefaultClient, nil
|
||||
}
|
||||
|
||||
proxyClientLock.Lock()
|
||||
if client, ok := proxyClients[proxyURL]; ok {
|
||||
proxyClientLock.Unlock()
|
||||
if client, ok := proxyClients.get(trimmedProxyURL); ok {
|
||||
return client, nil
|
||||
}
|
||||
proxyClientLock.Unlock()
|
||||
|
||||
parsedURL, err := url.Parse(proxyURL)
|
||||
parsedURL, legacySuffixStripped, err := common.ParseProxyURLRuntime(trimmedProxyURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
config := newProxyURLConfig(parsedURL)
|
||||
if legacySuffixStripped {
|
||||
warnLegacyProxyURLOnce(config)
|
||||
}
|
||||
return proxyClients.getOrCreate(trimmedProxyURL, config)
|
||||
}
|
||||
|
||||
switch parsedURL.Scheme {
|
||||
case "http", "https":
|
||||
transport := &http.Transport{
|
||||
MaxIdleConns: common.RelayMaxIdleConns,
|
||||
MaxIdleConnsPerHost: common.RelayMaxIdleConnsPerHost,
|
||||
IdleConnTimeout: time.Duration(common.RelayIdleConnTimeout) * time.Second,
|
||||
ForceAttemptHTTP2: true,
|
||||
Proxy: http.ProxyURL(parsedURL),
|
||||
}
|
||||
if common.TLSInsecureSkipVerify {
|
||||
transport.TLSClientConfig = common.InsecureTLSConfig
|
||||
}
|
||||
client := &http.Client{
|
||||
Transport: transport,
|
||||
CheckRedirect: checkRedirect,
|
||||
}
|
||||
client.Timeout = time.Duration(common.RelayTimeout) * time.Second
|
||||
proxyClientLock.Lock()
|
||||
proxyClients[proxyURL] = client
|
||||
proxyClientLock.Unlock()
|
||||
return client, nil
|
||||
|
||||
case "socks5", "socks5h":
|
||||
// 获取认证信息
|
||||
var auth *proxy.Auth
|
||||
if parsedURL.User != nil {
|
||||
auth = &proxy.Auth{
|
||||
User: parsedURL.User.Username(),
|
||||
Password: "",
|
||||
}
|
||||
if password, ok := parsedURL.User.Password(); ok {
|
||||
auth.Password = password
|
||||
}
|
||||
}
|
||||
|
||||
// 创建 SOCKS5 代理拨号器
|
||||
// proxy.SOCKS5 使用 tcp 参数,所有 TCP 连接包括 DNS 查询都将通过代理进行。行为与 socks5h 相同
|
||||
dialer, err := proxy.SOCKS5("tcp", parsedURL.Host, auth, proxy.Direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
transport := &http.Transport{
|
||||
MaxIdleConns: common.RelayMaxIdleConns,
|
||||
MaxIdleConnsPerHost: common.RelayMaxIdleConnsPerHost,
|
||||
IdleConnTimeout: time.Duration(common.RelayIdleConnTimeout) * time.Second,
|
||||
ForceAttemptHTTP2: true,
|
||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return dialer.Dial(network, addr)
|
||||
},
|
||||
}
|
||||
if common.TLSInsecureSkipVerify {
|
||||
transport.TLSClientConfig = common.InsecureTLSConfig
|
||||
}
|
||||
|
||||
client := &http.Client{Transport: transport, CheckRedirect: checkRedirect}
|
||||
client.Timeout = time.Duration(common.RelayTimeout) * time.Second
|
||||
proxyClientLock.Lock()
|
||||
proxyClients[proxyURL] = client
|
||||
proxyClientLock.Unlock()
|
||||
return client, nil
|
||||
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported proxy scheme: %s, must be http, https, socks5 or socks5h", parsedURL.Scheme)
|
||||
// InvalidateProxyClient removes one proxy client and closes its idle connections.
|
||||
func InvalidateProxyClient(rawProxyURL string) {
|
||||
parsedURL, legacySuffixStripped, err := common.ParseProxyURLRuntime(rawProxyURL)
|
||||
if err != nil || parsedURL == nil {
|
||||
return
|
||||
}
|
||||
config := newProxyURLConfig(parsedURL)
|
||||
if legacySuffixStripped {
|
||||
warnLegacyProxyURLOnce(config)
|
||||
}
|
||||
if client := proxyClients.remove(config.cacheKey); client != nil {
|
||||
client.CloseIdleConnections()
|
||||
}
|
||||
}
|
||||
|
||||
// ResetProxyClientCache clears all cached proxy clients.
|
||||
func ResetProxyClientCache() {
|
||||
for _, client := range proxyClients.reset() {
|
||||
client.CloseIdleConnections()
|
||||
}
|
||||
}
|
||||
|
||||
// NewProxyHttpClient is kept for compatibility.
|
||||
// Deprecated: use GetHttpClientWithProxy.
|
||||
func NewProxyHttpClient(proxyURL string) (*http.Client, error) {
|
||||
return GetHttpClientWithProxy(proxyURL)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user