feat: add per-channel HTTP transport controls
This commit is contained in:
+176
-40
@@ -2,6 +2,7 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
@@ -12,6 +13,7 @@ import (
|
||||
|
||||
"github.com/QuantumNous/new-api/common"
|
||||
"github.com/QuantumNous/new-api/logger"
|
||||
"github.com/QuantumNous/new-api/relaykit/dto"
|
||||
"github.com/QuantumNous/new-api/setting/system_setting"
|
||||
|
||||
"golang.org/x/net/proxy"
|
||||
@@ -30,7 +32,7 @@ var (
|
||||
type proxyHTTPClientCache struct {
|
||||
mutex sync.RWMutex
|
||||
clients map[string]*http.Client
|
||||
aliases map[string]string
|
||||
aliases map[string]string // rawProxyURL -> canonicalProxyURL
|
||||
}
|
||||
|
||||
type proxyURLConfig struct {
|
||||
@@ -96,7 +98,7 @@ func newRelayHTTPTransport() *http.Transport {
|
||||
return transport
|
||||
}
|
||||
|
||||
func newRelayHTTPClient(transport *http.Transport) *http.Client {
|
||||
func newRelayHTTPClient(transport http.RoundTripper) *http.Client {
|
||||
client := &http.Client{
|
||||
Transport: transport,
|
||||
CheckRedirect: checkRedirect,
|
||||
@@ -107,10 +109,14 @@ func newRelayHTTPClient(transport *http.Transport) *http.Client {
|
||||
return client
|
||||
}
|
||||
|
||||
func clientCacheKey(proxyCacheKey string, policy HTTPTransportPolicy) string {
|
||||
return proxyCacheKey + "\x00" + policy.cacheKeyPart()
|
||||
}
|
||||
|
||||
func InitHttpClient() {
|
||||
transport := newRelayHTTPTransport()
|
||||
transport.Proxy = http.ProxyFromEnvironment
|
||||
httpClient = newRelayHTTPClient(transport)
|
||||
policy := defaultHTTPTransportPolicy()
|
||||
httpClient = newDirectHTTPClient(policy, nil)
|
||||
proxyClients.store(clientCacheKey("", policy), httpClient)
|
||||
ssrfProtectedHTTPClient = newProtectedFetchHTTPClient()
|
||||
}
|
||||
|
||||
@@ -177,45 +183,73 @@ func ValidateProxyURL(rawProxyURL string) error {
|
||||
return err
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) get(rawCacheKey string) (*http.Client, bool) {
|
||||
func (cache *proxyHTTPClientCache) store(fullKey string, client *http.Client) {
|
||||
cache.mutex.Lock()
|
||||
defer cache.mutex.Unlock()
|
||||
cache.clients[fullKey] = client
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) resolveProxyKey(rawProxyURL string) string {
|
||||
if canonicalKey, ok := cache.aliases[rawProxyURL]; ok {
|
||||
return canonicalKey
|
||||
}
|
||||
return rawProxyURL
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) get(rawProxyURL string, policy HTTPTransportPolicy) (*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]
|
||||
proxyKey := cache.resolveProxyKey(rawProxyURL)
|
||||
client, ok := cache.clients[clientCacheKey(proxyKey, policy)]
|
||||
return client, ok
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) getOrCreate(rawCacheKey string, config *proxyURLConfig) (*http.Client, error) {
|
||||
func (cache *proxyHTTPClientCache) getOrCreate(
|
||||
rawProxyURL string,
|
||||
config *proxyURLConfig,
|
||||
policy HTTPTransportPolicy,
|
||||
factory func() (*http.Client, error),
|
||||
) (*http.Client, error) {
|
||||
cache.mutex.Lock()
|
||||
defer cache.mutex.Unlock()
|
||||
if client, ok := cache.clients[config.cacheKey]; ok {
|
||||
cache.aliases[rawCacheKey] = config.cacheKey
|
||||
|
||||
proxyKey := ""
|
||||
if config != nil {
|
||||
proxyKey = config.cacheKey
|
||||
cache.aliases[rawProxyURL] = proxyKey
|
||||
} else if rawProxyURL != "" {
|
||||
proxyKey = cache.resolveProxyKey(rawProxyURL)
|
||||
}
|
||||
fullKey := clientCacheKey(proxyKey, policy)
|
||||
if client, ok := cache.clients[fullKey]; ok {
|
||||
return client, nil
|
||||
}
|
||||
|
||||
client, err := newProxyHTTPClient(config.parsedURL)
|
||||
client, err := factory()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cache.clients[config.cacheKey] = client
|
||||
cache.aliases[rawCacheKey] = config.cacheKey
|
||||
cache.clients[fullKey] = client
|
||||
return client, nil
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) remove(cacheKey string) *http.Client {
|
||||
func (cache *proxyHTTPClientCache) removeProxy(proxyCacheKey string) []*http.Client {
|
||||
cache.mutex.Lock()
|
||||
defer cache.mutex.Unlock()
|
||||
client := cache.clients[cacheKey]
|
||||
delete(cache.clients, cacheKey)
|
||||
removed := make([]*http.Client, 0)
|
||||
prefix := proxyCacheKey + "\x00"
|
||||
for key, client := range cache.clients {
|
||||
if strings.HasPrefix(key, prefix) {
|
||||
removed = append(removed, client)
|
||||
delete(cache.clients, key)
|
||||
}
|
||||
}
|
||||
for alias, canonicalKey := range cache.aliases {
|
||||
if canonicalKey == cacheKey {
|
||||
if canonicalKey == proxyCacheKey {
|
||||
delete(cache.aliases, alias)
|
||||
}
|
||||
}
|
||||
return client
|
||||
return removed
|
||||
}
|
||||
|
||||
func (cache *proxyHTTPClientCache) reset() map[string]*http.Client {
|
||||
@@ -227,13 +261,11 @@ func (cache *proxyHTTPClientCache) reset() map[string]*http.Client {
|
||||
return oldClients
|
||||
}
|
||||
|
||||
func newProxyHTTPClient(proxyURL *url.URL) (*http.Client, error) {
|
||||
transport := newRelayHTTPTransport()
|
||||
|
||||
func configureProxyTransport(transport *http.Transport, proxyURL *url.URL) error {
|
||||
switch proxyURL.Scheme {
|
||||
case "http", "https":
|
||||
transport.Proxy = http.ProxyURL(proxyURL)
|
||||
|
||||
return nil
|
||||
case "socks5", "socks5h":
|
||||
transport.Proxy = nil
|
||||
forwardDialer := &net.Dialer{
|
||||
@@ -242,31 +274,104 @@ func newProxyHTTPClient(proxyURL *url.URL) (*http.Client, error) {
|
||||
}
|
||||
dialer, err := proxy.FromURL(proxyURL, forwardDialer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
contextDialer, ok := dialer.(proxy.ContextDialer)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("SOCKS proxy dialer does not support context cancellation")
|
||||
return fmt.Errorf("SOCKS proxy dialer does not support context cancellation")
|
||||
}
|
||||
transport.DialContext = contextDialer.DialContext
|
||||
|
||||
return nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported proxy scheme")
|
||||
return fmt.Errorf("unsupported proxy scheme")
|
||||
}
|
||||
}
|
||||
|
||||
return newRelayHTTPClient(transport), nil
|
||||
func newTransportFactory(proxyURL *url.URL, tlsConfig *tls.Config) (func() *http.Transport, error) {
|
||||
// Validate proxy configuration once before creating shard transports.
|
||||
if proxyURL != nil {
|
||||
probe := newRelayHTTPTransport()
|
||||
if err := configureProxyTransport(probe, proxyURL); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return func() *http.Transport {
|
||||
transport := newRelayHTTPTransport()
|
||||
if proxyURL != nil {
|
||||
_ = configureProxyTransport(transport, proxyURL)
|
||||
} else {
|
||||
transport.Proxy = http.ProxyFromEnvironment
|
||||
}
|
||||
if tlsConfig != nil {
|
||||
transport.TLSClientConfig = tlsConfig.Clone()
|
||||
}
|
||||
return transport
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newHTTPClientFromPolicy(policy HTTPTransportPolicy, proxyURL *url.URL, tlsConfig *tls.Config) (*http.Client, error) {
|
||||
factory, err := newTransportFactory(proxyURL, tlsConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return newHTTPClientFromTransportFactory(policy, factory), nil
|
||||
}
|
||||
|
||||
func newHTTPClientFromTransportFactory(policy HTTPTransportPolicy, factory func() *http.Transport) *http.Client {
|
||||
if policy.Shards < 1 {
|
||||
policy.Shards = 1
|
||||
}
|
||||
if policy.Protocol == dto.HTTPProtocolHTTP1 || policy.Shards == 1 {
|
||||
transport := factory()
|
||||
applyHTTPTransportPolicy(transport, policy)
|
||||
return newRelayHTTPClient(transport)
|
||||
}
|
||||
shardedFactory := func() *http.Transport {
|
||||
transport := factory()
|
||||
applyHTTPTransportPolicy(transport, policy)
|
||||
return transport
|
||||
}
|
||||
return newRelayHTTPClient(newShardedRoundTripper(policy, shardedFactory))
|
||||
}
|
||||
|
||||
func newDirectHTTPClient(policy HTTPTransportPolicy, tlsConfig *tls.Config) *http.Client {
|
||||
client, err := newHTTPClientFromPolicy(policy, nil, tlsConfig)
|
||||
if err != nil {
|
||||
// Direct clients cannot fail proxy configuration.
|
||||
transport := newRelayHTTPTransport()
|
||||
applyHTTPTransportPolicy(transport, policy)
|
||||
return newRelayHTTPClient(transport)
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
// newHTTPClientWithPolicyAndTLS is a test seam that builds a never-used transport
|
||||
// stack with the given policy and TLS config (for httptest certificate trust).
|
||||
func newHTTPClientWithPolicyAndTLS(policy HTTPTransportPolicy, tlsConfig *tls.Config) *http.Client {
|
||||
return newDirectHTTPClient(policy, tlsConfig)
|
||||
}
|
||||
|
||||
func newProxyHTTPClient(proxyURL *url.URL) (*http.Client, error) {
|
||||
return newHTTPClientFromPolicy(defaultHTTPTransportPolicy(), proxyURL, nil)
|
||||
}
|
||||
|
||||
// GetHttpClientWithProxy returns the default client or a cached proxy-enabled client.
|
||||
func GetHttpClientWithProxy(rawProxyURL string) (*http.Client, error) {
|
||||
return GetHttpClientWithProxySettings(rawProxyURL, dto.ChannelSettings{})
|
||||
}
|
||||
|
||||
// GetHttpClientWithProxySettings returns a cached HTTP client for the proxy URL and
|
||||
// channel transport settings. Default auto + 1 shard shares the same client pool as
|
||||
// GetHttpClientWithProxy / GetHttpClient for the empty-proxy case.
|
||||
func GetHttpClientWithProxySettings(rawProxyURL string, settings dto.ChannelSettings) (*http.Client, error) {
|
||||
policy := NormalizeHTTPTransportPolicy(settings)
|
||||
trimmedProxyURL := strings.TrimSpace(rawProxyURL)
|
||||
|
||||
if trimmedProxyURL == "" {
|
||||
if client := GetHttpClient(); client != nil {
|
||||
return client, nil
|
||||
}
|
||||
return http.DefaultClient, nil
|
||||
return getOrCreateDirectClient(policy)
|
||||
}
|
||||
if client, ok := proxyClients.get(trimmedProxyURL); ok {
|
||||
|
||||
if client, ok := proxyClients.get(trimmedProxyURL, policy); ok {
|
||||
return client, nil
|
||||
}
|
||||
|
||||
@@ -278,10 +383,31 @@ func GetHttpClientWithProxy(rawProxyURL string) (*http.Client, error) {
|
||||
if legacySuffixStripped {
|
||||
warnLegacyProxyURLOnce(config)
|
||||
}
|
||||
return proxyClients.getOrCreate(trimmedProxyURL, config)
|
||||
return proxyClients.getOrCreate(trimmedProxyURL, config, policy, func() (*http.Client, error) {
|
||||
return newHTTPClientFromPolicy(policy, config.parsedURL, nil)
|
||||
})
|
||||
}
|
||||
|
||||
// InvalidateProxyClient removes one proxy client and closes its idle connections.
|
||||
func getOrCreateDirectClient(policy HTTPTransportPolicy) (*http.Client, error) {
|
||||
defaultPolicy := defaultHTTPTransportPolicy()
|
||||
if policy == defaultPolicy {
|
||||
if client := GetHttpClient(); client != nil {
|
||||
return client, nil
|
||||
}
|
||||
// Compatibility with pre-init callers: never assign httpClient outside InitHttpClient.
|
||||
return http.DefaultClient, nil
|
||||
}
|
||||
|
||||
if client, ok := proxyClients.get("", policy); ok {
|
||||
return client, nil
|
||||
}
|
||||
return proxyClients.getOrCreate("", nil, policy, func() (*http.Client, error) {
|
||||
return newDirectHTTPClient(policy, nil), nil
|
||||
})
|
||||
}
|
||||
|
||||
// InvalidateProxyClient removes every cached policy variant for one proxy and
|
||||
// closes their idle connections (including all HTTP/2 shards).
|
||||
func InvalidateProxyClient(rawProxyURL string) {
|
||||
parsedURL, legacySuffixStripped, err := common.ParseProxyURLRuntime(rawProxyURL)
|
||||
if err != nil || parsedURL == nil {
|
||||
@@ -291,16 +417,26 @@ func InvalidateProxyClient(rawProxyURL string) {
|
||||
if legacySuffixStripped {
|
||||
warnLegacyProxyURLOnce(config)
|
||||
}
|
||||
if client := proxyClients.remove(config.cacheKey); client != nil {
|
||||
for _, client := range proxyClients.removeProxy(config.cacheKey) {
|
||||
client.CloseIdleConnections()
|
||||
}
|
||||
}
|
||||
|
||||
// ResetProxyClientCache clears all cached proxy clients.
|
||||
// ResetProxyClientCache clears cached proxy and non-default direct policy clients
|
||||
// and closes idle connections on every transport/shard. The package-level default
|
||||
// httpClient pointer stays stable after InitHttpClient; it is only closed and
|
||||
// re-registered in the policy cache so concurrent GetHttpClient readers never race
|
||||
// a pointer replacement.
|
||||
func ResetProxyClientCache() {
|
||||
defaultClient := httpClient
|
||||
for _, client := range proxyClients.reset() {
|
||||
client.CloseIdleConnections()
|
||||
}
|
||||
if defaultClient == nil {
|
||||
return
|
||||
}
|
||||
defaultClient.CloseIdleConnections()
|
||||
proxyClients.store(clientCacheKey("", defaultHTTPTransportPolicy()), defaultClient)
|
||||
}
|
||||
|
||||
// NewProxyHttpClient is kept for compatibility.
|
||||
|
||||
Reference in New Issue
Block a user