Add flow dispatcher

This commit is contained in:
世界 2026-07-06 11:49:18 +08:00
parent 47bdde06c3
commit ed63adda33
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
27 changed files with 2469 additions and 963 deletions

32
flow.go Normal file
View file

@ -0,0 +1,32 @@
package tun
import "net/netip"
type FlowVerdict struct {
Action FlowAction
Port Port
Destination netip.AddrPort
}
type FlowAction uint8
const (
ActionAccept FlowAction = iota
ActionFlow
ActionReject
ActionDrop
ActionBypass
)
type Port interface {
PortAddresses() (v4 netip.Addr, v6 netip.Addr)
PortMTU() uint32
AttachReturn(returnPath Return) error
DetachReturn(returnPath Return) error
WritePackets(packets [][]byte) error
}
type Return interface {
ReturnHeadroom() int
ReturnPackets(packets [][]byte) [][]byte
}

631
flow_dispatch.go Normal file
View file

@ -0,0 +1,631 @@
package tun
import (
"net/netip"
"sync/atomic"
"time"
"github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing-tun/gtcpip/header"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
)
const (
tcpEstablishedTimeout = 2*time.Hour + 4*time.Minute
tcpTransitoryTimeout = 4 * time.Minute
defaultUDPTimeout = 5 * time.Minute
defaultICMPTimeout = time.Minute
flowTombstoneTimeout = 4 * time.Minute
flowTableCapacity = 16384
flowSweepInterval = 30 * time.Second
flowSweepLimit = flowTableCapacity / int(flowTombstoneTimeout/flowSweepInterval)
)
type ForwardWriteback interface {
ReturnHeadroom() int
WriteReturnPackets(packets [][]byte) error
}
type flowEntry struct {
action FlowAction
deadline int64
idle time.Duration
flow *forwardFlow
}
type forwardFlow struct {
nat *portNAT
reverseKey flowKey
forwardRule rewriteRule
reverseRule rewriteRule
effectiveMTU uint32
protocol uint8
clientAddress netip.Addr
clientSelector uint16
clientDestinationAddress netip.Addr
clientDestinationPort uint16
serverAddress netip.Addr
dnatAddress bool
dnatPort bool
finForward bool
established atomic.Bool
finReverse atomic.Bool
closed atomic.Bool
lastReverse atomic.Int64
}
func (f *forwardFlow) observeReverse(packet *forwardPacket, now int64) {
f.lastReverse.Store(now)
if packet.protocol != uint8(header.TCPProtocolNumber) {
return
}
f.established.Store(true)
if packet.tcpFlags&header.TCPFlagRst != 0 {
f.closed.Store(true)
return
}
if packet.tcpFlags&header.TCPFlagFin != 0 {
f.finReverse.Store(true)
}
}
type ForwardDispatcher struct {
epoch time.Time
handler Handler
writeback ForwardWriteback
logger logger.Logger
udpTimeout time.Duration
icmpTimeout time.Duration
table map[flowKey]*flowEntry
lastSweep int64
ports map[Port]*portNAT
natList atomic.Pointer[[]*portNAT]
activeNATs []*portNAT
writebackBatch [][]byte
returnPath forwardReturn
segmentBuffers [][]byte
segmentSizes []int
}
func NewForwardDispatcher(handler Handler, writeback ForwardWriteback, logger logger.Logger, udpTimeout time.Duration, icmpTimeout time.Duration) *ForwardDispatcher {
dispatcher := &ForwardDispatcher{
epoch: time.Now(),
handler: handler,
writeback: writeback,
logger: logger,
udpTimeout: udpTimeout,
icmpTimeout: icmpTimeout,
table: make(map[flowKey]*flowEntry),
ports: make(map[Port]*portNAT),
}
if dispatcher.udpTimeout <= 0 {
dispatcher.udpTimeout = defaultUDPTimeout
}
if dispatcher.icmpTimeout <= 0 {
dispatcher.icmpTimeout = defaultICMPTimeout
}
dispatcher.returnPath.dispatcher = dispatcher
return dispatcher
}
func (d *ForwardDispatcher) now() int64 {
return int64(time.Since(d.epoch))
}
func (d *ForwardDispatcher) Close() {
if d == nil {
return
}
d.returnPath.closed.Store(true)
for port, nat := range d.ports {
if nat != nil {
port.DetachReturn(&d.returnPath)
}
}
}
func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
if d == nil {
return false
}
parsed, ok := parseForwardPacket(packet)
if !ok || parsed.fragment || !parsed.hasFlow {
return false
}
key := parsed.flowKey()
now := d.now()
entry, loaded := d.table[key]
if loaded && d.entryExpired(entry, now) {
d.removeEntry(key, entry)
loaded = false
}
if loaded {
return d.handleHit(key, entry, &parsed, packet, now)
}
if parsed.protocol == uint8(header.TCPProtocolNumber) &&
(parsed.tcpFlags&header.TCPFlagSyn == 0 || parsed.tcpFlags&header.TCPFlagAck != 0) {
return false
}
return d.judgeAndInstall(key, &parsed, packet, now)
}
func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *forwardPacket, raw []byte, now int64) bool {
switch entry.action {
case ActionFlow:
flow := entry.flow
if flow.closed.Load() {
d.tombstoneEntry(entry, now)
return true
}
if packet.protocol == uint8(header.TCPProtocolNumber) {
if packet.tcpFlags&header.TCPFlagRst != 0 {
d.forwardToPort(flow, packet, raw)
flow.closed.Store(true)
d.tombstoneEntry(entry, now)
return true
}
if packet.tcpFlags&header.TCPFlagFin != 0 {
flow.finForward = true
}
}
entry.idle = d.flowIdle(flow)
entry.deadline = now + int64(entry.idle)
d.forwardToPort(flow, packet, raw)
return true
case ActionAccept:
entry.deadline = now + int64(entry.idle)
if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 {
d.removeEntry(key, entry)
}
return false
case ActionReject:
entry.deadline = now + int64(entry.idle)
d.stageReject(packet)
return true
default:
entry.deadline = now + int64(entry.idle)
return true
}
}
func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool {
verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination)
switch verdict.Action {
case ActionFlow:
if verdict.Port != nil {
flow, created := d.createFlow(packet, verdict)
if created {
entry := &flowEntry{action: ActionFlow, flow: flow, idle: d.flowIdle(flow)}
entry.deadline = now + int64(entry.idle)
d.insertEntry(key, entry, now)
d.forwardToPort(flow, packet, raw)
return true
}
}
d.installSimple(key, ActionAccept, packet.protocol, now)
return false
case ActionReject:
d.installSimple(key, ActionReject, packet.protocol, now)
d.stageReject(packet)
return true
case ActionDrop:
d.installSimple(key, ActionDrop, packet.protocol, now)
return true
default:
d.installSimple(key, ActionAccept, packet.protocol, now)
return false
}
}
func (d *ForwardDispatcher) installSimple(key flowKey, action FlowAction, protocol uint8, now int64) {
entry := &flowEntry{action: action, idle: d.idleTimeout(protocol, false)}
entry.deadline = now + int64(entry.idle)
d.insertEntry(key, entry, now)
}
func (d *ForwardDispatcher) idleTimeout(protocol uint8, established bool) time.Duration {
switch protocol {
case uint8(header.TCPProtocolNumber):
if established {
return tcpEstablishedTimeout
}
return tcpTransitoryTimeout
case uint8(header.UDPProtocolNumber):
return d.udpTimeout
default:
return d.icmpTimeout
}
}
func (d *ForwardDispatcher) flowIdle(flow *forwardFlow) time.Duration {
established := flow.established.Load() && !flow.finForward && !flow.finReverse.Load()
return d.idleTimeout(flow.protocol, established)
}
func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdict) (*forwardFlow, bool) {
var portAddress netip.Addr
inet4Address, inet6Address := verdict.Port.PortAddresses()
if packet.ipVersion == 6 {
portAddress = inet6Address
} else {
portAddress = inet4Address
}
if !portAddress.IsValid() {
return nil, false
}
effectiveMTU := verdict.Port.PortMTU()
if packet.ipVersion == 6 && effectiveMTU != 0 && effectiveMTU < header.IPv6MinimumMTU {
return nil, false
}
isICMP := isICMPProtocol(packet.protocol)
clientDestinationAddress := packet.destination.Addr()
clientDestinationPort := packet.destination.Port()
serverAddress := clientDestinationAddress
serverPort := clientDestinationPort
if verdict.Destination.Addr().IsValid() {
serverAddress = verdict.Destination.Addr()
}
if verdict.Destination.Port() != 0 && !isICMP {
serverPort = verdict.Destination.Port()
}
nat := d.natFor(verdict.Port)
if nat == nil {
return nil, false
}
selector, reverseKey, allocated := nat.allocateSelector(packet.protocol, portAddress, serverAddress, serverPort, packet.source.Port())
if !allocated {
return nil, false
}
flow := &forwardFlow{
nat: nat,
reverseKey: reverseKey,
effectiveMTU: effectiveMTU,
protocol: packet.protocol,
clientAddress: packet.source.Addr(),
clientSelector: packet.source.Port(),
clientDestinationAddress: clientDestinationAddress,
clientDestinationPort: clientDestinationPort,
serverAddress: serverAddress,
dnatAddress: serverAddress != clientDestinationAddress,
dnatPort: serverPort != clientDestinationPort && !isICMP,
}
flow.forwardRule = rewriteRule{
sourceAddress: tcpip.AddrFromSlice(portAddress.AsSlice()),
sourcePort: selector,
rewriteSourcePort: true,
}
if flow.dnatAddress {
flow.forwardRule.destinationAddress = tcpip.AddrFromSlice(serverAddress.AsSlice())
}
if flow.dnatPort {
flow.forwardRule.destinationPort = serverPort
flow.forwardRule.rewriteDestinationPort = true
}
flow.reverseRule = rewriteRule{
destinationAddress: tcpip.AddrFromSlice(flow.clientAddress.AsSlice()),
destinationPort: flow.clientSelector,
rewriteDestinationPort: true,
}
if flow.dnatAddress {
flow.reverseRule.sourceAddress = tcpip.AddrFromSlice(clientDestinationAddress.AsSlice())
}
if flow.dnatPort {
flow.reverseRule.sourcePort = clientDestinationPort
flow.reverseRule.rewriteSourcePort = true
}
nat.insert(reverseKey, flow)
return flow, true
}
func (d *ForwardDispatcher) natFor(port Port) *portNAT {
nat, loaded := d.ports[port]
if loaded {
return nat
}
err := port.AttachReturn(&d.returnPath)
if err != nil {
d.logger.Trace(E.Cause(err, "attach return path"))
d.ports[port] = nil
return nil
}
nat = newPortNAT(port)
d.ports[port] = nat
var natList []*portNAT
current := d.natList.Load()
if current != nil {
natList = append(natList, *current...)
}
natList = append(natList, nat)
d.natList.Store(&natList)
return nat
}
func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPacket, raw []byte) {
if flow.effectiveMTU != 0 && uint32(len(raw)) > flow.effectiveMTU {
if packet.protocol == uint8(header.TCPProtocolNumber) {
d.rewriteForward(flow, packet)
d.resegmentTCP(flow, packet, raw)
return
}
if packet.ipVersion == 4 {
ipHdr := packet.network.(header.IPv4)
if ipHdr.Flags()&header.IPv4FlagDontFragment == 0 {
d.rewriteForward(flow, packet)
fragments, ok := fragmentIPv4Packet(ipHdr, flow.effectiveMTU)
if ok {
for _, fragment := range fragments {
d.stagePort(flow.nat, fragment)
}
}
return
}
reply, ok := buildFragmentationNeeded(ipHdr, flow.effectiveMTU, d.writeback.ReturnHeadroom())
if ok {
d.writebackBatch = append(d.writebackBatch, reply)
}
return
}
reply, ok := buildPacketTooBig(packet.network.(header.IPv6), flow.effectiveMTU, d.writeback.ReturnHeadroom())
if ok {
d.writebackBatch = append(d.writebackBatch, reply)
}
return
}
d.rewriteForward(flow, packet)
d.stagePort(flow.nat, raw)
}
func (d *ForwardDispatcher) rewriteForward(flow *forwardFlow, packet *forwardPacket) {
if packet.isTCPSyn() {
applyRewriteRaw(packet, &flow.forwardRule)
clampTCPMSS(packet, flow.effectiveMTU)
recomputeChecksums(packet)
} else {
applyRewrite(packet, &flow.forwardRule)
}
}
func (d *ForwardDispatcher) stagePort(nat *portNAT, packet []byte) {
if len(nat.pending) == 0 {
d.activeNATs = append(d.activeNATs, nat)
}
nat.pending = append(nat.pending, packet)
}
func (d *ForwardDispatcher) flushPort(nat *portNAT) {
if len(nat.pending) == 0 {
return
}
err := nat.port.WritePackets(nat.pending)
if err != nil {
d.logger.Trace(E.Cause(err, "forward packets"))
}
nat.pending = nat.pending[:0]
}
func (d *ForwardDispatcher) stageReject(packet *forwardPacket) {
reply, ok := buildReject(packet, d.writeback.ReturnHeadroom())
if ok {
d.writebackBatch = append(d.writebackBatch, reply)
}
}
func (d *ForwardDispatcher) Flush() {
if d == nil {
return
}
for _, nat := range d.activeNATs {
d.flushPort(nat)
}
d.activeNATs = d.activeNATs[:0]
if len(d.writebackBatch) > 0 {
err := d.writeback.WriteReturnPackets(d.writebackBatch)
if err != nil {
d.logger.Trace(E.Cause(err, "write back packets"))
}
d.writebackBatch = d.writebackBatch[:0]
}
d.maybeSweep(d.now())
}
func (d *ForwardDispatcher) entryExpired(entry *flowEntry, now int64) bool {
if now <= entry.deadline {
return false
}
if entry.action == ActionFlow {
lastReverse := entry.flow.lastReverse.Load()
reverseDeadline := lastReverse + int64(entry.idle)
if lastReverse != 0 && now <= reverseDeadline {
entry.deadline = reverseDeadline
return false
}
}
return true
}
func (d *ForwardDispatcher) tombstoneEntry(entry *flowEntry, now int64) {
entry.action = ActionDrop
entry.idle = flowTombstoneTimeout
entry.deadline = now + int64(entry.idle)
}
func (d *ForwardDispatcher) removeEntry(key flowKey, entry *flowEntry) {
delete(d.table, key)
if entry.flow != nil {
entry.flow.closed.Store(true)
entry.flow.nat.delete(entry.flow.reverseKey)
}
}
func (d *ForwardDispatcher) insertEntry(key flowKey, entry *flowEntry, now int64) {
if len(d.table) >= flowTableCapacity {
d.evictEntries(now)
}
d.table[key] = entry
}
func (d *ForwardDispatcher) evictEntries(now int64) {
var (
freed int
visited int
oldestKey flowKey
oldest *flowEntry
)
for key, entry := range d.table {
if d.entryExpired(entry, now) {
d.removeEntry(key, entry)
freed++
} else if oldest == nil || entry.deadline < oldest.deadline {
oldestKey = key
oldest = entry
}
visited++
if visited >= flowSweepLimit {
break
}
}
if freed == 0 && oldest != nil {
d.removeEntry(oldestKey, oldest)
}
}
func (d *ForwardDispatcher) maybeSweep(now int64) {
if now-d.lastSweep < int64(flowSweepInterval) {
return
}
d.lastSweep = now
visited := 0
for key, entry := range d.table {
if entry.action == ActionFlow && entry.flow.closed.Load() {
d.tombstoneEntry(entry, now)
} else if d.entryExpired(entry, now) {
d.removeEntry(key, entry)
}
visited++
if visited >= flowSweepLimit {
break
}
}
}
func isICMPProtocol(protocol uint8) bool {
return protocol == uint8(header.ICMPv4ProtocolNumber) || protocol == uint8(header.ICMPv6ProtocolNumber)
}
var _ Return = (*forwardReturn)(nil)
type forwardReturn struct {
dispatcher *ForwardDispatcher
closed atomic.Bool
}
func (r *forwardReturn) ReturnHeadroom() int {
return r.dispatcher.writeback.ReturnHeadroom()
}
func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
if r.closed.Load() {
return packets
}
natListPtr := r.dispatcher.natList.Load()
if natListPtr == nil {
return packets
}
natList := *natListPtr
headroom := r.dispatcher.writeback.ReturnHeadroom()
unconsumed := packets[:0]
var writeBatch [][]byte
now := r.dispatcher.now()
for _, raw := range packets {
if len(raw) < headroom+header.IPv4MinimumSize {
unconsumed = append(unconsumed, raw)
continue
}
parsed, ok := parseForwardPacket(raw[headroom:])
if !ok || parsed.fragment {
unconsumed = append(unconsumed, raw)
continue
}
if !parsed.hasFlow {
if parsed.isICMPError() && returnICMPError(natList, &parsed) {
writeBatch = append(writeBatch, raw)
} else {
unconsumed = append(unconsumed, raw)
}
continue
}
var flow *forwardFlow
for _, nat := range natList {
flow = nat.lookup(parsed.flowKey())
if flow != nil {
break
}
}
if flow == nil {
unconsumed = append(unconsumed, raw)
continue
}
if flow.closed.Load() {
continue
}
flow.observeReverse(&parsed, now)
if parsed.isTCPSyn() {
applyRewriteRaw(&parsed, &flow.reverseRule)
clampTCPMSS(&parsed, flow.effectiveMTU)
recomputeChecksums(&parsed)
} else {
applyRewrite(&parsed, &flow.reverseRule)
}
writeBatch = append(writeBatch, raw)
}
if len(writeBatch) > 0 {
err := r.dispatcher.writeback.WriteReturnPackets(writeBatch)
if err != nil {
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
}
}
return unconsumed
}
func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
inner, ok := parsed.icmpErrorInner()
if !ok {
return false
}
embedded, parsedInner := parseEmbedded(inner)
if !parsedInner {
return false
}
key := embedded.flowKey().reversed()
var flow *forwardFlow
for _, nat := range natList {
flow = nat.lookup(key)
if flow != nil {
break
}
}
if flow == nil || flow.closed.Load() {
return false
}
rewriteEmbeddedSource(&embedded, tcpip.AddrFromSlice(flow.clientAddress.AsSlice()), flow.clientSelector, true)
if flow.dnatAddress || flow.dnatPort {
rewriteEmbeddedDestination(&embedded, tcpip.AddrFromSlice(flow.clientDestinationAddress.AsSlice()), flow.clientDestinationPort, flow.dnatPort)
}
parsed.network.SetDestinationAddr(flow.clientAddress)
if parsed.network.SourceAddr() == flow.serverAddress {
parsed.network.SetSourceAddr(flow.clientDestinationAddress)
}
recomputeChecksums(parsed)
return true
}

159
flow_mtu.go Normal file
View file

@ -0,0 +1,159 @@
package tun
import (
"github.com/sagernet/sing-tun/gtcpip/header"
E "github.com/sagernet/sing/common/exceptions"
)
const segmentScratchCount = 128
// Linux delivers TSO aggregates to the TUN even with IFF_VNET_HDR off
// (observed on 6.x: the pre-segmentation skb is handed to the fd as-is).
func (d *ForwardDispatcher) resegmentTCP(flow *forwardFlow, packet *forwardPacket, raw []byte) {
if len(packet.transport) < header.TCPMinimumSize {
return
}
headerLength := len(raw) - len(packet.transport)
if packet.ipVersion == 6 && headerLength != header.IPv6MinimumSize {
reply, ok := buildPacketTooBig(packet.network.(header.IPv6), flow.effectiveMTU, d.writeback.ReturnHeadroom())
if ok {
d.writebackBatch = append(d.writebackBatch, reply)
}
return
}
tcpHeaderLength := int(header.TCP(packet.transport).DataOffset())
if tcpHeaderLength < header.TCPMinimumSize || tcpHeaderLength > len(packet.transport) {
return
}
totalHeaderLength := headerLength + tcpHeaderLength
segmentSize := int(flow.effectiveMTU) - totalHeaderLength
if segmentSize <= 0 {
return
}
gsoType := GSOTCPv4
if packet.ipVersion == 6 {
gsoType = GSOTCPv6
}
neededSegments := max((len(raw)-totalHeaderLength+segmentSize-1)/segmentSize, 1)
if d.segmentBuffers == nil || len(d.segmentBuffers) < neededSegments || len(d.segmentBuffers[0]) < int(flow.effectiveMTU) {
bufferSize := int(flow.effectiveMTU)
if d.segmentBuffers != nil && len(d.segmentBuffers[0]) > bufferSize {
bufferSize = len(d.segmentBuffers[0])
}
segmentCount := max(neededSegments, segmentScratchCount, len(d.segmentBuffers))
d.segmentBuffers = make([][]byte, segmentCount)
for i := range d.segmentBuffers {
d.segmentBuffers[i] = make([]byte, bufferSize)
}
d.segmentSizes = make([]int, segmentCount)
}
n, err := GSOSplit(raw, GSOOptions{
GSOType: gsoType,
HdrLen: uint16(totalHeaderLength),
CsumStart: uint16(headerLength),
CsumOffset: header.TCPChecksumOffset,
GSOSize: uint16(segmentSize),
}, d.segmentBuffers, d.segmentSizes, 0)
if err != nil {
d.logger.Trace(E.Cause(err, "resegment packet"))
return
}
for i := range n {
d.stagePort(flow.nat, d.segmentBuffers[i][:d.segmentSizes[i]])
}
d.flushPort(flow.nat)
}
const synthesizedTTL = 64
func fragmentIPv4Packet(packet header.IPv4, effectiveMTU uint32) ([][]byte, bool) {
headerLength := int(packet.HeaderLength())
if headerLength < header.IPv4MinimumSize || headerLength >= len(packet) {
return nil, false
}
payload := packet[headerLength:]
maxFragmentPayload := (int(effectiveMTU) - headerLength) &^ 7
if maxFragmentPayload <= 0 {
return nil, false
}
baseOffset := packet.FragmentOffset()
originalMore := packet.Flags()&header.IPv4FlagMoreFragments != 0
baseFlags := packet.Flags() &^ header.IPv4FlagMoreFragments
var fragments [][]byte
for start := 0; start < len(payload); start += maxFragmentPayload {
end := min(start+maxFragmentPayload, len(payload))
fragment := header.IPv4(make([]byte, headerLength+end-start))
copy(fragment, packet[:headerLength])
copy(fragment[headerLength:], payload[start:end])
flags := baseFlags
if originalMore || end < len(payload) {
flags |= header.IPv4FlagMoreFragments
}
fragment.SetFlagsFragmentOffset(flags, baseOffset+uint16(start))
fragment.SetTotalLength(uint16(len(fragment)))
fragment.SetChecksum(0)
fragment.SetChecksum(^fragment.CalculateChecksum())
fragments = append(fragments, fragment)
}
return fragments, true
}
func buildFragmentationNeeded(packet header.IPv4, effectiveMTU uint32, headroom int) ([]byte, bool) {
advertised := max(effectiveMTU, header.IPv4MinimumMTU)
originalLength := min(int(packet.TotalLength()), len(packet))
minPayloadLength := int(packet.HeaderLength()) + header.ICMPv4MinimumErrorPayloadSize
if originalLength < minPayloadLength {
return nil, false
}
maxPayloadLength := header.IPv4MinimumProcessableDatagramSize - header.IPv4MinimumSize - header.ICMPv4MinimumSize
payloadLength := min(originalLength, maxPayloadLength)
size := header.IPv4MinimumSize + header.ICMPv4MinimumSize + payloadLength
buffer := make([]byte, headroom+size)
response := header.IPv4(buffer[headroom:])
response.Encode(&header.IPv4Fields{
TotalLength: uint16(size),
TTL: synthesizedTTL,
Protocol: uint8(header.ICMPv4ProtocolNumber),
SrcAddr: packet.DestinationAddr(),
DstAddr: packet.SourceAddr(),
})
response.SetChecksum(^response.CalculateChecksum())
icmpHdr := header.ICMPv4(response.Payload())
icmpHdr.SetType(header.ICMPv4DstUnreachable)
icmpHdr.SetCode(header.ICMPv4FragmentationNeeded)
icmpHdr.SetMTU(uint16(min(advertised, uint32(0xffff))))
copy(icmpHdr.Payload(), packet[:payloadLength])
icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0))
return buffer, true
}
func buildPacketTooBig(packet header.IPv6, effectiveMTU uint32, headroom int) ([]byte, bool) {
advertised := max(effectiveMTU, header.IPv6MinimumMTU)
originalLength := min(header.IPv6MinimumSize+int(packet.PayloadLength()), len(packet))
if originalLength < header.IPv6MinimumSize {
return nil, false
}
maxPayloadLength := header.IPv6MinimumMTU - header.IPv6MinimumSize - header.ICMPv6PacketTooBigMinimumSize
payloadLength := min(originalLength, maxPayloadLength)
size := header.IPv6MinimumSize + header.ICMPv6PacketTooBigMinimumSize + payloadLength
buffer := make([]byte, headroom+size)
response := header.IPv6(buffer[headroom:])
response.Encode(&header.IPv6Fields{
PayloadLength: uint16(header.ICMPv6PacketTooBigMinimumSize + payloadLength),
TransportProtocol: header.ICMPv6ProtocolNumber,
HopLimit: synthesizedTTL,
SrcAddr: packet.DestinationAddr(),
DstAddr: packet.SourceAddr(),
})
icmpHdr := header.ICMPv6(response.Payload())
icmpHdr.SetType(header.ICMPv6PacketTooBig)
icmpHdr.SetCode(header.ICMPv6UnusedCode)
icmpHdr.SetMTU(advertised)
copy(icmpHdr.Payload(), packet[:payloadLength])
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr,
Src: response.SourceAddressSlice(),
Dst: response.DestinationAddressSlice(),
}))
return buffer, true
}

106
flow_nat.go Normal file
View file

@ -0,0 +1,106 @@
package tun
import (
"net/netip"
"runtime"
"sync"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/contrab/maphash"
)
const (
natSelectorMin = 49152
natSelectorMax = 65535
)
type portNAT struct {
port Port
hasher maphash.Hasher[flowKey]
shardMask uint32
shards []natShard
counter uint32
pending [][]byte
}
type natShard struct {
access sync.RWMutex
flows map[flowKey]*forwardFlow
}
func newPortNAT(port Port) *portNAT {
shardCount := 1
for shardCount < runtime.GOMAXPROCS(0) {
shardCount <<= 1
}
nat := &portNAT{
port: port,
hasher: maphash.NewHasher[flowKey](),
shardMask: uint32(shardCount - 1),
shards: make([]natShard, shardCount),
}
for i := range nat.shards {
nat.shards[i].flows = make(map[flowKey]*forwardFlow)
}
return nat
}
func (n *portNAT) shard(key flowKey) *natShard {
return &n.shards[n.hasher.Hash32(key)&n.shardMask]
}
func (n *portNAT) lookup(key flowKey) *forwardFlow {
shard := n.shard(key)
shard.access.RLock()
flow := shard.flows[key]
shard.access.RUnlock()
return flow
}
func (n *portNAT) insert(key flowKey, flow *forwardFlow) {
shard := n.shard(key)
shard.access.Lock()
shard.flows[key] = flow
shard.access.Unlock()
}
func (n *portNAT) delete(key flowKey) {
shard := n.shard(key)
shard.access.Lock()
delete(shard.flows, key)
shard.access.Unlock()
}
func (n *portNAT) reverseKeyFor(protocol uint8, portAddress, serverAddress netip.Addr, serverPort, selector uint16) flowKey {
if protocol == uint8(header.ICMPv4ProtocolNumber) || protocol == uint8(header.ICMPv6ProtocolNumber) {
return flowKey{
protocol: protocol,
source: netip.AddrPortFrom(serverAddress, selector),
destination: netip.AddrPortFrom(portAddress, selector),
}
}
return flowKey{
protocol: protocol,
source: netip.AddrPortFrom(serverAddress, serverPort),
destination: netip.AddrPortFrom(portAddress, selector),
}
}
func (n *portNAT) allocateSelector(protocol uint8, portAddress, serverAddress netip.Addr, serverPort, clientSelector uint16) (uint16, flowKey, bool) {
if clientSelector != 0 {
key := n.reverseKeyFor(protocol, portAddress, serverAddress, serverPort, clientSelector)
if n.lookup(key) == nil {
return clientSelector, key, true
}
}
for range natSelectorMax - natSelectorMin + 1 {
n.counter++
candidate := uint16(natSelectorMin + n.counter%(natSelectorMax-natSelectorMin+1))
key := n.reverseKeyFor(protocol, portAddress, serverAddress, serverPort, candidate)
if n.lookup(key) == nil {
return candidate, key, true
}
}
return 0, flowKey{}, false
}

256
flow_parse.go Normal file
View file

@ -0,0 +1,256 @@
package tun
import (
"encoding/binary"
"net/netip"
"github.com/sagernet/sing-tun/gtcpip/header"
)
type flowKey struct {
protocol uint8
source netip.AddrPort
destination netip.AddrPort
}
func (k flowKey) reversed() flowKey {
return flowKey{protocol: k.protocol, source: k.destination, destination: k.source}
}
type forwardPacket struct {
ipVersion uint8
protocol uint8
network header.Network
transport []byte
source netip.AddrPort
destination netip.AddrPort
tcpFlags header.TCPFlags
icmpType uint8
fragment bool
hasFlow bool
}
func (p *forwardPacket) flowKey() flowKey {
return flowKey{protocol: p.protocol, source: p.source, destination: p.destination}
}
func (p *forwardPacket) isTCPSyn() bool {
return p.protocol == uint8(header.TCPProtocolNumber) && p.tcpFlags&header.TCPFlagSyn != 0
}
func parseForwardPacket(packet []byte) (forwardPacket, bool) {
switch header.IPVersion(packet) {
case header.IPv4Version:
ipHdr := header.IPv4(packet)
if !ipHdr.IsValid(len(packet)) {
return forwardPacket{}, false
}
parsed := forwardPacket{
ipVersion: 4,
protocol: uint8(ipHdr.TransportProtocol()),
network: ipHdr,
source: netip.AddrPortFrom(ipHdr.SourceAddr(), 0),
destination: netip.AddrPortFrom(ipHdr.DestinationAddr(), 0),
}
if ipHdr.More() || ipHdr.FragmentOffset() != 0 {
parsed.fragment = true
return parsed, true
}
parsed.parseTransport(ipHdr.Payload())
return parsed, true
case header.IPv6Version:
ipHdr := header.IPv6(packet)
if !ipHdr.IsValid(len(packet)) {
return forwardPacket{}, false
}
protocol, payload, fragment, transportPresent := skipIPv6ExtensionHeaders(uint8(ipHdr.TransportProtocol()), ipHdr.Payload())
parsed := forwardPacket{
ipVersion: 6,
protocol: protocol,
network: ipHdr,
source: netip.AddrPortFrom(ipHdr.SourceAddr(), 0),
destination: netip.AddrPortFrom(ipHdr.DestinationAddr(), 0),
fragment: fragment,
}
if fragment || !transportPresent {
return parsed, true
}
parsed.parseTransport(payload)
return parsed, true
default:
return forwardPacket{}, false
}
}
func skipIPv6ExtensionHeaders(protocol uint8, payload []byte) (uint8, []byte, bool, bool) {
for {
switch header.IPv6ExtensionHeaderIdentifier(protocol) {
case header.IPv6HopByHopOptionsExtHdrIdentifier, header.IPv6RoutingExtHdrIdentifier, header.IPv6DestinationOptionsExtHdrIdentifier:
if len(payload) < 2 {
return protocol, payload, false, false
}
extensionLength := (int(payload[1]) + 1) * 8
if len(payload) < extensionLength {
return protocol, payload, false, false
}
protocol = payload[0]
payload = payload[extensionLength:]
case header.IPv6FragmentExtHdrIdentifier:
return protocol, payload, true, false
default:
return protocol, payload, false, true
}
}
}
func (p *forwardPacket) parseTransport(payload []byte) {
p.transport = payload
switch p.protocol {
case uint8(header.TCPProtocolNumber):
if len(payload) < header.TCPMinimumSize {
return
}
tcpHdr := header.TCP(payload)
p.source = netip.AddrPortFrom(p.source.Addr(), tcpHdr.SourcePort())
p.destination = netip.AddrPortFrom(p.destination.Addr(), tcpHdr.DestinationPort())
p.tcpFlags = tcpHdr.Flags()
p.hasFlow = true
case uint8(header.UDPProtocolNumber):
if len(payload) < header.UDPMinimumSize {
return
}
udpHdr := header.UDP(payload)
p.source = netip.AddrPortFrom(p.source.Addr(), udpHdr.SourcePort())
p.destination = netip.AddrPortFrom(p.destination.Addr(), udpHdr.DestinationPort())
p.hasFlow = true
case uint8(header.ICMPv4ProtocolNumber):
if len(payload) < header.ICMPv4MinimumSize {
return
}
icmpHdr := header.ICMPv4(payload)
p.icmpType = uint8(icmpHdr.Type())
switch icmpHdr.Type() {
case header.ICMPv4Echo, header.ICMPv4EchoReply:
identifier := icmpHdr.Ident()
p.source = netip.AddrPortFrom(p.source.Addr(), identifier)
p.destination = netip.AddrPortFrom(p.destination.Addr(), identifier)
p.hasFlow = true
}
case uint8(header.ICMPv6ProtocolNumber):
if len(payload) < header.ICMPv6MinimumSize {
return
}
icmpHdr := header.ICMPv6(payload)
p.icmpType = uint8(icmpHdr.Type())
switch icmpHdr.Type() {
case header.ICMPv6EchoRequest, header.ICMPv6EchoReply:
identifier := icmpHdr.Ident()
p.source = netip.AddrPortFrom(p.source.Addr(), identifier)
p.destination = netip.AddrPortFrom(p.destination.Addr(), identifier)
p.hasFlow = true
}
}
}
func (p *forwardPacket) isICMPError() bool {
switch p.protocol {
case uint8(header.ICMPv4ProtocolNumber):
switch header.ICMPv4Type(p.icmpType) {
case header.ICMPv4DstUnreachable, header.ICMPv4SrcQuench, header.ICMPv4Redirect, header.ICMPv4TimeExceeded, header.ICMPv4ParamProblem:
return true
}
return false
case uint8(header.ICMPv6ProtocolNumber):
return header.ICMPv6Type(p.icmpType).IsErrorType()
default:
return false
}
}
func (p *forwardPacket) icmpErrorInner() ([]byte, bool) {
var innerOffset int
switch p.protocol {
case uint8(header.ICMPv4ProtocolNumber):
innerOffset = header.ICMPv4MinimumSize
case uint8(header.ICMPv6ProtocolNumber):
innerOffset = header.ICMPv6ErrorHeaderSize
default:
return nil, false
}
if len(p.transport) <= innerOffset {
return nil, false
}
return p.transport[innerOffset:], true
}
type embeddedPacket struct {
network header.Network
payload []byte
protocol uint8
source netip.AddrPort
destination netip.AddrPort
}
func (p *embeddedPacket) flowKey() flowKey {
return flowKey{protocol: p.protocol, source: p.source, destination: p.destination}
}
func parseEmbedded(inner []byte) (embeddedPacket, bool) {
switch header.IPVersion(inner) {
case header.IPv4Version:
if len(inner) < header.IPv4MinimumSize {
return embeddedPacket{}, false
}
ipHdr := header.IPv4(inner)
headerLength := int(ipHdr.HeaderLength())
if headerLength < header.IPv4MinimumSize || headerLength > len(inner) {
return embeddedPacket{}, false
}
return parseEmbeddedTransport(ipHdr, inner[headerLength:], uint8(ipHdr.TransportProtocol()), ipHdr.SourceAddr(), ipHdr.DestinationAddr())
case header.IPv6Version:
if len(inner) < header.IPv6MinimumSize {
return embeddedPacket{}, false
}
ipHdr := header.IPv6(inner)
protocol, payload, _, transportPresent := skipIPv6ExtensionHeaders(uint8(ipHdr.TransportProtocol()), inner[header.IPv6MinimumSize:])
if !transportPresent {
return embeddedPacket{}, false
}
return parseEmbeddedTransport(ipHdr, payload, protocol, ipHdr.SourceAddr(), ipHdr.DestinationAddr())
default:
return embeddedPacket{}, false
}
}
func parseEmbeddedTransport(network header.Network, payload []byte, protocol uint8, source, destination netip.Addr) (embeddedPacket, bool) {
embedded := embeddedPacket{
network: network,
payload: payload,
protocol: protocol,
}
switch protocol {
case uint8(header.TCPProtocolNumber), uint8(header.UDPProtocolNumber):
if len(payload) < 4 {
return embeddedPacket{}, false
}
embedded.source = netip.AddrPortFrom(source, binary.BigEndian.Uint16(payload[0:]))
embedded.destination = netip.AddrPortFrom(destination, binary.BigEndian.Uint16(payload[2:]))
case uint8(header.ICMPv4ProtocolNumber):
if len(payload) < header.ICMPv4MinimumSize {
return embeddedPacket{}, false
}
identifier := header.ICMPv4(payload).Ident()
embedded.source = netip.AddrPortFrom(source, identifier)
embedded.destination = netip.AddrPortFrom(destination, identifier)
case uint8(header.ICMPv6ProtocolNumber):
if len(payload) < header.ICMPv6MinimumSize {
return embeddedPacket{}, false
}
identifier := header.ICMPv6(payload).Ident()
embedded.source = netip.AddrPortFrom(source, identifier)
embedded.destination = netip.AddrPortFrom(destination, identifier)
default:
return embeddedPacket{}, false
}
return embedded, true
}

210
flow_reject.go Normal file
View file

@ -0,0 +1,210 @@
package tun
import (
"net/netip"
"github.com/sagernet/sing-tun/gtcpip/checksum"
"github.com/sagernet/sing-tun/gtcpip/header"
)
func buildReject(packet *forwardPacket, headroom int) ([]byte, bool) {
switch packet.protocol {
case uint8(header.TCPProtocolNumber):
if len(packet.transport) < header.TCPMinimumSize {
return nil, false
}
tcpHdr := header.TCP(packet.transport)
switch ipHdr := packet.network.(type) {
case header.IPv4:
return buildResetIPv4(ipHdr, tcpHdr, headroom), true
case header.IPv6:
return buildResetIPv6(ipHdr, tcpHdr, headroom), true
default:
return nil, false
}
case uint8(header.UDPProtocolNumber):
switch ipHdr := packet.network.(type) {
case header.IPv4:
return buildRejectICMPv4(ipHdr, header.ICMPv4PortUnreachable, ipHdr.DestinationAddr(), headroom)
case header.IPv6:
return buildRejectICMPv6(ipHdr, header.ICMPv6PortUnreachable, ipHdr.DestinationAddr(), headroom)
default:
return nil, false
}
default:
switch ipHdr := packet.network.(type) {
case header.IPv4:
return buildRejectICMPv4(ipHdr, header.ICMPv4HostUnreachable, ipHdr.DestinationAddr(), headroom)
case header.IPv6:
return buildRejectICMPv6(ipHdr, header.ICMPv6AddressUnreachable, ipHdr.DestinationAddr(), headroom)
default:
return nil, false
}
}
}
func buildResetIPv4(origIPHdr header.IPv4, origTCPHdr header.TCP, headroom int) []byte {
size := header.IPv4MinimumSize + header.TCPMinimumSize
buffer := make([]byte, headroom+size)
ipHdr := header.IPv4(buffer[headroom:])
ipHdr.Encode(&header.IPv4Fields{
TotalLength: uint16(size),
TTL: synthesizedTTL,
Protocol: uint8(header.TCPProtocolNumber),
SrcAddr: origIPHdr.DestinationAddr(),
DstAddr: origIPHdr.SourceAddr(),
})
tcpHdr := header.TCP(ipHdr.Payload())
encodeResetTCP(tcpHdr, origTCPHdr)
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
return buffer
}
func buildResetIPv6(origIPHdr header.IPv6, origTCPHdr header.TCP, headroom int) []byte {
size := header.IPv6MinimumSize + header.TCPMinimumSize
buffer := make([]byte, headroom+size)
ipHdr := header.IPv6(buffer[headroom:])
ipHdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(header.TCPMinimumSize),
TransportProtocol: header.TCPProtocolNumber,
HopLimit: synthesizedTTL,
SrcAddr: origIPHdr.DestinationAddr(),
DstAddr: origIPHdr.SourceAddr(),
})
tcpHdr := header.TCP(ipHdr.Payload())
encodeResetTCP(tcpHdr, origTCPHdr)
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
return buffer
}
func encodeResetTCP(tcpHdr header.TCP, origTCPHdr header.TCP) {
fields := header.TCPFields{
SrcPort: origTCPHdr.DestinationPort(),
DstPort: origTCPHdr.SourcePort(),
DataOffset: header.TCPMinimumSize,
Flags: header.TCPFlagRst,
}
if origTCPHdr.Flags()&header.TCPFlagAck != 0 {
fields.SeqNum = origTCPHdr.AckNumber()
} else {
fields.Flags |= header.TCPFlagAck
ackNumber := origTCPHdr.SequenceNumber() + uint32(len(origTCPHdr.Payload()))
if origTCPHdr.Flags()&header.TCPFlagSyn != 0 {
ackNumber++
}
if origTCPHdr.Flags()&header.TCPFlagFin != 0 {
ackNumber++
}
fields.AckNum = ackNumber
}
tcpHdr.Encode(&fields)
}
func buildRejectICMPv4(ipHdr header.IPv4, code header.ICMPv4Code, source netip.Addr, headroom int) ([]byte, bool) {
const maxIPData = header.IPv4MinimumProcessableDatagramSize - header.IPv4MinimumSize
available := maxIPData - header.ICMPv4MinimumSize
if len(ipHdr) < header.ICMPv4MinimumErrorPayloadSize {
return nil, false
}
payload := []byte(ipHdr)
if len(payload) > available {
payload = payload[:available]
}
size := header.IPv4MinimumSize + header.ICMPv4MinimumSize + len(payload)
buffer := make([]byte, headroom+size)
newIPHdr := header.IPv4(buffer[headroom:])
newIPHdr.Encode(&header.IPv4Fields{
TotalLength: uint16(size),
TTL: synthesizedTTL,
Protocol: uint8(header.ICMPv4ProtocolNumber),
SrcAddr: source,
DstAddr: ipHdr.SourceAddr(),
})
newIPHdr.SetChecksum(^newIPHdr.CalculateChecksum())
icmpHdr := header.ICMPv4(newIPHdr.Payload())
icmpHdr.SetType(header.ICMPv4DstUnreachable)
icmpHdr.SetCode(code)
copy(icmpHdr.Payload(), payload)
icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr[:header.ICMPv4MinimumSize], checksum.Checksum(payload, 0)))
return buffer, true
}
func buildRejectICMPv6(ipHdr header.IPv6, code header.ICMPv6Code, source netip.Addr, headroom int) ([]byte, bool) {
const maxIPv6Data = header.IPv6MinimumMTU - header.IPv6FixedHeaderSize
available := maxIPv6Data - header.ICMPv6ErrorHeaderSize
if available < header.IPv6MinimumSize {
return nil, false
}
payload := []byte(ipHdr)
if len(payload) > available {
payload = payload[:available]
}
size := header.IPv6MinimumSize + header.ICMPv6DstUnreachableMinimumSize + len(payload)
buffer := make([]byte, headroom+size)
newIPHdr := header.IPv6(buffer[headroom:])
newIPHdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(header.ICMPv6DstUnreachableMinimumSize + len(payload)),
TransportProtocol: header.ICMPv6ProtocolNumber,
HopLimit: synthesizedTTL,
SrcAddr: source,
DstAddr: ipHdr.SourceAddr(),
})
icmpHdr := header.ICMPv6(newIPHdr.Payload())
icmpHdr.SetType(header.ICMPv6DstUnreachable)
icmpHdr.SetCode(code)
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr[:header.ICMPv6DstUnreachableMinimumSize],
Src: newIPHdr.SourceAddressSlice(),
Dst: newIPHdr.DestinationAddressSlice(),
PayloadCsum: checksum.Checksum(payload, 0),
PayloadLen: len(payload),
}))
copy(icmpHdr.Payload(), payload)
return buffer, true
}
func BuildUnreachable(packet []byte, source netip.Addr, headroom int) ([]byte, bool) {
switch header.IPVersion(packet) {
case header.IPv4Version:
ipHdr := header.IPv4(packet)
if !ipHdr.IsValid(len(packet)) || ipHdr.FragmentOffset() != 0 {
return nil, false
}
sourceAddr := ipHdr.SourceAddr()
if sourceAddr.IsUnspecified() || sourceAddr.IsMulticast() {
return nil, false
}
if ipHdr.TransportProtocol() == header.ICMPv4ProtocolNumber {
if len(ipHdr.Payload()) < header.ICMPv4MinimumSize || header.ICMPv4(ipHdr.Payload()).Type() != header.ICMPv4Echo {
return nil, false
}
}
replySource := ipHdr.DestinationAddr()
if source.Is4() {
replySource = source
}
return buildRejectICMPv4(ipHdr, header.ICMPv4HostUnreachable, replySource, headroom)
case header.IPv6Version:
ipHdr := header.IPv6(packet)
if !ipHdr.IsValid(len(packet)) {
return nil, false
}
sourceAddr := ipHdr.SourceAddr()
if sourceAddr.IsUnspecified() || sourceAddr.IsMulticast() {
return nil, false
}
if ipHdr.TransportProtocol() == header.ICMPv6ProtocolNumber {
if len(ipHdr.Payload()) < header.ICMPv6MinimumSize || header.ICMPv6(ipHdr.Payload()).Type() != header.ICMPv6EchoRequest {
return nil, false
}
}
replySource := ipHdr.DestinationAddr()
if source.Is6() {
replySource = source
}
return buildRejectICMPv6(ipHdr, header.ICMPv6NetworkUnreachable, replySource, headroom)
default:
return nil, false
}
}

350
flow_rewrite.go Normal file
View file

@ -0,0 +1,350 @@
package tun
import (
"encoding/binary"
"github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing-tun/gtcpip/checksum"
"github.com/sagernet/sing-tun/gtcpip/header"
)
type rewriteRule struct {
sourceAddress tcpip.Address
sourcePort uint16
rewriteSourcePort bool
destinationAddress tcpip.Address
destinationPort uint16
rewriteDestinationPort bool
}
func applyRewrite(packet *forwardPacket, rule *rewriteRule) {
oldSource := packet.network.SourceAddress()
oldDestination := packet.network.DestinationAddress()
newSource := oldSource
newDestination := oldDestination
if rule.sourceAddress.Len() > 0 {
newSource = rule.sourceAddress
}
if rule.destinationAddress.Len() > 0 {
newDestination = rule.destinationAddress
}
if ipHdr, isIPv4 := packet.network.(header.IPv4); isIPv4 {
if newSource != oldSource {
ipHdr.SetSourceAddressWithChecksumUpdate(newSource)
}
if newDestination != oldDestination {
ipHdr.SetDestinationAddressWithChecksumUpdate(newDestination)
}
} else {
if newSource != oldSource {
packet.network.SetSourceAddress(newSource)
}
if newDestination != oldDestination {
packet.network.SetDestinationAddress(newDestination)
}
}
transport := packet.transport
switch packet.protocol {
case uint8(header.TCPProtocolNumber):
if len(transport) < header.TCPMinimumSize {
return
}
tcpHdr := header.TCP(transport)
if newSource != oldSource {
tcpHdr.UpdateChecksumPseudoHeaderAddress(oldSource, newSource, true)
}
if newDestination != oldDestination {
tcpHdr.UpdateChecksumPseudoHeaderAddress(oldDestination, newDestination, true)
}
if rule.rewriteSourcePort {
tcpHdr.SetSourcePortWithChecksumUpdate(rule.sourcePort)
}
if rule.rewriteDestinationPort {
tcpHdr.SetDestinationPortWithChecksumUpdate(rule.destinationPort)
}
case uint8(header.UDPProtocolNumber):
if len(transport) < header.UDPMinimumSize {
return
}
udpHdr := header.UDP(transport)
if packet.ipVersion == 4 && udpHdr.Checksum() == 0 {
if rule.rewriteSourcePort {
udpHdr.SetSourcePort(rule.sourcePort)
}
if rule.rewriteDestinationPort {
udpHdr.SetDestinationPort(rule.destinationPort)
}
return
}
if newSource != oldSource {
udpHdr.UpdateChecksumPseudoHeaderAddress(oldSource, newSource, true)
}
if newDestination != oldDestination {
udpHdr.UpdateChecksumPseudoHeaderAddress(oldDestination, newDestination, true)
}
if rule.rewriteSourcePort {
udpHdr.SetSourcePortWithChecksumUpdate(rule.sourcePort)
}
if rule.rewriteDestinationPort {
udpHdr.SetDestinationPortWithChecksumUpdate(rule.destinationPort)
}
case uint8(header.ICMPv4ProtocolNumber):
if len(transport) < header.ICMPv4MinimumSize {
return
}
icmpHdr := header.ICMPv4(transport)
if rule.rewriteSourcePort {
icmpHdr.SetIdentWithChecksumUpdate(rule.sourcePort)
} else if rule.rewriteDestinationPort {
icmpHdr.SetIdentWithChecksumUpdate(rule.destinationPort)
}
case uint8(header.ICMPv6ProtocolNumber):
if len(transport) < header.ICMPv6MinimumSize {
return
}
icmpHdr := header.ICMPv6(transport)
if newSource != oldSource {
icmpHdr.UpdateChecksumPseudoHeaderAddress(oldSource, newSource)
}
if newDestination != oldDestination {
icmpHdr.UpdateChecksumPseudoHeaderAddress(oldDestination, newDestination)
}
if rule.rewriteSourcePort {
icmpHdr.SetIdentWithChecksumUpdate(rule.sourcePort)
} else if rule.rewriteDestinationPort {
icmpHdr.SetIdentWithChecksumUpdate(rule.destinationPort)
}
}
}
func applyRewriteRaw(packet *forwardPacket, rule *rewriteRule) {
if rule.sourceAddress.Len() > 0 {
packet.network.SetSourceAddress(rule.sourceAddress)
}
if rule.destinationAddress.Len() > 0 {
packet.network.SetDestinationAddress(rule.destinationAddress)
}
transport := packet.transport
switch packet.protocol {
case uint8(header.TCPProtocolNumber):
if len(transport) < header.TCPMinimumSize {
return
}
tcpHdr := header.TCP(transport)
if rule.rewriteSourcePort {
tcpHdr.SetSourcePort(rule.sourcePort)
}
if rule.rewriteDestinationPort {
tcpHdr.SetDestinationPort(rule.destinationPort)
}
case uint8(header.UDPProtocolNumber):
if len(transport) < header.UDPMinimumSize {
return
}
udpHdr := header.UDP(transport)
if rule.rewriteSourcePort {
udpHdr.SetSourcePort(rule.sourcePort)
}
if rule.rewriteDestinationPort {
udpHdr.SetDestinationPort(rule.destinationPort)
}
case uint8(header.ICMPv4ProtocolNumber):
if len(transport) < header.ICMPv4MinimumSize {
return
}
icmpHdr := header.ICMPv4(transport)
if rule.rewriteSourcePort {
icmpHdr.SetIdent(rule.sourcePort)
} else if rule.rewriteDestinationPort {
icmpHdr.SetIdent(rule.destinationPort)
}
case uint8(header.ICMPv6ProtocolNumber):
if len(transport) < header.ICMPv6MinimumSize {
return
}
icmpHdr := header.ICMPv6(transport)
if rule.rewriteSourcePort {
icmpHdr.SetIdent(rule.sourcePort)
} else if rule.rewriteDestinationPort {
icmpHdr.SetIdent(rule.destinationPort)
}
}
}
func recomputeChecksums(packet *forwardPacket) {
if ipHdr, isIPv4 := packet.network.(header.IPv4); isIPv4 {
ipHdr.SetChecksum(0)
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
}
transport := packet.transport
switch packet.protocol {
case uint8(header.TCPProtocolNumber):
if len(transport) < header.TCPMinimumSize {
return
}
tcpHdr := header.TCP(transport)
tcpHdr.SetChecksum(0)
payloadChecksum := checksum.Checksum(tcpHdr.Payload(), 0)
pseudoChecksum := header.PseudoHeaderChecksum(header.TCPProtocolNumber, packet.network.SourceAddressSlice(), packet.network.DestinationAddressSlice(), uint16(len(transport)))
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(checksum.Combine(pseudoChecksum, payloadChecksum)))
case uint8(header.UDPProtocolNumber):
if len(transport) < header.UDPMinimumSize {
return
}
udpHdr := header.UDP(transport)
if packet.ipVersion == 4 && udpHdr.Checksum() == 0 {
return
}
udpHdr.SetChecksum(0)
payloadChecksum := checksum.Checksum(udpHdr.Payload(), 0)
pseudoChecksum := header.PseudoHeaderChecksum(header.UDPProtocolNumber, packet.network.SourceAddressSlice(), packet.network.DestinationAddressSlice(), udpHdr.Length())
udpChecksum := ^udpHdr.CalculateChecksum(checksum.Combine(pseudoChecksum, payloadChecksum))
if udpChecksum == 0 {
udpChecksum = 0xffff
}
udpHdr.SetChecksum(udpChecksum)
case uint8(header.ICMPv4ProtocolNumber):
if len(transport) < header.ICMPv4MinimumSize {
return
}
icmpHdr := header.ICMPv4(transport)
icmpHdr.SetChecksum(0)
icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0))
case uint8(header.ICMPv6ProtocolNumber):
if len(transport) < header.ICMPv6MinimumSize {
return
}
icmpHdr := header.ICMPv6(transport)
icmpHdr.SetChecksum(0)
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr,
Src: packet.network.SourceAddressSlice(),
Dst: packet.network.DestinationAddressSlice(),
}))
}
}
func clampTCPMSS(packet *forwardPacket, effectiveMTU uint32) {
if effectiveMTU == 0 || packet.protocol != uint8(header.TCPProtocolNumber) {
return
}
transport := packet.transport
if len(transport) < header.TCPMinimumSize {
return
}
tcpHdr := header.TCP(transport)
tcpHeaderLength := int(tcpHdr.DataOffset())
if tcpHeaderLength < header.TCPMinimumSize || tcpHeaderLength > len(transport) {
return
}
var networkHeaderLength int
switch packet.ipVersion {
case 4:
networkHeaderLength = len(packet.network.(header.IPv4)) - len(transport)
default:
networkHeaderLength = len(packet.network.(header.IPv6)) - len(transport)
}
if effectiveMTU <= uint32(networkHeaderLength+header.TCPMinimumSize) {
return
}
maxMSS := min(effectiveMTU-uint32(networkHeaderLength+header.TCPMinimumSize), header.TCPMaximumMSS)
options := tcpHdr.Options()
for i := 0; i < len(options); {
switch options[i] {
case header.TCPOptionEOL:
return
case header.TCPOptionNOP:
i++
continue
case header.TCPOptionMSS:
if i+header.TCPOptionMSSLength > len(options) || options[i+1] != header.TCPOptionMSSLength {
return
}
currentMSS := binary.BigEndian.Uint16(options[i+2:])
if uint32(currentMSS) <= maxMSS {
return
}
binary.BigEndian.PutUint16(options[i+2:], uint16(maxMSS))
return
default:
if i+2 > len(options) {
return
}
optionLength := int(options[i+1])
if optionLength < 2 || i+optionLength > len(options) {
return
}
i += optionLength
}
}
}
func rewriteEmbeddedDestination(embedded *embeddedPacket, destination tcpip.Address, selector uint16, remapSelector bool) {
oldDestination := embedded.network.DestinationAddress()
if ipHdr, isIPv4 := embedded.network.(header.IPv4); isIPv4 {
ipHdr.SetDestinationAddressWithChecksumUpdate(destination)
} else {
embedded.network.SetDestinationAddress(destination)
}
rewriteEmbeddedSelector(embedded, oldDestination, destination, selector, remapSelector, true)
}
func rewriteEmbeddedSource(embedded *embeddedPacket, source tcpip.Address, selector uint16, remapSelector bool) {
oldSource := embedded.network.SourceAddress()
if ipHdr, isIPv4 := embedded.network.(header.IPv4); isIPv4 {
ipHdr.SetSourceAddressWithChecksumUpdate(source)
} else {
embedded.network.SetSourceAddress(source)
}
rewriteEmbeddedSelector(embedded, oldSource, source, selector, remapSelector, false)
}
func rewriteEmbeddedSelector(embedded *embeddedPacket, oldAddress, newAddress tcpip.Address, selector uint16, remapSelector bool, destinationSide bool) {
if !remapSelector {
return
}
payload := embedded.payload
_, isIPv4 := embedded.network.(header.IPv4)
switch embedded.protocol {
case uint8(header.TCPProtocolNumber):
if len(payload) >= 4 {
if destinationSide {
binary.BigEndian.PutUint16(payload[2:], selector)
} else {
binary.BigEndian.PutUint16(payload[0:], selector)
}
}
case uint8(header.UDPProtocolNumber):
if len(payload) >= header.UDPMinimumSize {
udpHdr := header.UDP(payload)
if isIPv4 && udpHdr.Checksum() == 0 {
if destinationSide {
udpHdr.SetDestinationPort(selector)
} else {
udpHdr.SetSourcePort(selector)
}
} else {
if oldAddress != newAddress {
udpHdr.UpdateChecksumPseudoHeaderAddress(oldAddress, newAddress, true)
}
if destinationSide {
udpHdr.SetDestinationPortWithChecksumUpdate(selector)
} else {
udpHdr.SetSourcePortWithChecksumUpdate(selector)
}
}
}
case uint8(header.ICMPv4ProtocolNumber):
if len(payload) >= header.ICMPv4MinimumSize {
header.ICMPv4(payload).SetIdentWithChecksumUpdate(selector)
}
case uint8(header.ICMPv6ProtocolNumber):
if len(payload) >= header.ICMPv6MinimumSize {
icmpHdr := header.ICMPv6(payload)
if oldAddress != newAddress {
icmpHdr.UpdateChecksumPseudoHeaderAddress(oldAddress, newAddress)
}
icmpHdr.SetIdentWithChecksumUpdate(selector)
}
}
}

View file

@ -4,15 +4,12 @@ package tun
import (
"context"
"errors"
"net/netip"
"sync/atomic"
"github.com/sagernet/sing-tun/gtcpip/header"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/florianl/go-nfqueue/v2"
"github.com/mdlayher/netlink"
@ -105,9 +102,9 @@ const ipv6AuthenticationHeaderIdentifier header.IPv6ExtensionHeaderIdentifier =
type preMatchPacket struct {
protocol uint8
network string
source M.Socksaddr
destination M.Socksaddr
source netip.AddrPort
destination netip.AddrPort
firstPacket []byte
}
func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
@ -161,20 +158,23 @@ func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
if !flags.Contains(header.TCPFlagSyn) || flags.Contains(header.TCPFlagAck) {
return preMatchPacket{}, false
}
parsed.network = N.NetworkTCP
parsed.source = M.SocksaddrFrom(source, tcpHdr.SourcePort())
parsed.destination = M.SocksaddrFrom(destination, tcpHdr.DestinationPort())
parsed.source = netip.AddrPortFrom(source, tcpHdr.SourcePort())
parsed.destination = netip.AddrPortFrom(destination, tcpHdr.DestinationPort())
case uint8(header.UDPProtocolNumber):
if len(transport) < header.UDPMinimumSize {
return preMatchPacket{}, false
}
udpHdr := header.UDP(transport)
if int(udpHdr.Length()) < header.UDPMinimumSize {
udpLength := int(udpHdr.Length())
if udpLength < header.UDPMinimumSize {
return preMatchPacket{}, false
}
parsed.network = N.NetworkUDP
parsed.source = M.SocksaddrFrom(source, udpHdr.SourcePort())
parsed.destination = M.SocksaddrFrom(destination, udpHdr.DestinationPort())
if udpLength < len(transport) {
transport = transport[:udpLength]
}
parsed.source = netip.AddrPortFrom(source, udpHdr.SourcePort())
parsed.destination = netip.AddrPortFrom(destination, udpHdr.DestinationPort())
parsed.firstPacket = header.UDP(transport).Payload()
case uint8(header.ICMPv4ProtocolNumber):
if !source.Is4() || len(transport) < header.ICMPv4MinimumSize {
return preMatchPacket{}, false
@ -183,9 +183,9 @@ func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
return preMatchPacket{}, false
}
parsed.network = N.NetworkICMP
parsed.source = M.SocksaddrFrom(source, 0)
parsed.destination = M.SocksaddrFrom(destination, 0)
identifier := icmpHdr.Ident()
parsed.source = netip.AddrPortFrom(source, identifier)
parsed.destination = netip.AddrPortFrom(destination, identifier)
case uint8(header.ICMPv6ProtocolNumber):
if !source.Is6() || len(transport) < header.ICMPv6MinimumSize {
return preMatchPacket{}, false
@ -194,9 +194,9 @@ func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return preMatchPacket{}, false
}
parsed.network = N.NetworkICMP
parsed.source = M.SocksaddrFrom(source, 0)
parsed.destination = M.SocksaddrFrom(destination, 0)
identifier := icmpHdr.Ident()
parsed.source = netip.AddrPortFrom(source, identifier)
parsed.destination = netip.AddrPortFrom(destination, identifier)
default:
return preMatchPacket{}, false
}
@ -265,22 +265,26 @@ func (h *nfqueueHandler) handlePacket(attr nfqueue.Attribute) int {
return 0
}
_, pErr := h.handler.PrepareConnection(packet.network, packet.source, packet.destination, nil, 0)
verdict := h.handler.JudgeFlow(
packet.protocol,
packet.source,
packet.destination,
)
// Use NfRepeat for bypass/reset so the packet re-enters the chain
// from the beginning, allowing mark-checking rules to save the mark
// to conntrack. NfAccept is a terminal verdict in nftables — it exits
// the chain immediately, skipping any rules after the queue statement.
switch {
case errors.Is(pErr, ErrBypass):
switch verdict.Action {
case ActionBypass:
h.setVerdict(packetID, nfqueue.NfRepeat, h.outputMark)
case errors.Is(pErr, ErrReset):
case ActionReject:
if packet.protocol == uint8(unix.IPPROTO_TCP) {
h.setVerdict(packetID, nfqueue.NfRepeat, h.resetMark)
} else {
h.setVerdict(packetID, nfqueue.NfAccept, 0)
}
case errors.Is(pErr, ErrDrop):
case ActionDrop:
h.setVerdict(packetID, nfqueue.NfDrop, 0)
default:
h.setVerdict(packetID, nfqueue.NfAccept, 0)

View file

@ -9,7 +9,6 @@ import (
"sync"
"time"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control"
@ -20,14 +19,16 @@ import (
// 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
var _ tun.DirectRouteDestination = (*Destination)(nil)
type PacketWriter interface {
WritePacket(packet []byte) error
}
type Destination struct {
conn *Conn
ctx context.Context
logger logger.ContextLogger
destination netip.Addr
routeContext tun.DirectRouteContext
writer PacketWriter
timeout time.Duration
requestAccess sync.Mutex
requests map[pingRequest]time.Time
@ -45,9 +46,9 @@ func ConnectDestination(
logger logger.ContextLogger,
controlFunc control.Func,
destination netip.Addr,
routeContext tun.DirectRouteContext,
writer PacketWriter,
timeout time.Duration,
) (tun.DirectRouteDestination, error) {
) (*Destination, error) {
var (
conn *Conn
err error
@ -69,7 +70,7 @@ func ConnectDestination(
ctx: ctx,
logger: logger,
destination: destination,
routeContext: routeContext,
writer: writer,
timeout: timeout,
requests: make(map[pingRequest]time.Time),
}
@ -158,7 +159,7 @@ func (d *Destination) loopRead() {
}
d.logger.TraceContext(d.ctx, "read ICMPv6 echo reply from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence())
}
err = d.routeContext.WritePacket(buffer.Bytes())
err = d.writer.WritePacket(buffer.Bytes())
if err != nil {
d.logger.ErrorContext(d.ctx, E.Cause(err, "write ICMP echo reply"))
}

View file

@ -1,143 +0,0 @@
//go:build with_gvisor
package ping
import (
"context"
"net/netip"
"time"
"github.com/sagernet/gvisor/pkg/tcpip"
"github.com/sagernet/gvisor/pkg/tcpip/adapters/gonet"
"github.com/sagernet/gvisor/pkg/tcpip/header"
"github.com/sagernet/gvisor/pkg/tcpip/stack"
"github.com/sagernet/gvisor/pkg/tcpip/transport"
"github.com/sagernet/gvisor/pkg/waiter"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
)
var _ tun.DirectRouteDestination = (*GVisorDestination)(nil)
type GVisorDestination struct {
ctx context.Context
logger logger.ContextLogger
endpoint tcpip.Endpoint
conn *gonet.TCPConn
rewriter *SourceRewriter
timeout time.Duration
lastActive common.TypedValue[time.Time]
}
func ConnectGVisor(
ctx context.Context, logger logger.ContextLogger,
sourceAddress, destinationAddress netip.Addr,
routeContext tun.DirectRouteContext,
stack *stack.Stack,
bindAddress4, bindAddress6 netip.Addr,
timeout time.Duration,
) (*GVisorDestination, error) {
var (
bindAddress tcpip.Address
wq waiter.Queue
endpoint tcpip.Endpoint
gErr tcpip.Error
)
if !destinationAddress.Is6() {
if !bindAddress4.IsValid() {
return nil, E.New("missing IPv4 interface address")
}
bindAddress = tun.AddressFromAddr(bindAddress4)
endpoint, gErr = stack.NewRawEndpoint(header.ICMPv4ProtocolNumber, header.IPv4ProtocolNumber, &wq, true)
} else {
if !bindAddress6.IsValid() {
return nil, E.New("missing IPv6 interface address")
}
bindAddress = tun.AddressFromAddr(bindAddress6)
endpoint, gErr = stack.NewRawEndpoint(header.ICMPv6ProtocolNumber, header.IPv6ProtocolNumber, &wq, true)
}
if gErr != nil {
return nil, gonet.TranslateNetstackError(gErr)
}
gErr = endpoint.Bind(tcpip.FullAddress{
NIC: 1,
Addr: bindAddress,
})
if gErr != nil {
return nil, gonet.TranslateNetstackError(gErr)
}
gErr = endpoint.Connect(tcpip.FullAddress{
NIC: 1,
Addr: tun.AddressFromAddr(destinationAddress),
})
if gErr != nil {
return nil, gonet.TranslateNetstackError(gErr)
}
endpoint.SocketOptions().SetHeaderIncluded(true)
rewriter := NewSourceRewriter(ctx, logger, bindAddress4, bindAddress6)
rewriter.CreateSession(tun.DirectRouteSession{Source: sourceAddress, Destination: destinationAddress}, routeContext)
destination := &GVisorDestination{
ctx: ctx,
logger: logger,
endpoint: endpoint,
conn: gonet.NewTCPConn(&wq, endpoint),
rewriter: rewriter,
timeout: timeout,
}
destination.lastActive.Store(time.Now())
go destination.loopRead()
return destination, nil
}
func (d *GVisorDestination) loopRead() {
defer d.endpoint.Close()
for {
deadline := d.lastActive.Load().Add(d.timeout)
if !time.Now().Before(deadline) {
return
}
err := d.conn.SetReadDeadline(deadline)
if err != nil {
d.logger.ErrorContext(d.ctx, E.Cause(err, "set read deadline for ICMP conn"))
}
buffer := buf.NewSize(maxICMPPacketSize)
n, err := d.conn.Read(buffer.FreeBytes())
if err != nil {
buffer.Release()
if E.IsTimeout(err) {
continue
}
if !E.IsClosed(err) {
d.logger.ErrorContext(d.ctx, E.Cause(err, "receive ICMP echo reply"))
}
return
}
buffer.Truncate(n)
var matched bool
matched, err = d.rewriter.WriteBack(buffer.Bytes())
if err != nil {
d.logger.ErrorContext(d.ctx, E.Cause(err, "write ICMP echo reply"))
}
if matched {
d.lastActive.Store(time.Now())
}
buffer.Release()
}
}
func (d *GVisorDestination) WritePacket(packet *buf.Buffer) error {
d.lastActive.Store(time.Now())
d.rewriter.RewritePacket(packet.Bytes())
return common.Error(d.conn.Write(packet.Bytes()))
}
func (d *GVisorDestination) Close() error {
return d.conn.Close()
}
func (d *GVisorDestination) IsClosed() bool {
return transport.DatagramEndpointState(d.endpoint.State()) == transport.DatagramEndpointStateClosed
}

View file

@ -1,79 +0,0 @@
package ping
import (
"net/netip"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/buf"
)
type DestinationWriter struct {
tun.DirectRouteDestination
destination netip.Addr
}
func NewDestinationWriter(routeDestination tun.DirectRouteDestination, destination netip.Addr) *DestinationWriter {
return &DestinationWriter{routeDestination, destination}
}
func (w *DestinationWriter) WritePacket(packet *buf.Buffer) error {
var ipHdr header.Network
switch header.IPVersion(packet.Bytes()) {
case header.IPv4Version:
ipHdr = header.IPv4(packet.Bytes())
case header.IPv6Version:
ipHdr = header.IPv6(packet.Bytes())
default:
return w.DirectRouteDestination.WritePacket(packet)
}
ipHdr.SetDestinationAddr(w.destination)
if ipHdr4, isIPv4 := ipHdr.(header.IPv4); isIPv4 {
ipHdr4.SetChecksum(^ipHdr4.CalculateChecksum())
}
if ipHdr.TransportProtocol() == header.ICMPv6ProtocolNumber {
icmpHdr := header.ICMPv6(ipHdr.Payload())
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr,
Src: ipHdr.SourceAddressSlice(),
Dst: ipHdr.DestinationAddressSlice(),
}))
}
return w.DirectRouteDestination.WritePacket(packet)
}
type ContextDestinationWriter struct {
tun.DirectRouteContext
destination netip.Addr
}
func NewContextDestinationWriter(context tun.DirectRouteContext, destination netip.Addr) *ContextDestinationWriter {
return &ContextDestinationWriter{
context, destination,
}
}
func (w *ContextDestinationWriter) WritePacket(packet []byte) error {
var ipHdr header.Network
switch header.IPVersion(packet) {
case header.IPv4Version:
ipHdr = header.IPv4(packet)
case header.IPv6Version:
ipHdr = header.IPv6(packet)
default:
return w.DirectRouteContext.WritePacket(packet)
}
ipHdr.SetSourceAddr(w.destination)
if ipHdr4, isIPv4 := ipHdr.(header.IPv4); isIPv4 {
ipHdr4.SetChecksum(^ipHdr4.CalculateChecksum())
}
if ipHdr.TransportProtocol() == header.ICMPv6ProtocolNumber {
icmpHdr := header.ICMPv6(ipHdr.Payload())
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr,
Src: ipHdr.SourceAddressSlice(),
Dst: ipHdr.DestinationAddressSlice(),
}))
}
return w.DirectRouteContext.WritePacket(packet)
}

194
ping/port.go Normal file
View file

@ -0,0 +1,194 @@
package ping
import (
"context"
"net/netip"
"slices"
"sync"
"time"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
)
const defaultFlowTimeout = time.Minute
type Port struct {
ctx context.Context
logger logger.ContextLogger
controlFunc func(destination netip.Addr) control.Func
timeout time.Duration
returnAccess sync.Mutex
returnPaths []tun.Return
flowAccess sync.Mutex
flows map[flowKey]*Destination
lastSweep time.Time
}
type flowKey struct {
source netip.Addr
destination netip.Addr
identifier uint16
}
func NewPort(ctx context.Context, logger logger.ContextLogger, controlFunc func(destination netip.Addr) control.Func, timeout time.Duration) *Port {
if timeout <= 0 {
timeout = defaultFlowTimeout
}
return &Port{
ctx: ctx,
logger: logger,
controlFunc: controlFunc,
timeout: timeout,
flows: make(map[flowKey]*Destination),
}
}
func (p *Port) PortAddresses() (netip.Addr, netip.Addr) {
return netip.IPv4Unspecified(), netip.IPv6Unspecified()
}
func (p *Port) PortMTU() uint32 {
return 0
}
func (p *Port) AttachReturn(returnPath tun.Return) error {
p.returnAccess.Lock()
defer p.returnAccess.Unlock()
if slices.Contains(p.returnPaths, returnPath) {
return nil
}
p.returnPaths = append(p.returnPaths[:len(p.returnPaths):len(p.returnPaths)], returnPath)
return nil
}
func (p *Port) DetachReturn(returnPath tun.Return) error {
p.returnAccess.Lock()
defer p.returnAccess.Unlock()
returnPaths := make([]tun.Return, 0, len(p.returnPaths))
for _, existing := range p.returnPaths {
if existing != returnPath {
returnPaths = append(returnPaths, existing)
}
}
p.returnPaths = returnPaths
return nil
}
func (p *Port) WritePackets(packets [][]byte) error {
var errs []error
for _, packet := range packets {
err := p.writePacket(packet)
if err != nil {
errs = append(errs, err)
}
}
return E.Errors(errs...)
}
func (p *Port) writePacket(packet []byte) error {
var (
source netip.Addr
destination netip.Addr
identifier uint16
)
switch header.IPVersion(packet) {
case header.IPv4Version:
ipHdr := header.IPv4(packet)
if !ipHdr.IsValid(len(packet)) || ipHdr.TransportProtocol() != header.ICMPv4ProtocolNumber || ipHdr.PayloadLength() < header.ICMPv4MinimumSize {
return nil
}
icmpHdr := header.ICMPv4(ipHdr.Payload())
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
return nil
}
source = ipHdr.SourceAddr()
destination = ipHdr.DestinationAddr()
identifier = icmpHdr.Ident()
case header.IPv6Version:
ipHdr := header.IPv6(packet)
if !ipHdr.IsValid(len(packet)) || ipHdr.TransportProtocol() != header.ICMPv6ProtocolNumber || ipHdr.PayloadLength() < header.ICMPv6MinimumSize {
return nil
}
icmpHdr := header.ICMPv6(ipHdr.Payload())
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return nil
}
source = ipHdr.SourceAddr()
destination = ipHdr.DestinationAddr()
identifier = icmpHdr.Ident()
default:
return nil
}
flow, err := p.flowFor(source, destination, identifier)
if err != nil {
return E.Cause(err, "connect ICMP flow to ", destination)
}
return flow.WritePacket(buf.As(packet))
}
func (p *Port) flowFor(source netip.Addr, destination netip.Addr, identifier uint16) (*Destination, error) {
key := flowKey{source: source, destination: destination, identifier: identifier}
p.flowAccess.Lock()
defer p.flowAccess.Unlock()
now := time.Now()
if now.Sub(p.lastSweep) >= p.timeout {
p.lastSweep = now
for oldKey, oldFlow := range p.flows {
if oldFlow.IsClosed() {
delete(p.flows, oldKey)
}
}
}
flow, loaded := p.flows[key]
if loaded && !flow.IsClosed() {
return flow, nil
}
var controlFunc control.Func
if p.controlFunc != nil {
controlFunc = p.controlFunc(destination)
}
flow, err := ConnectDestination(p.ctx, p.logger, controlFunc, destination, portWriter{p}, p.timeout)
if err != nil {
return nil, err
}
p.flows[key] = flow
return flow, nil
}
type portWriter struct {
port *Port
}
func (w portWriter) WritePacket(packet []byte) error {
w.port.returnAccess.Lock()
returnPaths := w.port.returnPaths
w.port.returnAccess.Unlock()
for _, returnPath := range returnPaths {
headroom := returnPath.ReturnHeadroom()
buffer := make([]byte, headroom+len(packet))
copy(buffer[headroom:], packet)
unconsumed := returnPath.ReturnPackets([][]byte{buffer})
if len(unconsumed) == 0 {
return nil
}
}
return nil
}
func (p *Port) Close() error {
p.flowAccess.Lock()
defer p.flowAccess.Unlock()
var errs []error
for key, flow := range p.flows {
errs = append(errs, flow.Close())
delete(p.flows, key)
}
return E.Errors(errs...)
}

View file

@ -1,150 +0,0 @@
package ping
import (
"context"
"net/netip"
"sync"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/logger"
)
type SourceRewriter struct {
ctx context.Context
logger logger.ContextLogger
access sync.RWMutex
sessions map[tun.DirectRouteSession]tun.DirectRouteContext
sourceAddress map[uint16]netip.Addr
inet4Address netip.Addr
inet6Address netip.Addr
}
func NewSourceRewriter(ctx context.Context, logger logger.ContextLogger, inet4Address netip.Addr, inet6Address netip.Addr) *SourceRewriter {
return &SourceRewriter{
ctx: ctx,
logger: logger,
sessions: make(map[tun.DirectRouteSession]tun.DirectRouteContext),
sourceAddress: make(map[uint16]netip.Addr),
inet4Address: inet4Address,
inet6Address: inet6Address,
}
}
func (m *SourceRewriter) CreateSession(session tun.DirectRouteSession, context tun.DirectRouteContext) {
m.access.Lock()
m.sessions[session] = context
m.access.Unlock()
}
func (m *SourceRewriter) DeleteSession(session tun.DirectRouteSession) {
m.access.Lock()
delete(m.sessions, session)
m.access.Unlock()
}
func (m *SourceRewriter) RewritePacket(packet []byte) {
var ipHdr header.Network
var bindAddr netip.Addr
switch header.IPVersion(packet) {
case header.IPv4Version:
ipHdr = header.IPv4(packet)
bindAddr = m.inet4Address
case header.IPv6Version:
ipHdr = header.IPv6(packet)
bindAddr = m.inet6Address
default:
return
}
sourceAddr := ipHdr.SourceAddr()
ipHdr.SetSourceAddr(bindAddr)
if ipHdr4, isIPv4 := ipHdr.(header.IPv4); isIPv4 {
ipHdr4.SetChecksum(^ipHdr4.CalculateChecksum())
}
switch ipHdr.TransportProtocol() {
case header.ICMPv4ProtocolNumber:
icmpHdr := header.ICMPv4(ipHdr.Payload())
m.access.Lock()
m.sourceAddress[icmpHdr.Ident()] = sourceAddr
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:
icmpHdr := header.ICMPv6(ipHdr.Payload())
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr,
Src: ipHdr.SourceAddressSlice(),
Dst: ipHdr.DestinationAddressSlice(),
}))
m.access.Lock()
m.sourceAddress[icmpHdr.Ident()] = sourceAddr
m.access.Unlock()
m.logger.TraceContext(m.ctx, "write ICMPv6 echo request from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence())
}
}
func (m *SourceRewriter) WriteBack(packet []byte) (bool, error) {
var ipHdr header.Network
var routeSession tun.DirectRouteSession
switch header.IPVersion(packet) {
case header.IPv4Version:
ipHdr = header.IPv4(packet)
routeSession.Destination = ipHdr.SourceAddr()
case header.IPv6Version:
ipHdr = header.IPv6(packet)
routeSession.Destination = ipHdr.SourceAddr()
default:
return false, nil
}
switch ipHdr.TransportProtocol() {
case header.ICMPv4ProtocolNumber:
icmpHdr := header.ICMPv4(ipHdr.Payload())
m.access.Lock()
ident := icmpHdr.Ident()
source, loaded := m.sourceAddress[ident]
if !loaded {
m.access.Unlock()
return false, nil
}
delete(m.sourceAddress, icmpHdr.Ident())
m.access.Unlock()
routeSession.Source = source
case header.ICMPv6ProtocolNumber:
icmpHdr := header.ICMPv6(ipHdr.Payload())
m.access.Lock()
ident := icmpHdr.Ident()
source, loaded := m.sourceAddress[ident]
if !loaded {
m.access.Unlock()
return false, nil
}
delete(m.sourceAddress, icmpHdr.Ident())
m.access.Unlock()
routeSession.Source = source
default:
return false, nil
}
m.access.RLock()
context, loaded := m.sessions[routeSession]
m.access.RUnlock()
if !loaded {
return false, nil
}
ipHdr.SetDestinationAddr(routeSession.Source)
if ipHdr4, isIPv4 := ipHdr.(header.IPv4); isIPv4 {
ipHdr4.SetChecksum(^ipHdr4.CalculateChecksum())
}
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:
icmpHdr := header.ICMPv6(ipHdr.Payload())
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr,
Src: ipHdr.SourceAddressSlice(),
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)
}

View file

@ -139,6 +139,7 @@ func (r *autoRedirect) Start() error {
r.redirectServer = server
}
if r.useNFTables {
if r.handler != nil {
var handler *nfqueueHandler
handler, err = newNFQueueHandler(nfqueueOptions{
Context: r.ctx,
@ -156,6 +157,7 @@ func (r *autoRedirect) Start() error {
r.nfqueueHandler = handler
r.nfqueueEnabled = true
}
}
r.cleanupNFTables()
err = r.setupNFTables()
if err != nil {

View file

@ -1,61 +0,0 @@
package tun
import (
"net/netip"
"time"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/contrab/freelru"
"github.com/sagernet/sing/contrab/maphash"
)
type DirectRouteDestination interface {
WritePacket(packet *buf.Buffer) error
Close() error
IsClosed() bool
}
type DirectRouteSession struct {
// IPVersion uint8
// Network uint8
Source netip.Addr
Destination netip.Addr
}
type DirectRouteMapping struct {
mapping freelru.Cache[DirectRouteSession, DirectRouteDestination]
timeout time.Duration
}
func NewDirectRouteMapping(timeout time.Duration) *DirectRouteMapping {
mapping := common.Must1(freelru.NewSharded[DirectRouteSession, DirectRouteDestination](1024, maphash.NewHasher[DirectRouteSession]().Hash32))
mapping.SetHealthCheck(func(session DirectRouteSession, action DirectRouteDestination) bool {
if action != nil {
return !action.IsClosed()
}
return true
})
mapping.SetOnEvict(func(session DirectRouteSession, action DirectRouteDestination) {
if action != nil {
action.Close()
}
})
mapping.SetLifetime(timeout)
return &DirectRouteMapping{mapping, timeout}
}
func (m *DirectRouteMapping) Lookup(session DirectRouteSession, constructor func(timeout time.Duration) (DirectRouteDestination, error)) (DirectRouteDestination, error) {
var (
created DirectRouteDestination
err error
)
action, _, ok := m.mapping.GetAndRefreshOrAdd(session, func() (DirectRouteDestination, bool) {
created, err = constructor(m.timeout)
return created, err == nil
})
if !ok {
return nil, err
}
return action, nil
}

View file

@ -12,12 +12,6 @@ import (
"github.com/sagernet/sing/common/logger"
)
var (
ErrDrop = E.New("drop by rule")
ErrReset = E.New("reset by rule")
ErrBypass = E.New("bypass by rule")
)
type Stack interface {
Start() error
Close() error

View file

@ -6,8 +6,10 @@ import (
"context"
"net/netip"
"runtime"
"sync"
"time"
"github.com/sagernet/gvisor/pkg/buffer"
"github.com/sagernet/gvisor/pkg/tcpip"
"github.com/sagernet/gvisor/pkg/tcpip/adapters/gonet"
"github.com/sagernet/gvisor/pkg/tcpip/header"
@ -40,6 +42,8 @@ type GVisor struct {
logger logger.Logger
stack *stack.Stack
endpoint stack.LinkEndpoint
dispatcher *ForwardDispatcher
icmpForwarder *ICMPForwarder
}
type GVisorTun interface {
@ -88,23 +92,39 @@ func (t *GVisor) Start() error {
if err != nil {
return err
}
linkEndpoint = &LinkEndpointFilter{linkEndpoint, t.broadcastAddr, t.tun}
if t.handler != nil {
t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpTimeout, t.icmpTimeout)
}
linkEndpoint = &LinkEndpointFilter{
LinkEndpoint: linkEndpoint,
BroadcastAddress: t.broadcastAddr,
Writer: t.tun,
Dispatcher: t.dispatcher,
Inet4Address: t.inet4Address,
Inet6Address: t.inet6Address,
Inet4LoopbackAddress: t.inet4LoopbackAddress,
Inet6LoopbackAddress: t.inet6LoopbackAddress,
}
ipStack, err := newGVisorStack(linkEndpoint, nicOptions, false, true)
if err != nil {
return err
}
ipStack.SetTransportProtocolHandler(tcp.ProtocolNumber, NewTCPForwarderWithLoopback(t.ctx, ipStack, t.handler, t.inet4LoopbackAddress, t.inet6LoopbackAddress, t.tun).HandlePacket)
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpTimeout).HandlePacket)
icmpForwarder := NewICMPForwarder(t.ctx, ipStack, t.logger, t.handler, t.icmpTimeout)
icmpForwarder.SetLocalAddresses(t.inet4Address, t.inet6Address)
icmpForwarder := NewICMPForwarder(ipStack, t.handler, t.logger)
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket)
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket)
t.icmpForwarder = icmpForwarder
t.stack = ipStack
t.endpoint = linkEndpoint
return nil
}
func (t *GVisor) Close() error {
t.dispatcher.Close()
if t.icmpForwarder != nil {
t.icmpForwarder.Close()
}
if t.stack == nil {
return nil
}
@ -116,6 +136,37 @@ func (t *GVisor) Close() error {
return nil
}
type gvisorWriteback struct {
tun GVisorTun
access sync.Mutex
}
func (w *gvisorWriteback) ReturnHeadroom() int {
return 0
}
func (w *gvisorWriteback) WriteReturnPackets(packets [][]byte) error {
w.access.Lock()
defer w.access.Unlock()
var writeErrors []error
for _, packet := range packets {
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(packet),
})
if header.IPVersion(packet) == header.IPv6Version {
pkt.NetworkProtocolNumber = header.IPv6ProtocolNumber
} else {
pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber
}
_, err := w.tun.WritePacket(pkt)
pkt.DecRef()
if err != nil {
writeErrors = append(writeErrors, err)
}
}
return E.Errors(writeErrors...)
}
func AddressFromAddr(destination netip.Addr) tcpip.Address {
if destination.Is6() {
return tcpip.AddrFrom16(destination.As16())

View file

@ -16,10 +16,24 @@ type LinkEndpointFilter struct {
stack.LinkEndpoint
BroadcastAddress netip.Addr
Writer GVisorTun
Dispatcher *ForwardDispatcher
Inet4Address netip.Addr
Inet6Address netip.Addr
Inet4LoopbackAddress []netip.Addr
Inet6LoopbackAddress []netip.Addr
}
func (w *LinkEndpointFilter) Attach(dispatcher stack.NetworkDispatcher) {
w.LinkEndpoint.Attach(&networkDispatcherFilter{dispatcher, w.BroadcastAddress, w.Writer})
w.LinkEndpoint.Attach(&networkDispatcherFilter{
NetworkDispatcher: dispatcher,
broadcastAddress: w.BroadcastAddress,
writer: w.Writer,
dispatcher: w.Dispatcher,
inet4Address: w.Inet4Address,
inet6Address: w.Inet6Address,
inet4LoopbackAddress: w.Inet4LoopbackAddress,
inet6LoopbackAddress: w.Inet6LoopbackAddress,
})
}
var _ stack.NetworkDispatcher = (*networkDispatcherFilter)(nil)
@ -28,6 +42,11 @@ type networkDispatcherFilter struct {
stack.NetworkDispatcher
broadcastAddress netip.Addr
writer GVisorTun
dispatcher *ForwardDispatcher
inet4Address netip.Addr
inet6Address netip.Addr
inet4LoopbackAddress []netip.Addr
inet6LoopbackAddress []netip.Addr
}
func (w *networkDispatcherFilter) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) {
@ -50,5 +69,45 @@ func (w *networkDispatcherFilter) DeliverNetworkPacket(protocol tcpip.NetworkPro
w.writer.WritePacket(pkt)
return
}
if w.dispatcher != nil && pkt.GSOOptions.Type == stack.GSONone && !pkt.GSOOptions.NeedsCsum {
if view, loaded := pkt.Data().PullUp(pkt.Data().Size()); loaded {
consumed := w.dispatch(protocol, destination, view)
w.dispatcher.Flush()
if consumed {
return
}
}
}
w.NetworkDispatcher.DeliverNetworkPacket(protocol, pkt)
}
func (w *networkDispatcherFilter) dispatch(protocol tcpip.NetworkProtocolNumber, destination netip.Addr, view []byte) bool {
if protocol == header.IPv4ProtocolNumber {
switch header.IPv4(view).TransportProtocol() {
case header.TCPProtocolNumber:
for _, inet4LoopbackAddress := range w.inet4LoopbackAddress {
if destination == inet4LoopbackAddress {
return false
}
}
case header.ICMPv4ProtocolNumber:
if destination == w.inet4Address {
return false
}
}
} else {
switch header.IPv6(view).TransportProtocol() {
case header.TCPProtocolNumber:
for _, inet6LoopbackAddress := range w.inet6LoopbackAddress {
if destination == inet6LoopbackAddress {
return false
}
}
case header.ICMPv6ProtocolNumber:
if destination == w.inet6Address {
return false
}
}
}
return w.dispatcher.Dispatch(view)
}

View file

@ -3,10 +3,9 @@
package tun
import (
"context"
"errors"
"net/netip"
"sync"
"sync/atomic"
"time"
"github.com/sagernet/gvisor/pkg/buffer"
@ -17,42 +16,51 @@ import (
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv4"
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv6"
"github.com/sagernet/gvisor/pkg/tcpip/stack"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
type ICMPForwarder struct {
ctx context.Context
stack *stack.Stack
logger logger.Logger
inet4Address netip.Addr
inet6Address netip.Addr
handler Handler
mapping *DirectRouteMapping
logger logger.Logger
returnPath icmpForwarderReturn
flowAccess sync.Mutex
flows map[icmpFlowKey]time.Time
lastSweep time.Time
attachedPorts map[Port]bool
}
func NewICMPForwarder(
ctx context.Context,
stack *stack.Stack,
logger logger.Logger,
handler Handler,
timeout time.Duration,
) *ICMPForwarder {
return &ICMPForwarder{
ctx: ctx,
type icmpFlowKey struct {
v6 bool
source netip.Addr
destination netip.Addr
identifier uint16
}
func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) *ICMPForwarder {
forwarder := &ICMPForwarder{
stack: stack,
logger: logger,
handler: handler,
mapping: NewDirectRouteMapping(timeout),
logger: logger,
flows: make(map[icmpFlowKey]time.Time),
attachedPorts: make(map[Port]bool),
}
forwarder.returnPath.forwarder = forwarder
return forwarder
}
func (f *ICMPForwarder) SetLocalAddresses(inet4Address, inet6Address netip.Addr) {
f.inet4Address = inet4Address
f.inet6Address = inet6Address
func (f *ICMPForwarder) Close() error {
f.returnPath.closed.Store(true)
f.flowAccess.Lock()
defer f.flowAccess.Unlock()
for port := range f.attachedPorts {
port.DetachReturn(&f.returnPath)
delete(f.attachedPorts, port)
}
return nil
}
func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
@ -62,34 +70,17 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
return false
}
sourceAddr := M.AddrFromIP(ipHdr.SourceAddressSlice())
destinationAddr := M.AddrFromIP(ipHdr.DestinationAddressSlice())
if destinationAddr != f.inet4Address {
action, err := f.mapping.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
return f.handler.PrepareConnection(
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&ICMPBackWriter{
stack: f.stack,
packet: pkt,
source: ipHdr.SourceAddress(),
sourceNetwork: header.IPv4ProtocolNumber,
},
timeout,
identifier := icmpHdr.Ident()
verdict := f.handler.JudgeFlow(
uint8(header.ICMPv4ProtocolNumber),
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
)
})
if errors.Is(err, ErrReset) {
gWriteUnreachable(f.stack, pkt)
switch verdict.Action {
case ActionReject, ActionDrop:
return true
} else if errors.Is(err, ErrDrop) {
return true
}
if action != nil {
err = icmpWritePacketBuffer(action, pkt)
if err != nil {
f.logger.Error(E.Cause(err, "write ICMPv4 echo request"))
}
case ActionFlow:
if f.forwardFlow(verdict.Port, false, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
return true
}
}
@ -125,35 +116,17 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return false
}
sourceAddr := M.AddrFromIP(ipHdr.SourceAddressSlice())
destinationAddr := M.AddrFromIP(ipHdr.DestinationAddressSlice())
if destinationAddr != f.inet6Address {
action, err := f.mapping.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
return f.handler.PrepareConnection(
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&ICMPBackWriter{
stack: f.stack,
packet: pkt,
source: ipHdr.SourceAddress(),
sourceNetwork: header.IPv6ProtocolNumber,
},
timeout,
identifier := icmpHdr.Ident()
verdict := f.handler.JudgeFlow(
uint8(header.ICMPv6ProtocolNumber),
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
)
})
if errors.Is(err, ErrReset) {
gWriteUnreachable(f.stack, pkt)
switch verdict.Action {
case ActionReject, ActionDrop:
return true
} else if errors.Is(err, ErrDrop) {
return true
}
if action != nil {
pkt.IncRef()
err = icmpWritePacketBuffer(action, pkt)
if err != nil {
f.logger.Error(E.Cause(err, "write ICMPv6 echo request"))
}
case ActionFlow:
if f.forwardFlow(verdict.Port, true, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
return true
}
}
@ -190,64 +163,179 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
}
}
type ICMPBackWriter struct {
access sync.Mutex
stack *stack.Stack
packet *stack.PacketBuffer
source tcpip.Address
sourceNetwork tcpip.NetworkProtocolNumber
func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, destination netip.Addr, identifier uint16, pkt *stack.PacketBuffer) bool {
if port == nil {
return false
}
inet4Address, inet6Address := port.PortAddresses()
portAddress := inet4Address
if v6 {
portAddress = inet6Address
}
if !portAddress.IsValid() || !portAddress.IsUnspecified() {
return false
}
f.flowAccess.Lock()
if !f.attachedPorts[port] {
err := port.AttachReturn(&f.returnPath)
if err != nil {
f.flowAccess.Unlock()
f.logger.Trace(E.Cause(err, "attach ICMP return path"))
return false
}
f.attachedPorts[port] = true
}
now := time.Now()
if now.Sub(f.lastSweep) >= defaultICMPTimeout {
f.lastSweep = now
for key, deadline := range f.flows {
if now.After(deadline) {
delete(f.flows, key)
}
}
}
f.flows[icmpFlowKey{v6: v6, source: source, destination: destination, identifier: identifier}] = now.Add(defaultICMPTimeout)
f.flowAccess.Unlock()
networkSlice := pkt.NetworkHeader().Slice()
transportSlice := pkt.TransportHeader().Slice()
dataSlice := pkt.Data().AsRange().ToSlice()
packetSlice := make([]byte, 0, len(networkSlice)+len(transportSlice)+len(dataSlice))
packetSlice = append(packetSlice, networkSlice...)
packetSlice = append(packetSlice, transportSlice...)
packetSlice = append(packetSlice, dataSlice...)
err := port.WritePackets([][]byte{packetSlice})
if err != nil {
f.logger.Trace(E.Cause(err, "forward ICMP packet"))
}
return true
}
func (w *ICMPBackWriter) WritePacket(p []byte) error {
if w.sourceNetwork == header.IPv4ProtocolNumber {
route, err := w.stack.FindRoute(
DefaultNIC,
header.IPv4(p).SourceAddress(),
w.source,
w.sourceNetwork,
false,
)
if err != nil {
return gonet.TranslateNetstackError(err)
func (f *ICMPForwarder) lookupFlow(key icmpFlowKey) bool {
f.flowAccess.Lock()
defer f.flowAccess.Unlock()
deadline, loaded := f.flows[key]
if !loaded {
return false
}
now := time.Now()
if now.After(deadline) {
delete(f.flows, key)
return false
}
f.flows[key] = now.Add(defaultICMPTimeout)
return true
}
type icmpForwarderReturn struct {
forwarder *ICMPForwarder
closed atomic.Bool
}
func (r *icmpForwarderReturn) ReturnHeadroom() int {
return 0
}
func (r *icmpForwarderReturn) ReturnPackets(packets [][]byte) [][]byte {
if r.closed.Load() {
return packets
}
unconsumed := packets[:0]
for _, packet := range packets {
if !r.forwarder.returnPacket(packet) {
unconsumed = append(unconsumed, packet)
}
}
return unconsumed
}
func (f *ICMPForwarder) returnPacket(packet []byte) bool {
if len(packet) == 0 {
return false
}
switch header.IPVersion(packet) {
case header.IPv4Version:
ipHdr := header.IPv4(packet)
if !ipHdr.IsValid(len(packet)) || ipHdr.TransportProtocol() != header.ICMPv4ProtocolNumber || len(ipHdr.Payload()) < header.ICMPv4MinimumSize {
return false
}
icmpHdr := header.ICMPv4(ipHdr.Payload())
var key icmpFlowKey
switch icmpHdr.Type() {
case header.ICMPv4EchoReply:
key = icmpFlowKey{
source: AddrFromAddress(ipHdr.DestinationAddress()),
destination: AddrFromAddress(ipHdr.SourceAddress()),
identifier: icmpHdr.Ident(),
}
case header.ICMPv4TimeExceeded, header.ICMPv4DstUnreachable:
inner := icmpHdr.Payload()
if len(inner) < header.IPv4MinimumSize {
return false
}
innerIPHdr := header.IPv4(inner)
innerHeaderLength := int(innerIPHdr.HeaderLength())
if innerHeaderLength < header.IPv4MinimumSize || len(inner) < innerHeaderLength+header.ICMPv4MinimumSize {
return false
}
if innerIPHdr.TransportProtocol() != header.ICMPv4ProtocolNumber {
return false
}
innerICMPHdr := header.ICMPv4(inner[innerHeaderLength:])
key = icmpFlowKey{
source: AddrFromAddress(innerIPHdr.SourceAddress()),
destination: AddrFromAddress(innerIPHdr.DestinationAddress()),
identifier: innerICMPHdr.Ident(),
}
default:
return false
}
if !f.lookupFlow(key) {
return false
}
return f.writeBack(packet, header.IPv4ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
case header.IPv6Version:
ipHdr := header.IPv6(packet)
if !ipHdr.IsValid(len(packet)) || ipHdr.TransportProtocol() != header.ICMPv6ProtocolNumber || len(ipHdr.Payload()) < header.ICMPv6MinimumSize {
return false
}
icmpHdr := header.ICMPv6(ipHdr.Payload())
if icmpHdr.Type() != header.ICMPv6EchoReply {
return false
}
key := icmpFlowKey{
v6: true,
source: AddrFromAddress(ipHdr.DestinationAddress()),
destination: AddrFromAddress(ipHdr.SourceAddress()),
identifier: icmpHdr.Ident(),
}
if !f.lookupFlow(key) {
return false
}
return f.writeBack(packet, header.IPv6ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
default:
return false
}
}
func (f *ICMPForwarder) writeBack(packet []byte, protocol tcpip.NetworkProtocolNumber, localAddress tcpip.Address, remoteAddress tcpip.Address) bool {
route, gErr := f.stack.FindRoute(DefaultNIC, localAddress, remoteAddress, protocol, false)
if gErr != nil {
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "find route for ICMP reply"))
return true
}
defer route.Release()
packet := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(p),
packetBuffer := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(packet),
})
defer packet.DecRef()
parse.IPv4(packet)
err = route.WritePacketDirect(packet)
if err != nil {
return gonet.TranslateNetstackError(err)
}
defer packetBuffer.DecRef()
if protocol == header.IPv4ProtocolNumber {
parse.IPv4(packetBuffer)
} else {
route, err := w.stack.FindRoute(
DefaultNIC,
header.IPv6(p).SourceAddress(),
w.source,
w.sourceNetwork,
false,
)
if err != nil {
return gonet.TranslateNetstackError(err)
parse.IPv6(packetBuffer)
}
defer route.Release()
packet := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(p),
})
parse.IPv6(packet)
defer packet.DecRef()
err = route.WritePacketDirect(packet)
if err != nil {
return gonet.TranslateNetstackError(err)
gErr = route.WritePacketDirect(packetBuffer)
if gErr != nil {
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "write ICMP reply"))
}
}
return nil
}
func icmpWritePacketBuffer(action DirectRouteDestination, packetBuffer *stack.PacketBuffer) error {
packetSlice := packetBuffer.NetworkHeader().Slice()
packetSlice = append(packetSlice, packetBuffer.TransportHeader().Slice()...)
packetSlice = append(packetSlice, packetBuffer.Data().AsRange().ToSlice()...)
return action.WritePacket(buf.As(packetSlice).ToOwned())
return true
}

View file

@ -4,7 +4,6 @@ package tun
import (
"context"
"errors"
"net"
"os"
"sync"
@ -74,7 +73,7 @@ func (c *gLazyConn) HandshakeFailure(err error) error {
if c.handshakeDone {
return os.ErrInvalid
}
c.request.Complete(!errors.Is(err, ErrDrop))
c.request.Complete(true)
c.handshakeDone = true
c.handshakeErr = err
return nil

View file

@ -4,7 +4,6 @@ package tun
import (
"context"
"errors"
"net/netip"
"github.com/sagernet/gvisor/pkg/tcpip"
@ -14,7 +13,6 @@ import (
"github.com/sagernet/sing-tun/gtcpip/checksum"
"github.com/sagernet/sing/common"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
type TCPForwarder struct {
@ -79,9 +77,12 @@ func (f *TCPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac
func (f *TCPForwarder) Forward(r *tcp.ForwarderRequest) {
source := M.SocksaddrFrom(AddrFromAddress(r.ID().RemoteAddress), r.ID().RemotePort)
destination := M.SocksaddrFrom(AddrFromAddress(r.ID().LocalAddress), r.ID().LocalPort)
_, pErr := f.handler.PrepareConnection(N.NetworkTCP, source, destination, nil, 0)
if pErr != nil {
r.Complete(!errors.Is(pErr, ErrDrop))
switch f.handler.JudgeFlow(uint8(header.TCPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action {
case ActionReject:
r.Complete(true)
return
case ActionDrop:
r.Complete(false)
return
}
conn := &gLazyConn{

View file

@ -4,7 +4,6 @@ package tun
import (
"context"
"errors"
"math"
"net/netip"
"os"
@ -58,11 +57,11 @@ func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac
func rangeIterate(r stack.Range, fn func(*buffer.View))
func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) {
_, pErr := f.handler.PrepareConnection(N.NetworkUDP, source, destination, nil, 0)
if pErr != nil {
if !errors.Is(pErr, ErrDrop) {
switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action {
case ActionReject:
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
}
return false, nil, nil, nil
case ActionDrop:
return false, nil, nil, nil
}
var sourceNetwork tcpip.NetworkProtocolNumber

View file

@ -73,8 +73,6 @@ func (m *Mixed) tunLoop() {
return
}
if linuxTUN, isLinuxTUN := m.tun.(LinuxTUN); isLinuxTUN {
m.frontHeadroom = linuxTUN.FrontHeadroom()
m.txChecksumOffload = linuxTUN.TXChecksumOffload()
batchSize := linuxTUN.BatchSize()
if batchSize > 1 {
m.batchLoopLinux(linuxTUN, batchSize)
@ -105,6 +103,7 @@ func (m *Mixed) tunLoop() {
m.logger.Trace(E.Cause(err, "write packet"))
}
}
m.dispatcher.Flush()
}
}
@ -124,6 +123,7 @@ func (m *Mixed) wintunLoop(winTun WinTun) {
m.logger.Trace(E.Cause(err, "write packet"))
}
}
m.dispatcher.Flush()
release()
}
}
@ -164,11 +164,13 @@ func (m *Mixed) batchLoopLinux(linuxTUN LinuxTUN, batchSize int) {
}
writeBuffers = writeBuffers[:0]
}
m.dispatcher.Flush()
}
}
func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
var writeBuffers []*buf.Buffer
var releaseBuffers []*buf.Buffer
for {
buffers, err := darwinTUN.BatchRead()
if err != nil {
@ -181,6 +183,7 @@ func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
continue
}
writeBuffers = writeBuffers[:0]
releaseBuffers = releaseBuffers[:0]
for _, buffer := range buffers {
packetSize := buffer.Len()
if packetSize < header.IPv4MinimumSize {
@ -190,7 +193,7 @@ func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
if m.processPacket(buffer.Bytes()) {
writeBuffers = append(writeBuffers, buffer)
} else {
buffer.Release()
releaseBuffers = append(releaseBuffers, buffer)
}
}
if len(writeBuffers) > 0 {
@ -200,6 +203,8 @@ func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
}
buf.ReleaseMulti(writeBuffers)
}
m.dispatcher.Flush()
buf.ReleaseMulti(releaseBuffers)
}
}
@ -229,6 +234,9 @@ func (m *Mixed) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
if destination == m.broadcastAddr || !destination.IsGlobalUnicast() {
return
}
if m.dispatchIPv4(ipHdr, destination) {
return false, nil
}
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
writeBack, err = m.processIPv4TCP(ipHdr, ipHdr.Payload())
@ -249,9 +257,13 @@ func (m *Mixed) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
func (m *Mixed) processIPv6(ipHdr header.IPv6) (writeBack bool, err error) {
writeBack = true
if !ipHdr.DestinationAddr().IsGlobalUnicast() {
destination := ipHdr.DestinationAddr()
if !destination.IsGlobalUnicast() {
return
}
if m.dispatchIPv6(ipHdr, destination) {
return false, nil
}
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
writeBack, err = m.processIPv6TCP(ipHdr, ipHdr.Payload())

View file

@ -5,6 +5,7 @@ import (
"errors"
"net"
"net/netip"
"slices"
"syscall"
"time"
@ -46,7 +47,7 @@ type System struct {
tcpPort6 uint16
tcpNat *TCPNat
udpNat *udpnat.Service
directNat *DirectRouteMapping
dispatcher *ForwardDispatcher
bindInterface bool
interfaceFinder control.InterfaceFinder
frontHeadroom int
@ -101,6 +102,7 @@ func NewSystem(options StackOptions) (Stack, error) {
}
func (s *System) Close() error {
s.dispatcher.Close()
return common.Close(
s.tcpListener,
s.tcpListener6,
@ -162,7 +164,13 @@ func (s *System) start() error {
}
s.tcpNat = NewNat(s.ctx, s.udpTimeout)
s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false)
s.directNat = NewDirectRouteMapping(s.icmpTimeout)
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
s.frontHeadroom = linuxTUN.FrontHeadroom()
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
}
if s.handler != nil {
s.dispatcher = NewForwardDispatcher(s.handler, newSystemWriteback(s.tun, s.frontHeadroom), s.logger, s.udpTimeout, s.icmpTimeout)
}
return nil
}
@ -172,8 +180,6 @@ func (s *System) tunLoop() {
return
}
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
s.frontHeadroom = linuxTUN.FrontHeadroom()
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
batchSize := linuxTUN.BatchSize()
if batchSize > 1 {
s.batchLoopLinux(linuxTUN, batchSize)
@ -204,6 +210,7 @@ func (s *System) tunLoop() {
s.logger.Trace(E.Cause(err, "write packet"))
}
}
s.dispatcher.Flush()
}
}
@ -223,6 +230,7 @@ func (s *System) wintunLoop(winTun WinTun) {
s.logger.Trace(E.Cause(err, "write packet"))
}
}
s.dispatcher.Flush()
release()
}
}
@ -263,11 +271,13 @@ func (s *System) batchLoopLinux(linuxTUN LinuxTUN, batchSize int) {
}
writeBuffers = writeBuffers[:0]
}
s.dispatcher.Flush()
}
}
func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
var writeBuffers []*buf.Buffer
var releaseBuffers []*buf.Buffer
for {
buffers, err := darwinTUN.BatchRead()
if err != nil {
@ -280,6 +290,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
continue
}
writeBuffers = writeBuffers[:0]
releaseBuffers = releaseBuffers[:0]
for _, buffer := range buffers {
packetSize := buffer.Len()
if packetSize < header.IPv4MinimumSize {
@ -289,7 +300,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
if s.processPacket(buffer.Bytes()) {
writeBuffers = append(writeBuffers, buffer)
} else {
buffer.Release()
releaseBuffers = append(releaseBuffers, buffer)
}
}
if len(writeBuffers) > 0 {
@ -299,6 +310,8 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
}
buf.ReleaseMulti(writeBuffers)
}
s.dispatcher.Flush()
buf.ReleaseMulti(releaseBuffers)
}
}
@ -338,11 +351,53 @@ func (s *System) acceptLoop(listener net.Listener) {
}
}
func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool {
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
if slices.Contains(s.inet4LoopbackAddress, destination) {
return false
}
if ipHdr.SourceAddr() == s.inet4Address &&
ipHdr.FragmentOffset() == 0 &&
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort {
return false
}
case header.ICMPv4ProtocolNumber:
if destination == s.inet4Address {
return false
}
}
return s.dispatcher.Dispatch(ipHdr)
}
func (s *System) dispatchIPv6(ipHdr header.IPv6, destination netip.Addr) bool {
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
if slices.Contains(s.inet6LoopbackAddress, destination) {
return false
}
if ipHdr.SourceAddr() == s.inet6Address &&
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort6 {
return false
}
case header.ICMPv6ProtocolNumber:
if destination == s.inet6Address {
return false
}
}
return s.dispatcher.Dispatch(ipHdr)
}
func (s *System) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
destination := ipHdr.DestinationAddr()
if destination == s.broadcastAddr || !destination.IsGlobalUnicast() {
return
}
if s.dispatchIPv4(ipHdr, destination) {
return false, nil
}
writeBack = true
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
@ -360,9 +415,13 @@ func (s *System) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
}
func (s *System) processIPv6(ipHdr header.IPv6) (writeBack bool, err error) {
if !ipHdr.DestinationAddr().IsGlobalUnicast() {
destination := ipHdr.DestinationAddr()
if !destination.IsGlobalUnicast() {
return
}
if s.dispatchIPv6(ipHdr, destination) {
return false, nil
}
writeBack = true
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
@ -404,14 +463,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
}
}
if !loopback {
natPort, err := s.tcpNat.Lookup(source, destination, s.handler)
if err != nil {
if errors.Is(err, ErrDrop) {
return false, nil
} else {
return false, s.resetIPv4TCP(ipHdr, tcpHdr)
}
}
natPort := s.tcpNat.Lookup(source, destination)
ipHdr.SetSourceAddr(s.inet4NextAddress)
tcpHdr.SetSourcePort(natPort)
ipHdr.SetDestinationAddr(s.inet4Address)
@ -429,51 +481,6 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
return true, nil
}
func (s *System) resetIPv4TCP(origIPHdr header.IPv4, origTCPHdr header.TCP) error {
frontHeadroom := s.frontHeadroom + PacketOffset
newPacket := buf.NewSize(frontHeadroom + header.IPv4MinimumSize + header.TCPMinimumSize)
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv4MinimumSize+header.TCPMinimumSize)
ipHdr := header.IPv4(newPacket.Bytes())
ipHdr.Encode(&header.IPv4Fields{
TotalLength: uint16(newPacket.Len()),
Protocol: uint8(header.TCPProtocolNumber),
SrcAddr: origIPHdr.DestinationAddr(),
DstAddr: origIPHdr.SourceAddr(),
})
tcpHdr := header.TCP(ipHdr.Payload())
fields := header.TCPFields{
SrcPort: origTCPHdr.DestinationPort(),
DstPort: origTCPHdr.SourcePort(),
DataOffset: header.TCPMinimumSize,
Flags: header.TCPFlagRst,
}
if origTCPHdr.Flags()&header.TCPFlagAck != 0 {
fields.SeqNum = origTCPHdr.AckNumber()
} else {
fields.Flags |= header.TCPFlagAck
ackNum := origTCPHdr.SequenceNumber() + uint32(len(origTCPHdr.Payload()))
if origTCPHdr.Flags()&header.TCPFlagSyn != 0 {
ackNum++
}
if origTCPHdr.Flags()&header.TCPFlagFin != 0 {
ackNum++
}
fields.AckNum = ackNum
}
tcpHdr.Encode(&fields)
if !s.txChecksumOffload {
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
}
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, error) {
source := netip.AddrPortFrom(ipHdr.SourceAddr(), tcpHdr.SourcePort())
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort())
@ -499,14 +506,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
}
}
if !loopback {
natPort, err := s.tcpNat.Lookup(source, destination, s.handler)
if err != nil {
if errors.Is(err, ErrDrop) {
return false, nil
} else {
return false, s.resetIPv6TCP(ipHdr, tcpHdr)
}
}
natPort := s.tcpNat.Lookup(source, destination)
ipHdr.SetSourceAddr(s.inet6NextAddress)
tcpHdr.SetSourcePort(natPort)
ipHdr.SetDestinationAddr(s.inet6Address)
@ -523,50 +523,6 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
return true, nil
}
func (s *System) resetIPv6TCP(origIPHdr header.IPv6, origTCPHdr header.TCP) error {
frontHeadroom := s.frontHeadroom + PacketOffset
newPacket := buf.NewSize(frontHeadroom + header.IPv6MinimumSize + header.TCPMinimumSize)
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv6MinimumSize+header.TCPMinimumSize)
ipHdr := header.IPv6(newPacket.Bytes())
ipHdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(header.TCPMinimumSize),
TransportProtocol: header.TCPProtocolNumber,
SrcAddr: origIPHdr.DestinationAddr(),
DstAddr: origIPHdr.SourceAddr(),
})
tcpHdr := header.TCP(ipHdr.Payload())
fields := header.TCPFields{
SrcPort: origTCPHdr.DestinationPort(),
DstPort: origTCPHdr.SourcePort(),
DataOffset: header.TCPMinimumSize,
Flags: header.TCPFlagRst,
}
if origTCPHdr.Flags()&header.TCPFlagAck != 0 {
fields.SeqNum = origTCPHdr.AckNumber()
} else {
fields.Flags |= header.TCPFlagAck
ackNum := origTCPHdr.SequenceNumber() + uint32(len(origTCPHdr.Payload()))
if origTCPHdr.Flags()&header.TCPFlagSyn != 0 {
ackNum++
}
if origTCPHdr.Flags()&header.TCPFlagFin != 0 {
ackNum++
}
fields.AckNum = ackNum
}
tcpHdr.Encode(&fields)
if !s.txChecksumOffload {
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
}
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
func (s *System) processIPv4UDP(ipHdr header.IPv4, udpHdr header.UDP) error {
if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 {
return E.New("ipv4: fragment dropped")
@ -594,19 +550,6 @@ func (s *System) processIPv6UDP(ipHdr header.IPv6, udpHdr header.UDP) error {
}
func (s *System) preparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) {
_, pErr := s.handler.PrepareConnection(N.NetworkUDP, source, destination, nil, 0)
if pErr != nil {
if !errors.Is(pErr, ErrDrop) {
if source.IsIPv4() {
ipHdr := userData.(header.IPv4)
s.rejectIPv4WithICMP(ipHdr, header.ICMPv4PortUnreachable)
} else {
ipHdr := userData.(header.IPv6)
s.rejectIPv6WithICMP(ipHdr, header.ICMPv6PortUnreachable)
}
}
return false, nil, nil, nil
}
var writer N.PacketWriter
if source.IsIPv4() {
packet := userData.(header.IPv4)
@ -640,29 +583,6 @@ func (s *System) processIPv4ICMP(ipHdr header.IPv4, icmpHdr header.ICMPv4) (bool
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
return false, nil
}
sourceAddr := ipHdr.SourceAddr()
destinationAddr := ipHdr.DestinationAddr()
if destinationAddr != s.inet4Address {
action, err := s.directNat.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
return s.handler.PrepareConnection(
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&systemICMPDirectPacketWriter4{s.tun, s.frontHeadroom + PacketOffset, sourceAddr},
timeout,
)
})
if err != nil {
if errors.Is(err, ErrReset) {
return false, s.rejectIPv4WithICMP(ipHdr, header.ICMPv4HostUnreachable)
} else if errors.Is(err, ErrDrop) {
return false, nil
}
}
if action != nil {
return false, action.WritePacket(buf.As(ipHdr).ToOwned())
}
}
icmpHdr.SetType(header.ICMPv4EchoReply)
sourceAddress := ipHdr.SourceAddr()
ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
@ -672,70 +592,10 @@ func (s *System) processIPv4ICMP(ipHdr header.IPv4, icmpHdr header.ICMPv4) (bool
return true, nil
}
func (s *System) rejectIPv4WithICMP(ipHdr header.IPv4, code header.ICMPv4Code) error {
frontHeadroom := s.frontHeadroom + PacketOffset
mtu := s.mtu
const maxIPData = header.IPv4MinimumProcessableDatagramSize - header.IPv4MinimumSize
if mtu > maxIPData {
mtu = maxIPData
}
available := mtu - header.ICMPv4MinimumSize
if available < len(ipHdr)+header.ICMPv4MinimumErrorPayloadSize {
return nil
}
payload := ipHdr
if len(payload) > available {
payload = payload[:available]
}
newPacket := buf.NewSize(frontHeadroom + header.IPv4MinimumSize + header.ICMPv4MinimumSize + len(payload))
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv4MinimumSize+header.ICMPv4MinimumSize+len(payload))
newIPHdr := header.IPv4(newPacket.Bytes())
newIPHdr.Encode(&header.IPv4Fields{
TotalLength: uint16(newPacket.Len()),
Protocol: uint8(header.ICMPv4ProtocolNumber),
SrcAddr: ipHdr.DestinationAddr(),
DstAddr: ipHdr.SourceAddr(),
})
newIPHdr.SetChecksum(^newIPHdr.CalculateChecksum())
icmpHdr := header.ICMPv4(newIPHdr.Payload())
icmpHdr.SetType(header.ICMPv4DstUnreachable)
icmpHdr.SetCode(code)
icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr[:header.ICMPv4MinimumSize], checksum.Checksum(ipHdr.Payload(), 0)))
copy(icmpHdr.Payload(), payload)
if PacketOffset > 0 {
newPacket.ExtendHeader(PacketOffset)[3] = syscall.AF_INET
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool, error) {
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return false, nil
}
sourceAddr := ipHdr.SourceAddr()
destinationAddr := ipHdr.DestinationAddr()
if destinationAddr != s.inet6Address {
action, err := s.directNat.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
return s.handler.PrepareConnection(
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&systemICMPDirectPacketWriter6{s.tun, s.frontHeadroom + PacketOffset, sourceAddr},
timeout,
)
})
if errors.Is(err, ErrReset) {
return false, s.rejectIPv6WithICMP(ipHdr, header.ICMPv6AddressUnreachable)
} else if errors.Is(err, ErrDrop) {
return false, nil
}
if action != nil {
return false, action.WritePacket(buf.As(ipHdr).ToOwned())
}
}
icmpHdr.SetType(header.ICMPv6EchoReply)
sourceAddress := ipHdr.SourceAddr()
ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
@ -748,50 +608,6 @@ func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool
return true, nil
}
func (s *System) rejectIPv6WithICMP(ipHdr header.IPv6, code header.ICMPv6Code) error {
frontHeadroom := s.frontHeadroom + PacketOffset
mtu := s.mtu
const maxIPv6Data = header.IPv6MinimumMTU - header.IPv6FixedHeaderSize
if mtu > maxIPv6Data {
mtu = maxIPv6Data
}
available := mtu - header.ICMPv6ErrorHeaderSize
if available < header.IPv6MinimumSize {
return nil
}
payload := ipHdr
if len(payload) > available {
payload = payload[:available]
}
newPacket := buf.NewSize(frontHeadroom + header.IPv6MinimumSize + header.ICMPv6DstUnreachableMinimumSize + len(payload))
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv6MinimumSize+header.ICMPv6DstUnreachableMinimumSize+len(payload))
newIPHdr := header.IPv6(newPacket.Bytes())
newIPHdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(header.ICMPv6DstUnreachableMinimumSize + len(payload)),
TransportProtocol: header.ICMPv6ProtocolNumber,
SrcAddr: ipHdr.DestinationAddr(),
DstAddr: ipHdr.SourceAddr(),
})
icmpHdr := header.ICMPv6(newIPHdr.Payload())
icmpHdr.SetType(header.ICMPv6DstUnreachable)
icmpHdr.SetCode(code)
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr[:header.ICMPv6DstUnreachableMinimumSize],
Src: newIPHdr.SourceAddressSlice(),
Dst: newIPHdr.DestinationAddressSlice(),
PayloadCsum: checksum.Checksum(payload, 0),
PayloadLen: len(payload),
}))
copy(icmpHdr.Payload(), payload)
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
type systemUDPPacketWriter4 struct {
tun Tun
frontHeadroom int
@ -868,45 +684,37 @@ func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.S
return common.Error(w.tun.Write(newPacket.Bytes()))
}
type systemICMPDirectPacketWriter4 struct {
type systemWriteback struct {
tun Tun
linuxTUN LinuxTUN
frontHeadroom int
source netip.Addr
}
func (w *systemICMPDirectPacketWriter4) WritePacket(p []byte) error {
newPacket := buf.NewSize(w.frontHeadroom + len(p))
defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(p)
ipHdr := header.IPv4(newPacket.Bytes())
ipHdr.SetDestinationAddr(w.source)
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
} else {
newPacket.Advance(-w.frontHeadroom)
func newSystemWriteback(tunInterface Tun, frontHeadroom int) *systemWriteback {
writeback := &systemWriteback{tun: tunInterface, frontHeadroom: frontHeadroom}
if linuxTUN, isLinuxTUN := tunInterface.(LinuxTUN); isLinuxTUN {
writeback.linuxTUN = linuxTUN
}
return common.Error(w.tun.Write(newPacket.Bytes()))
return writeback
}
type systemICMPDirectPacketWriter6 struct {
tun Tun
frontHeadroom int
source netip.Addr
func (w *systemWriteback) ReturnHeadroom() int {
return w.frontHeadroom + PacketOffset
}
func (w *systemICMPDirectPacketWriter6) WritePacket(p []byte) error {
newPacket := buf.NewSize(w.frontHeadroom + len(p))
defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(p)
ipHdr := header.IPv6(newPacket.Bytes())
ipHdr.SetDestinationAddr(w.source)
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} else {
newPacket.Advance(-w.frontHeadroom)
func (w *systemWriteback) WriteReturnPackets(packets [][]byte) error {
if w.linuxTUN != nil {
return common.Error(w.linuxTUN.BatchWrite(packets, w.frontHeadroom))
}
return common.Error(w.tun.Write(newPacket.Bytes()))
var writeErrors []error
for _, packet := range packets {
if PacketOffset > 0 {
PacketFillHeader(packet, header.IPVersion(packet[PacketOffset:]))
}
_, err := w.tun.Write(packet)
if err != nil {
writeErrors = append(writeErrors, err)
}
}
return E.Errors(writeErrors...)
}

View file

@ -5,9 +5,6 @@ import (
"net/netip"
"sync"
"time"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
type TCPNat struct {
@ -85,17 +82,13 @@ func (n *TCPNat) LookupBack(port uint16) *TCPSession {
return session
}
func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort, handler Handler) (uint16, error) {
func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort) uint16 {
key := tcpNatKey{Source: source, Destination: destination}
n.addrAccess.RLock()
port, loaded := n.addrMap[key]
n.addrAccess.RUnlock()
if loaded {
return port, nil
}
_, pErr := handler.PrepareConnection(N.NetworkTCP, M.SocksaddrFromNetIP(source), M.SocksaddrFromNetIP(destination), nil, 0)
if pErr != nil {
return 0, pErr
return port
}
n.addrAccess.Lock()
nextPort := n.portIndex
@ -114,5 +107,5 @@ func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort, handl
LastActive: time.Now(),
}
n.portAccess.Unlock()
return nextPort, nil
return nextPort
}

14
tun.go
View file

@ -7,7 +7,6 @@ import (
"runtime"
"strconv"
"strings"
"time"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
@ -15,27 +14,16 @@ import (
E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/ranges"
)
type Handler interface {
PrepareConnection(
network string,
source M.Socksaddr,
destination M.Socksaddr,
routeContext DirectRouteContext,
timeout time.Duration,
) (DirectRouteDestination, error)
JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort) FlowVerdict
N.TCPConnectionHandlerEx
N.UDPConnectionHandlerEx
}
type DirectRouteContext interface {
WritePacket(packet []byte) error
}
type Tun interface {
io.ReadWriter
Name() (string, error)

View file

@ -568,8 +568,10 @@ func (t *NativeTun) readNonblocking(buffer []byte) (int, error) {
func (t *NativeTun) BatchWrite(buffers [][]byte, offset int) (int, error) {
t.writeAccess.Lock()
defer func() {
if t.vnetHdr {
t.tcpGROTable.reset()
t.udpGROTable.reset()
}
t.writeAccess.Unlock()
}()
var (