respect per-call session dial timeouts

This commit is contained in:
ZacharyZcR
2026-05-18 16:35:58 +08:00
parent 13f7997d16
commit 856eeccd78
2 changed files with 65 additions and 15 deletions
+28 -13
View File
@@ -24,10 +24,10 @@ type ScanSession struct {
Params *FlagVars // 原始参数,只读 Params *FlagVars // 原始参数,只读
ResultSink ResultSink // 可选,覆盖全局输出 ResultSink ResultSink // 可选,覆盖全局输出
// 每会话 dialer(懒初始化,取决于代理配置) // 每会话 dialer(按 timeout 懒初始化,取决于代理配置)
dialerOnce sync.Once dialerMu sync.Mutex
dialer proxy.Dialer dialers map[time.Duration]proxy.Dialer
dialerErr error dialerErrs map[time.Duration]error
} }
// NewScanSession 从已构建的 Config、State 和 FlagVars 创建会话 // NewScanSession 从已构建的 Config、State 和 FlagVars 创建会话
@@ -96,7 +96,7 @@ func (s *ScanSession) DialTCP(ctx context.Context, network, address string, time
} }
// 获取 dialer // 获取 dialer
dialer, err := s.getDialer() dialer, err := s.getDialer(timeout)
if err != nil { if err != nil {
s.LogError(fmt.Sprintf("获取代理拨号器失败: %v", err)) s.LogError(fmt.Sprintf("获取代理拨号器失败: %v", err))
s.State.IncrementTCPFailedPacketCount() s.State.IncrementTCPFailedPacketCount()
@@ -119,18 +119,33 @@ func (s *ScanSession) DialTCP(ctx context.Context, network, address string, time
return conn, nil return conn, nil
} }
func (s *ScanSession) getDialer() (proxy.Dialer, error) { func (s *ScanSession) getDialer(timeout time.Duration) (proxy.Dialer, error) {
s.dialerOnce.Do(func() { if timeout <= 0 {
cfg := s.createProxyConfig() timeout = s.Config.Timeout
}
s.dialerMu.Lock()
defer s.dialerMu.Unlock()
if s.dialers == nil {
s.dialers = make(map[time.Duration]proxy.Dialer)
s.dialerErrs = make(map[time.Duration]error)
}
if dialer, ok := s.dialers[timeout]; ok {
return dialer, s.dialerErrs[timeout]
}
cfg := s.createProxyConfig(timeout)
manager := proxy.NewProxyManager(cfg) manager := proxy.NewProxyManager(cfg)
s.dialer, s.dialerErr = manager.GetDialer() dialer, err := manager.GetDialer()
}) s.dialers[timeout] = dialer
return s.dialer, s.dialerErr s.dialerErrs[timeout] = err
return dialer, err
} }
func (s *ScanSession) createProxyConfig() *proxy.ProxyConfig { func (s *ScanSession) createProxyConfig(timeout time.Duration) *proxy.ProxyConfig {
cfg := proxy.DefaultProxyConfig() cfg := proxy.DefaultProxyConfig()
cfg.Timeout = s.Config.Timeout cfg.Timeout = timeout
cfg.LocalAddr = s.Config.Network.Iface cfg.LocalAddr = s.Config.Network.Iface
// 优先 SOCKS5 // 优先 SOCKS5
+36 -1
View File
@@ -1,6 +1,9 @@
package common package common
import "testing" import (
"testing"
"time"
)
func TestScanSessionLogMethodsHonorSilentConfig(t *testing.T) { func TestScanSessionLogMethodsHonorSilentConfig(t *testing.T) {
loggerMu.Lock() loggerMu.Lock()
@@ -30,3 +33,35 @@ func TestScanSessionLogMethodsHonorSilentConfig(t *testing.T) {
t.Fatal("silent session log methods initialized global logger") t.Fatal("silent session log methods initialized global logger")
} }
} }
func TestScanSessionDialerCacheIsTimeoutAware(t *testing.T) {
cfg := NewConfig()
cfg.Timeout = 5 * time.Second
session := NewScanSession(cfg, NewState(), &FlagVars{})
shortTimeout := 100 * time.Millisecond
longTimeout := 2 * time.Second
shortDialer, err := session.getDialer(shortTimeout)
if err != nil {
t.Fatal(err)
}
shortDialerAgain, err := session.getDialer(shortTimeout)
if err != nil {
t.Fatal(err)
}
longDialer, err := session.getDialer(longTimeout)
if err != nil {
t.Fatal(err)
}
if shortDialer != shortDialerAgain {
t.Fatal("same timeout should reuse the session dialer")
}
if shortDialer == longDialer {
t.Fatal("different timeouts should not share one session dialer")
}
if got := session.createProxyConfig(shortTimeout).Timeout; got != shortTimeout {
t.Fatalf("proxy timeout = %v, want %v", got, shortTimeout)
}
}