Improve ping performance

This commit is contained in:
世界 2026-07-10 12:29:13 +08:00
parent 23a39c59be
commit d0d4ebd8db
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
2 changed files with 122 additions and 45 deletions

View file

@ -17,8 +17,11 @@ import (
"github.com/sagernet/sing/common/logger" "github.com/sagernet/sing/common/logger"
) )
const (
// Although its theoretical maximum may be 64k, I dont yet know of any practical use case for that. For memory-usage reasons, Im just using a 2k buffer. // Although its theoretical maximum may be 64k, I dont yet know of any practical use case for that. For memory-usage reasons, Im just using a 2k buffer.
const maxICMPPacketSize = 2048 maxICMPPacketSize = 2048
requestsLimit = 1024
)
type PacketWriter interface { type PacketWriter interface {
WritePacket(packet []byte) error WritePacket(packet []byte) error
@ -33,7 +36,11 @@ type Destination struct {
timeout time.Duration timeout time.Duration
lastActive common.TypedValue[time.Time] lastActive common.TypedValue[time.Time]
requestAccess sync.Mutex requestAccess sync.Mutex
requests map[pingRequest]time.Time requests map[pingRequest]int
requestSlots []trackedPingRequest
requestHead int
requestTail int
requestFree int
} }
type pingRequest struct { type pingRequest struct {
@ -43,6 +50,13 @@ type pingRequest struct {
Sequence uint16 Sequence uint16
} }
type trackedPingRequest struct {
request pingRequest
createdAt time.Time
previous int
next int
}
func ConnectDestination( func ConnectDestination(
ctx context.Context, ctx context.Context,
logger logger.ContextLogger, logger logger.ContextLogger,
@ -74,7 +88,10 @@ func ConnectDestination(
destination: destination, destination: destination,
writer: writer, writer: writer,
timeout: timeout, timeout: timeout,
requests: make(map[pingRequest]time.Time), requests: make(map[pingRequest]int),
requestHead: -1,
requestTail: -1,
requestFree: -1,
} }
d.lastActive.Store(time.Now()) d.lastActive.Store(time.Now())
go d.loopRead() go d.loopRead()
@ -120,10 +137,7 @@ func (d *Destination) loopRead() {
case header.ICMPv4EchoReply: case header.ICMPv4EchoReply:
request := pingRequest{Source: ipHdr.DestinationAddr(), Destination: ipHdr.SourceAddr(), Identifier: icmpHdr.Ident(), Sequence: icmpHdr.Sequence()} request := pingRequest{Source: ipHdr.DestinationAddr(), Destination: ipHdr.SourceAddr(), Identifier: icmpHdr.Ident(), Sequence: icmpHdr.Sequence()}
d.requestAccess.Lock() d.requestAccess.Lock()
_, loaded := d.requests[request] loaded := d.removeRequest(request)
if loaded {
delete(d.requests, request)
}
d.requestAccess.Unlock() d.requestAccess.Unlock()
if !loaded { if !loaded {
continue continue
@ -157,11 +171,7 @@ func (d *Destination) loopRead() {
var requestExists bool var requestExists bool
request := pingRequest{Source: ipHdr.DestinationAddr(), Destination: ipHdr.SourceAddr(), Identifier: icmpHdr.Ident(), Sequence: icmpHdr.Sequence()} request := pingRequest{Source: ipHdr.DestinationAddr(), Destination: ipHdr.SourceAddr(), Identifier: icmpHdr.Ident(), Sequence: icmpHdr.Sequence()}
d.requestAccess.Lock() d.requestAccess.Lock()
_, loaded := d.requests[request] requestExists = d.removeRequest(request)
if loaded {
requestExists = true
delete(d.requests, request)
}
d.requestAccess.Unlock() d.requestAccess.Unlock()
if !requestExists { if !requestExists {
continue continue
@ -254,26 +264,77 @@ func (d *Destination) needFilter() bool {
} }
func (d *Destination) registerRequest(request pingRequest) { func (d *Destination) registerRequest(request pingRequest) {
const requestsLimit = 1024
d.requestAccess.Lock() d.requestAccess.Lock()
defer d.requestAccess.Unlock() defer d.requestAccess.Unlock()
now := time.Now() now := time.Now()
var ( d.pruneRequests(now)
oldestRequest pingRequest if existing, loaded := d.requests[request]; loaded {
oldestCreateAt = now d.removeRequestAt(existing)
) }
for oldRequest, createdAt := range d.requests { if len(d.requests) >= requestsLimit {
if now.Sub(createdAt) > d.timeout { d.removeRequestAt(d.requestHead)
delete(d.requests, oldRequest) }
} else if createdAt.Before(oldestCreateAt) { var requestIndex int
oldestRequest = oldRequest if d.requestFree >= 0 {
oldestCreateAt = createdAt requestIndex = d.requestFree
d.requestFree = d.requestSlots[requestIndex].next
d.requestSlots[requestIndex] = trackedPingRequest{
request: request,
createdAt: now,
previous: d.requestTail,
next: -1,
}
} else {
requestIndex = len(d.requestSlots)
d.requestSlots = append(d.requestSlots, trackedPingRequest{
request: request,
createdAt: now,
previous: d.requestTail,
next: -1,
})
}
if d.requestTail >= 0 {
d.requestSlots[d.requestTail].next = requestIndex
} else {
d.requestHead = requestIndex
}
d.requestTail = requestIndex
d.requests[request] = requestIndex
}
func (d *Destination) pruneRequests(now time.Time) {
for d.requestHead >= 0 && now.Sub(d.requestSlots[d.requestHead].createdAt) > d.timeout {
d.removeRequestAt(d.requestHead)
} }
} }
if len(d.requests) > requestsLimit {
delete(d.requests, oldestRequest) func (d *Destination) removeRequest(request pingRequest) bool {
requestIndex, loaded := d.requests[request]
if !loaded {
return false
} }
d.requests[request] = now d.removeRequestAt(requestIndex)
return true
}
func (d *Destination) removeRequestAt(requestIndex int) {
trackedRequest := &d.requestSlots[requestIndex]
if trackedRequest.previous >= 0 {
d.requestSlots[trackedRequest.previous].next = trackedRequest.next
} else {
d.requestHead = trackedRequest.next
}
if trackedRequest.next >= 0 {
d.requestSlots[trackedRequest.next].previous = trackedRequest.previous
} else {
d.requestTail = trackedRequest.previous
}
delete(d.requests, trackedRequest.request)
trackedRequest.request = pingRequest{}
trackedRequest.createdAt = time.Time{}
trackedRequest.previous = -1
trackedRequest.next = d.requestFree
d.requestFree = requestIndex
} }
func (d *Destination) Close() error { func (d *Destination) Close() error {

View file

@ -6,6 +6,7 @@ import (
"net/netip" "net/netip"
"reflect" "reflect"
"runtime" "runtime"
"sync"
"sync/atomic" "sync/atomic"
"time" "time"
@ -30,6 +31,9 @@ type Conn struct {
closed atomic.Bool closed atomic.Bool
identFilter identFilterState identFilter identFilterState
readMsg func(b, oob []byte) (n, oobn int, addr netip.Addr, err error) readMsg func(b, oob []byte) (n, oobn int, addr netip.Addr, err error)
writeAccess sync.Mutex
lastTTL int
lastHopLimit int
} }
func Connect(ctx context.Context, privileged bool, controlFunc control.Func, destination netip.Addr, idleTimeout time.Duration) (*Conn, error) { func Connect(ctx context.Context, privileged bool, controlFunc control.Func, destination netip.Addr, idleTimeout time.Duration) (*Conn, error) {
@ -37,6 +41,8 @@ func Connect(ctx context.Context, privileged bool, controlFunc control.Func, des
ctx: ctx, ctx: ctx,
privileged: privileged, privileged: privileged,
destination: destination, destination: destination,
lastTTL: -1,
lastHopLimit: -1,
} }
err := c.connect(controlFunc, idleTimeout) err := c.connect(controlFunc, idleTimeout)
if err != nil { if err != nil {
@ -259,13 +265,19 @@ func (c *Conn) ReadICMP(buffer *buf.Buffer) error {
func (c *Conn) WriteIP(buffer *buf.Buffer) error { func (c *Conn) WriteIP(buffer *buf.Buffer) error {
defer buffer.Release() defer buffer.Release()
c.writeAccess.Lock()
defer c.writeAccess.Unlock()
if !c.destination.Is6() { if !c.destination.Is6() {
ipHdr := header.IPv4(buffer.Bytes()) ipHdr := header.IPv4(buffer.Bytes())
if !c.isLinuxUnprivileged() { if !c.isLinuxUnprivileged() {
err := ipv4.NewConn(c.controlConn).SetTTL(int(ipHdr.TTL())) ttl := int(ipHdr.TTL())
if ttl != c.lastTTL {
err := ipv4.NewConn(c.controlConn).SetTTL(ttl)
if err != nil { if err != nil {
return err return err
} }
c.lastTTL = ttl
}
icmpHdr := header.ICMPv4(ipHdr.Payload()) icmpHdr := header.ICMPv4(ipHdr.Payload())
icmpHdr.SetIdent(^icmpHdr.Ident()) icmpHdr.SetIdent(^icmpHdr.Ident())
icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0))
@ -276,10 +288,14 @@ func (c *Conn) WriteIP(buffer *buf.Buffer) error {
} else { } else {
ipHdr := header.IPv6(buffer.Bytes()) ipHdr := header.IPv6(buffer.Bytes())
if !c.isLinuxUnprivileged() { if !c.isLinuxUnprivileged() {
err := ipv6.NewConn(c.controlConn).SetHopLimit(int(ipHdr.HopLimit())) hopLimit := int(ipHdr.HopLimit())
if hopLimit != c.lastHopLimit {
err := ipv6.NewConn(c.controlConn).SetHopLimit(hopLimit)
if err != nil { if err != nil {
return err return err
} }
c.lastHopLimit = hopLimit
}
icmpHdr := header.ICMPv6(ipHdr.Payload()) icmpHdr := header.ICMPv6(ipHdr.Payload())
icmpHdr.SetIdent(^icmpHdr.Ident()) icmpHdr.SetIdent(^icmpHdr.Ident())
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{ icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{