fix: align atomic counters on arm

This commit is contained in:
ZacharyZcR
2026-06-01 03:03:40 +08:00
parent 3e4e2db722
commit 8b558b4f12
8 changed files with 161 additions and 150 deletions
+7 -14
View File
@@ -7,7 +7,6 @@ import (
"fmt"
"net"
"net/http"
"sync/atomic"
"time"
)
@@ -24,33 +23,27 @@ func (h *httpDialer) Dial(network, address string) (net.Conn, error) {
func (h *httpDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
start := time.Now()
atomic.AddInt64(&h.stats.TotalConnections, 1)
h.stats.addTotal(1)
// 连接到HTTP代理服务器
proxyConn, err := h.baseDial.DialContext(ctx, NetworkTCP, h.config.Address)
if err != nil {
atomic.AddInt64(&h.stats.FailedConnections, 1)
h.stats.mu.Lock()
h.stats.LastError = err.Error()
h.stats.mu.Unlock()
h.stats.addFailed(1)
h.stats.setLastError(err.Error())
return nil, NewProxyError(ErrTypeConnection, ErrMsgHTTPConnFailed, ErrCodeHTTPConnFailed, err)
}
// 发送CONNECT请求
if err := h.sendConnectRequest(proxyConn, address); err != nil {
_ = proxyConn.Close() // 错误处理路径,Close错误可忽略
atomic.AddInt64(&h.stats.FailedConnections, 1)
h.stats.mu.Lock()
h.stats.LastError = err.Error()
h.stats.mu.Unlock()
h.stats.addFailed(1)
h.stats.setLastError(err.Error())
return nil, err
}
duration := time.Since(start)
h.stats.mu.Lock()
h.stats.LastConnectTime = start
h.stats.mu.Unlock()
atomic.AddInt64(&h.stats.ActiveConnections, 1)
h.stats.setLastConnectTime(start)
h.stats.addActive(1)
h.updateAverageConnectTime(duration)
return &trackedConn{
+13 -36
View File
@@ -5,7 +5,6 @@ import (
"fmt"
"net"
"sync"
"sync/atomic"
"time"
"golang.org/x/net/proxy"
@@ -127,19 +126,7 @@ func (m *manager) Stats() *ProxyStats {
m.mu.RLock()
defer m.mu.RUnlock()
m.stats.mu.Lock()
defer m.stats.mu.Unlock()
return &ProxyStats{
TotalConnections: atomic.LoadInt64(&m.stats.TotalConnections),
ActiveConnections: atomic.LoadInt64(&m.stats.ActiveConnections),
FailedConnections: atomic.LoadInt64(&m.stats.FailedConnections),
AverageConnectTime: m.stats.AverageConnectTime,
LastConnectTime: m.stats.LastConnectTime,
LastError: m.stats.LastError,
ProxyType: m.stats.ProxyType,
ProxyAddress: m.stats.ProxyAddress,
}
return m.stats.snapshot()
}
// createDirectDialer 创建直连拨号器
@@ -246,7 +233,7 @@ func (d *directDialer) Dial(network, address string) (net.Conn, error) {
func (d *directDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
start := time.Now()
atomic.AddInt64(&d.stats.TotalConnections, 1)
d.stats.addTotal(1)
dialer := &net.Dialer{
Timeout: d.timeout,
@@ -263,19 +250,15 @@ func (d *directDialer) DialContext(ctx context.Context, network, address string)
duration := time.Since(start)
d.stats.mu.Lock()
d.stats.LastConnectTime = start
d.stats.mu.Unlock()
d.stats.setLastConnectTime(start)
if err != nil {
atomic.AddInt64(&d.stats.FailedConnections, 1)
d.stats.mu.Lock()
d.stats.LastError = err.Error()
d.stats.mu.Unlock()
d.stats.addFailed(1)
d.stats.setLastError(err.Error())
return nil, NewProxyError(ErrTypeConnection, ErrMsgDirectConnFailed, ErrCodeDirectConnFailed, err)
}
atomic.AddInt64(&d.stats.ActiveConnections, 1)
d.stats.addActive(1)
d.updateAverageConnectTime(duration)
return &trackedConn{
@@ -297,7 +280,7 @@ func (s *socks5Dialer) Dial(network, address string) (net.Conn, error) {
func (s *socks5Dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
start := time.Now()
atomic.AddInt64(&s.stats.TotalConnections, 1)
s.stats.addTotal(1)
// 创建一个带超时的上下文
dialCtx, cancel := context.WithTimeout(ctx, s.config.Timeout)
@@ -325,27 +308,21 @@ func (s *socks5Dialer) DialContext(ctx context.Context, network, address string)
select {
case <-dialCtx.Done():
atomic.AddInt64(&s.stats.FailedConnections, 1)
s.stats.mu.Lock()
s.stats.LastError = dialCtx.Err().Error()
s.stats.mu.Unlock()
s.stats.addFailed(1)
s.stats.setLastError(dialCtx.Err().Error())
return nil, NewProxyError(ErrTypeTimeout, ErrMsgSOCKS5ConnTimeout, ErrCodeSOCKS5ConnTimeout, dialCtx.Err())
case result := <-connChan:
duration := time.Since(start)
s.stats.mu.Lock()
s.stats.LastConnectTime = start
s.stats.mu.Unlock()
s.stats.setLastConnectTime(start)
if result.err != nil {
atomic.AddInt64(&s.stats.FailedConnections, 1)
s.stats.mu.Lock()
s.stats.LastError = result.err.Error()
s.stats.mu.Unlock()
s.stats.addFailed(1)
s.stats.setLastError(result.err.Error())
return nil, NewProxyError(ErrTypeConnection, ErrMsgSOCKS5ConnFailed, ErrCodeSOCKS5ConnFailed, result.err)
}
atomic.AddInt64(&s.stats.ActiveConnections, 1)
s.stats.addActive(1)
s.updateAverageConnectTime(duration)
return &trackedConn{
+8 -10
View File
@@ -49,10 +49,8 @@ func (t *tlsDialerWrapper) DialTLSContext(ctx context.Context, network, address
// 进行TLS握手
if err := tlsConn.Handshake(); err != nil {
_ = tcpConn.Close() // TLS握手失败,Close错误可忽略
atomic.AddInt64(&t.stats.FailedConnections, 1)
t.stats.mu.Lock()
t.stats.LastError = err.Error()
t.stats.mu.Unlock()
t.stats.addFailed(1)
t.stats.setLastError(err.Error())
return nil, NewProxyError(ErrTypeConnection, ErrMsgTLSHandshakeFailed, ErrCodeTLSHandshakeFailed, err)
}
@@ -84,16 +82,16 @@ func (t *tlsDialerWrapper) updateAverageConnectTime(duration time.Duration) {
// trackedConn 带统计的连接
type trackedConn struct {
bytesSent atomic.Int64
bytesRecv atomic.Int64
net.Conn
stats *ProxyStats
bytesSent int64
bytesRecv int64
stats *ProxyStats
}
func (tc *trackedConn) Read(b []byte) (n int, err error) {
n, err = tc.Conn.Read(b)
if n > 0 {
atomic.AddInt64(&tc.bytesRecv, int64(n))
tc.bytesRecv.Add(int64(n))
}
return n, err
}
@@ -101,13 +99,13 @@ func (tc *trackedConn) Read(b []byte) (n int, err error) {
func (tc *trackedConn) Write(b []byte) (n int, err error) {
n, err = tc.Conn.Write(b)
if n > 0 {
atomic.AddInt64(&tc.bytesSent, int64(n))
tc.bytesSent.Add(int64(n))
}
return n, err
}
func (tc *trackedConn) Close() error {
atomic.AddInt64(&tc.stats.ActiveConnections, -1)
tc.stats.addActive(-1)
return tc.Conn.Close()
}
+49 -3
View File
@@ -96,9 +96,9 @@ type ProxyManager interface {
//
//nolint:revive // 保持与现有代码的向后兼容性
type ProxyStats struct {
TotalConnections int64 `json:"total_connections"`
ActiveConnections int64 `json:"active_connections"`
FailedConnections int64 `json:"failed_connections"`
TotalConnections int64 `json:"total_connections"`
ActiveConnections int64 `json:"active_connections"`
FailedConnections int64 `json:"failed_connections"`
mu sync.Mutex `json:"-"`
AverageConnectTime time.Duration `json:"average_connect_time"`
LastConnectTime time.Time `json:"last_connect_time"`
@@ -107,6 +107,52 @@ type ProxyStats struct {
ProxyAddress string `json:"proxy_address"`
}
func (s *ProxyStats) addTotal(delta int64) {
s.mu.Lock()
s.TotalConnections += delta
s.mu.Unlock()
}
func (s *ProxyStats) addActive(delta int64) {
s.mu.Lock()
s.ActiveConnections += delta
s.mu.Unlock()
}
func (s *ProxyStats) addFailed(delta int64) {
s.mu.Lock()
s.FailedConnections += delta
s.mu.Unlock()
}
func (s *ProxyStats) setLastConnectTime(t time.Time) {
s.mu.Lock()
s.LastConnectTime = t
s.mu.Unlock()
}
func (s *ProxyStats) setLastError(err string) {
s.mu.Lock()
s.LastError = err
s.mu.Unlock()
}
func (s *ProxyStats) snapshot() *ProxyStats {
s.mu.Lock()
defer s.mu.Unlock()
return &ProxyStats{
TotalConnections: s.TotalConnections,
ActiveConnections: s.ActiveConnections,
FailedConnections: s.FailedConnections,
AverageConnectTime: s.AverageConnectTime,
LastConnectTime: s.LastConnectTime,
LastError: s.LastError,
ProxyType: s.ProxyType,
ProxyAddress: s.ProxyAddress,
}
}
// ProxyError 代理错误类型
//
//nolint:revive // 保持与现有代码的向后兼容性