ping: Fix rewriter

This commit is contained in:
世界 2025-08-25 22:04:16 +08:00
parent ce55929883
commit 59e42c0d1f
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
2 changed files with 33 additions and 21 deletions

View file

@ -76,7 +76,7 @@ func ConnectGVisor(
return nil, gonet.TranslateNetstackError(gErr) return nil, gonet.TranslateNetstackError(gErr)
} }
endpoint.SocketOptions().SetHeaderIncluded(true) endpoint.SocketOptions().SetHeaderIncluded(true)
rewriter := NewRewriter(bindAddress4, bindAddress6) rewriter := NewRewriter(ctx, logger, bindAddress4, bindAddress6)
rewriter.CreateSession(tun.DirectRouteSession{Source: sourceAddress, Destination: destinationAddress}, routeContext) rewriter.CreateSession(tun.DirectRouteSession{Source: sourceAddress, Destination: destinationAddress}, routeContext)
destination := &GVisorDestination{ destination := &GVisorDestination{
ctx: ctx, ctx: ctx,

View file

@ -1,25 +1,31 @@
package ping package ping
import ( import (
"context"
"net/netip" "net/netip"
"sync" "sync"
"github.com/sagernet/sing-tun" "github.com/sagernet/sing-tun"
"github.com/sagernet/sing-tun/internal/gtcpip/header" "github.com/sagernet/sing-tun/internal/gtcpip/header"
"github.com/sagernet/sing/common/logger"
) )
type Rewriter struct { type Rewriter struct {
ctx context.Context
logger logger.ContextLogger
access sync.RWMutex access sync.RWMutex
sessions map[tun.DirectRouteSession]tun.DirectRouteContext sessions map[tun.DirectRouteSession]tun.DirectRouteContext
source4Address map[uint16]netip.Addr sourceAddress map[uint16]netip.Addr
source6Address map[uint16]netip.Addr
inet4Address netip.Addr inet4Address netip.Addr
inet6Address netip.Addr inet6Address netip.Addr
} }
func NewRewriter(inet4Address netip.Addr, inet6Address netip.Addr) *Rewriter { func NewRewriter(ctx context.Context, logger logger.ContextLogger, inet4Address netip.Addr, inet6Address netip.Addr) *Rewriter {
return &Rewriter{ return &Rewriter{
ctx: ctx,
logger: logger,
sessions: make(map[tun.DirectRouteSession]tun.DirectRouteContext), sessions: make(map[tun.DirectRouteSession]tun.DirectRouteContext),
sourceAddress: make(map[uint16]netip.Addr),
inet4Address: inet4Address, inet4Address: inet4Address,
inet6Address: inet6Address, inet6Address: inet6Address,
} }
@ -60,8 +66,9 @@ func (m *Rewriter) RewritePacket(packet []byte) {
case header.ICMPv4ProtocolNumber: case header.ICMPv4ProtocolNumber:
icmpHdr := header.ICMPv4(ipHdr.Payload()) icmpHdr := header.ICMPv4(ipHdr.Payload())
m.access.Lock() m.access.Lock()
m.source4Address[icmpHdr.Ident()] = sourceAddr m.sourceAddress[icmpHdr.Ident()] = sourceAddr
m.access.Lock() m.access.Unlock()
m.logger.TraceContext(m.ctx, "write ICMPv4 echo request from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence())
case header.ICMPv6ProtocolNumber: case header.ICMPv6ProtocolNumber:
icmpHdr := header.ICMPv6(ipHdr.Payload()) icmpHdr := header.ICMPv6(ipHdr.Payload())
icmpHdr.SetChecksum(0) icmpHdr.SetChecksum(0)
@ -71,8 +78,9 @@ func (m *Rewriter) RewritePacket(packet []byte) {
Dst: ipHdr.DestinationAddressSlice(), Dst: ipHdr.DestinationAddressSlice(),
})) }))
m.access.Lock() m.access.Lock()
m.source6Address[icmpHdr.Ident()] = sourceAddr m.sourceAddress[icmpHdr.Ident()] = sourceAddr
m.access.Lock() m.access.Unlock()
m.logger.TraceContext(m.ctx, "write ICMPv6 echo request from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence())
} }
} }
@ -94,25 +102,25 @@ func (m *Rewriter) WriteBack(packet []byte) (bool, error) {
icmpHdr := header.ICMPv4(ipHdr.Payload()) icmpHdr := header.ICMPv4(ipHdr.Payload())
m.access.Lock() m.access.Lock()
ident := icmpHdr.Ident() ident := icmpHdr.Ident()
source, loaded := m.source4Address[ident] source, loaded := m.sourceAddress[ident]
if !loaded { if !loaded {
m.access.Unlock() m.access.Unlock()
return false, nil return false, nil
} }
delete(m.source4Address, icmpHdr.Ident()) delete(m.sourceAddress, icmpHdr.Ident())
m.access.Lock() m.access.Unlock()
routeSession.Source = source routeSession.Source = source
case header.ICMPv6ProtocolNumber: case header.ICMPv6ProtocolNumber:
icmpHdr := header.ICMPv6(ipHdr.Payload()) icmpHdr := header.ICMPv6(ipHdr.Payload())
m.access.Lock() m.access.Lock()
ident := icmpHdr.Ident() ident := icmpHdr.Ident()
source, loaded := m.source6Address[ident] source, loaded := m.sourceAddress[ident]
if !loaded { if !loaded {
m.access.Unlock() m.access.Unlock()
return false, nil return false, nil
} }
delete(m.source6Address, icmpHdr.Ident()) delete(m.sourceAddress, icmpHdr.Ident())
m.access.Lock() m.access.Unlock()
routeSession.Source = source routeSession.Source = source
default: default:
return false, nil return false, nil
@ -129,6 +137,9 @@ func (m *Rewriter) WriteBack(packet []byte) (bool, error) {
ipHdr4.SetChecksum(^ipHdr4.CalculateChecksum()) ipHdr4.SetChecksum(^ipHdr4.CalculateChecksum())
} }
switch ipHdr.TransportProtocol() { switch ipHdr.TransportProtocol() {
case header.ICMPv4ProtocolNumber:
icmpHdr := header.ICMPv4(ipHdr.Payload())
m.logger.TraceContext(m.ctx, "read ICMPv4 echo reply from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence())
case header.ICMPv6ProtocolNumber: case header.ICMPv6ProtocolNumber:
icmpHdr := header.ICMPv6(ipHdr.Payload()) icmpHdr := header.ICMPv6(ipHdr.Payload())
icmpHdr.SetChecksum(0) icmpHdr.SetChecksum(0)
@ -137,6 +148,7 @@ func (m *Rewriter) WriteBack(packet []byte) (bool, error) {
Src: ipHdr.SourceAddressSlice(), Src: ipHdr.SourceAddressSlice(),
Dst: ipHdr.DestinationAddressSlice(), Dst: ipHdr.DestinationAddressSlice(),
})) }))
m.logger.TraceContext(m.ctx, "read ICMPv6 echo reply from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence())
} }
return true, context.WritePacket(packet) return true, context.WritePacket(packet)
} }