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

View file

@ -9,7 +9,6 @@ import (
"sync" "sync"
"time" "time"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing-tun/gtcpip/header" "github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/buf" "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control" "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. // 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 const maxICMPPacketSize = 2048
var _ tun.DirectRouteDestination = (*Destination)(nil) type PacketWriter interface {
WritePacket(packet []byte) error
}
type Destination struct { type Destination struct {
conn *Conn conn *Conn
ctx context.Context ctx context.Context
logger logger.ContextLogger logger logger.ContextLogger
destination netip.Addr destination netip.Addr
routeContext tun.DirectRouteContext writer PacketWriter
timeout time.Duration timeout time.Duration
requestAccess sync.Mutex requestAccess sync.Mutex
requests map[pingRequest]time.Time requests map[pingRequest]time.Time
@ -45,9 +46,9 @@ func ConnectDestination(
logger logger.ContextLogger, logger logger.ContextLogger,
controlFunc control.Func, controlFunc control.Func,
destination netip.Addr, destination netip.Addr,
routeContext tun.DirectRouteContext, writer PacketWriter,
timeout time.Duration, timeout time.Duration,
) (tun.DirectRouteDestination, error) { ) (*Destination, error) {
var ( var (
conn *Conn conn *Conn
err error err error
@ -69,7 +70,7 @@ func ConnectDestination(
ctx: ctx, ctx: ctx,
logger: logger, logger: logger,
destination: destination, destination: destination,
routeContext: routeContext, writer: writer,
timeout: timeout, timeout: timeout,
requests: make(map[pingRequest]time.Time), 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()) 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 { if err != nil {
d.logger.ErrorContext(d.ctx, E.Cause(err, "write ICMP echo reply")) 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 r.redirectServer = server
} }
if r.useNFTables { if r.useNFTables {
if r.handler != nil {
var handler *nfqueueHandler var handler *nfqueueHandler
handler, err = newNFQueueHandler(nfqueueOptions{ handler, err = newNFQueueHandler(nfqueueOptions{
Context: r.ctx, Context: r.ctx,
@ -156,6 +157,7 @@ func (r *autoRedirect) Start() error {
r.nfqueueHandler = handler r.nfqueueHandler = handler
r.nfqueueEnabled = true r.nfqueueEnabled = true
} }
}
r.cleanupNFTables() r.cleanupNFTables()
err = r.setupNFTables() err = r.setupNFTables()
if err != nil { 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" "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 { type Stack interface {
Start() error Start() error
Close() error Close() error

View file

@ -6,8 +6,10 @@ import (
"context" "context"
"net/netip" "net/netip"
"runtime" "runtime"
"sync"
"time" "time"
"github.com/sagernet/gvisor/pkg/buffer"
"github.com/sagernet/gvisor/pkg/tcpip" "github.com/sagernet/gvisor/pkg/tcpip"
"github.com/sagernet/gvisor/pkg/tcpip/adapters/gonet" "github.com/sagernet/gvisor/pkg/tcpip/adapters/gonet"
"github.com/sagernet/gvisor/pkg/tcpip/header" "github.com/sagernet/gvisor/pkg/tcpip/header"
@ -40,6 +42,8 @@ type GVisor struct {
logger logger.Logger logger logger.Logger
stack *stack.Stack stack *stack.Stack
endpoint stack.LinkEndpoint endpoint stack.LinkEndpoint
dispatcher *ForwardDispatcher
icmpForwarder *ICMPForwarder
} }
type GVisorTun interface { type GVisorTun interface {
@ -88,23 +92,39 @@ func (t *GVisor) Start() error {
if err != nil { if err != nil {
return err 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) ipStack, err := newGVisorStack(linkEndpoint, nicOptions, false, true)
if err != nil { if err != nil {
return err return err
} }
ipStack.SetTransportProtocolHandler(tcp.ProtocolNumber, NewTCPForwarderWithLoopback(t.ctx, ipStack, t.handler, t.inet4LoopbackAddress, t.inet6LoopbackAddress, t.tun).HandlePacket) 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) 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 := NewICMPForwarder(ipStack, t.handler, t.logger)
icmpForwarder.SetLocalAddresses(t.inet4Address, t.inet6Address)
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket)
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket)
t.icmpForwarder = icmpForwarder
t.stack = ipStack t.stack = ipStack
t.endpoint = linkEndpoint t.endpoint = linkEndpoint
return nil return nil
} }
func (t *GVisor) Close() error { func (t *GVisor) Close() error {
t.dispatcher.Close()
if t.icmpForwarder != nil {
t.icmpForwarder.Close()
}
if t.stack == nil { if t.stack == nil {
return nil return nil
} }
@ -116,6 +136,37 @@ func (t *GVisor) Close() error {
return nil 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 { func AddressFromAddr(destination netip.Addr) tcpip.Address {
if destination.Is6() { if destination.Is6() {
return tcpip.AddrFrom16(destination.As16()) return tcpip.AddrFrom16(destination.As16())

View file

@ -16,10 +16,24 @@ type LinkEndpointFilter struct {
stack.LinkEndpoint stack.LinkEndpoint
BroadcastAddress netip.Addr BroadcastAddress netip.Addr
Writer GVisorTun Writer GVisorTun
Dispatcher *ForwardDispatcher
Inet4Address netip.Addr
Inet6Address netip.Addr
Inet4LoopbackAddress []netip.Addr
Inet6LoopbackAddress []netip.Addr
} }
func (w *LinkEndpointFilter) Attach(dispatcher stack.NetworkDispatcher) { 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) var _ stack.NetworkDispatcher = (*networkDispatcherFilter)(nil)
@ -28,6 +42,11 @@ type networkDispatcherFilter struct {
stack.NetworkDispatcher stack.NetworkDispatcher
broadcastAddress netip.Addr broadcastAddress netip.Addr
writer GVisorTun 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) { 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) w.writer.WritePacket(pkt)
return 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) 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 package tun
import ( import (
"context"
"errors"
"net/netip" "net/netip"
"sync" "sync"
"sync/atomic"
"time" "time"
"github.com/sagernet/gvisor/pkg/buffer" "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/ipv4"
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv6" "github.com/sagernet/gvisor/pkg/tcpip/network/ipv6"
"github.com/sagernet/gvisor/pkg/tcpip/stack" "github.com/sagernet/gvisor/pkg/tcpip/stack"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions" E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger" "github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
) )
type ICMPForwarder struct { type ICMPForwarder struct {
ctx context.Context
stack *stack.Stack stack *stack.Stack
logger logger.Logger
inet4Address netip.Addr
inet6Address netip.Addr
handler Handler 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( type icmpFlowKey struct {
ctx context.Context, v6 bool
stack *stack.Stack, source netip.Addr
logger logger.Logger, destination netip.Addr
handler Handler, identifier uint16
timeout time.Duration, }
) *ICMPForwarder {
return &ICMPForwarder{ func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) *ICMPForwarder {
ctx: ctx, forwarder := &ICMPForwarder{
stack: stack, stack: stack,
logger: logger,
handler: handler, 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) { func (f *ICMPForwarder) Close() error {
f.inet4Address = inet4Address f.returnPath.closed.Store(true)
f.inet6Address = inet6Address 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 { 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 { if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
return false return false
} }
sourceAddr := M.AddrFromIP(ipHdr.SourceAddressSlice()) identifier := icmpHdr.Ident()
destinationAddr := M.AddrFromIP(ipHdr.DestinationAddressSlice()) verdict := f.handler.JudgeFlow(
if destinationAddr != f.inet4Address { uint8(header.ICMPv4ProtocolNumber),
action, err := f.mapping.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) { netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
return f.handler.PrepareConnection( netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&ICMPBackWriter{
stack: f.stack,
packet: pkt,
source: ipHdr.SourceAddress(),
sourceNetwork: header.IPv4ProtocolNumber,
},
timeout,
) )
}) switch verdict.Action {
if errors.Is(err, ErrReset) { case ActionReject, ActionDrop:
gWriteUnreachable(f.stack, pkt)
return true return true
} else if errors.Is(err, ErrDrop) { case ActionFlow:
return true if f.forwardFlow(verdict.Port, false, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
}
if action != nil {
err = icmpWritePacketBuffer(action, pkt)
if err != nil {
f.logger.Error(E.Cause(err, "write ICMPv4 echo request"))
}
return true return true
} }
} }
@ -125,35 +116,17 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 { if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return false return false
} }
sourceAddr := M.AddrFromIP(ipHdr.SourceAddressSlice()) identifier := icmpHdr.Ident()
destinationAddr := M.AddrFromIP(ipHdr.DestinationAddressSlice()) verdict := f.handler.JudgeFlow(
if destinationAddr != f.inet6Address { uint8(header.ICMPv6ProtocolNumber),
action, err := f.mapping.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) { netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
return f.handler.PrepareConnection( netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&ICMPBackWriter{
stack: f.stack,
packet: pkt,
source: ipHdr.SourceAddress(),
sourceNetwork: header.IPv6ProtocolNumber,
},
timeout,
) )
}) switch verdict.Action {
if errors.Is(err, ErrReset) { case ActionReject, ActionDrop:
gWriteUnreachable(f.stack, pkt)
return true return true
} else if errors.Is(err, ErrDrop) { case ActionFlow:
return true if f.forwardFlow(verdict.Port, true, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
}
if action != nil {
pkt.IncRef()
err = icmpWritePacketBuffer(action, pkt)
if err != nil {
f.logger.Error(E.Cause(err, "write ICMPv6 echo request"))
}
return true return true
} }
} }
@ -190,64 +163,179 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
} }
} }
type ICMPBackWriter struct { func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, destination netip.Addr, identifier uint16, pkt *stack.PacketBuffer) bool {
access sync.Mutex if port == nil {
stack *stack.Stack return false
packet *stack.PacketBuffer }
source tcpip.Address inet4Address, inet6Address := port.PortAddresses()
sourceNetwork tcpip.NetworkProtocolNumber 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 { func (f *ICMPForwarder) lookupFlow(key icmpFlowKey) bool {
if w.sourceNetwork == header.IPv4ProtocolNumber { f.flowAccess.Lock()
route, err := w.stack.FindRoute( defer f.flowAccess.Unlock()
DefaultNIC, deadline, loaded := f.flows[key]
header.IPv4(p).SourceAddress(), if !loaded {
w.source, return false
w.sourceNetwork, }
false, now := time.Now()
) if now.After(deadline) {
if err != nil { delete(f.flows, key)
return gonet.TranslateNetstackError(err) 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() defer route.Release()
packet := stack.NewPacketBuffer(stack.PacketBufferOptions{ packetBuffer := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(p), Payload: buffer.MakeWithData(packet),
}) })
defer packet.DecRef() defer packetBuffer.DecRef()
parse.IPv4(packet) if protocol == header.IPv4ProtocolNumber {
err = route.WritePacketDirect(packet) parse.IPv4(packetBuffer)
if err != nil {
return gonet.TranslateNetstackError(err)
}
} else { } else {
route, err := w.stack.FindRoute( parse.IPv6(packetBuffer)
DefaultNIC,
header.IPv6(p).SourceAddress(),
w.source,
w.sourceNetwork,
false,
)
if err != nil {
return gonet.TranslateNetstackError(err)
} }
defer route.Release() gErr = route.WritePacketDirect(packetBuffer)
packet := stack.NewPacketBuffer(stack.PacketBufferOptions{ if gErr != nil {
Payload: buffer.MakeWithData(p), f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "write ICMP reply"))
})
parse.IPv6(packet)
defer packet.DecRef()
err = route.WritePacketDirect(packet)
if err != nil {
return gonet.TranslateNetstackError(err)
} }
} return true
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())
} }

View file

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

View file

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

View file

@ -4,7 +4,6 @@ package tun
import ( import (
"context" "context"
"errors"
"math" "math"
"net/netip" "net/netip"
"os" "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 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) { 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) switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action {
if pErr != nil { case ActionReject:
if !errors.Is(pErr, ErrDrop) {
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer)) gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
} return false, nil, nil, nil
case ActionDrop:
return false, nil, nil, nil return false, nil, nil, nil
} }
var sourceNetwork tcpip.NetworkProtocolNumber var sourceNetwork tcpip.NetworkProtocolNumber

View file

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

View file

@ -5,6 +5,7 @@ import (
"errors" "errors"
"net" "net"
"net/netip" "net/netip"
"slices"
"syscall" "syscall"
"time" "time"
@ -46,7 +47,7 @@ type System struct {
tcpPort6 uint16 tcpPort6 uint16
tcpNat *TCPNat tcpNat *TCPNat
udpNat *udpnat.Service udpNat *udpnat.Service
directNat *DirectRouteMapping dispatcher *ForwardDispatcher
bindInterface bool bindInterface bool
interfaceFinder control.InterfaceFinder interfaceFinder control.InterfaceFinder
frontHeadroom int frontHeadroom int
@ -101,6 +102,7 @@ func NewSystem(options StackOptions) (Stack, error) {
} }
func (s *System) Close() error { func (s *System) Close() error {
s.dispatcher.Close()
return common.Close( return common.Close(
s.tcpListener, s.tcpListener,
s.tcpListener6, s.tcpListener6,
@ -162,7 +164,13 @@ func (s *System) start() error {
} }
s.tcpNat = NewNat(s.ctx, s.udpTimeout) s.tcpNat = NewNat(s.ctx, s.udpTimeout)
s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false) 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 return nil
} }
@ -172,8 +180,6 @@ func (s *System) tunLoop() {
return return
} }
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN { if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
s.frontHeadroom = linuxTUN.FrontHeadroom()
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
batchSize := linuxTUN.BatchSize() batchSize := linuxTUN.BatchSize()
if batchSize > 1 { if batchSize > 1 {
s.batchLoopLinux(linuxTUN, batchSize) s.batchLoopLinux(linuxTUN, batchSize)
@ -204,6 +210,7 @@ func (s *System) tunLoop() {
s.logger.Trace(E.Cause(err, "write packet")) 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.logger.Trace(E.Cause(err, "write packet"))
} }
} }
s.dispatcher.Flush()
release() release()
} }
} }
@ -263,11 +271,13 @@ func (s *System) batchLoopLinux(linuxTUN LinuxTUN, batchSize int) {
} }
writeBuffers = writeBuffers[:0] writeBuffers = writeBuffers[:0]
} }
s.dispatcher.Flush()
} }
} }
func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) { func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
var writeBuffers []*buf.Buffer var writeBuffers []*buf.Buffer
var releaseBuffers []*buf.Buffer
for { for {
buffers, err := darwinTUN.BatchRead() buffers, err := darwinTUN.BatchRead()
if err != nil { if err != nil {
@ -280,6 +290,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
continue continue
} }
writeBuffers = writeBuffers[:0] writeBuffers = writeBuffers[:0]
releaseBuffers = releaseBuffers[:0]
for _, buffer := range buffers { for _, buffer := range buffers {
packetSize := buffer.Len() packetSize := buffer.Len()
if packetSize < header.IPv4MinimumSize { if packetSize < header.IPv4MinimumSize {
@ -289,7 +300,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
if s.processPacket(buffer.Bytes()) { if s.processPacket(buffer.Bytes()) {
writeBuffers = append(writeBuffers, buffer) writeBuffers = append(writeBuffers, buffer)
} else { } else {
buffer.Release() releaseBuffers = append(releaseBuffers, buffer)
} }
} }
if len(writeBuffers) > 0 { if len(writeBuffers) > 0 {
@ -299,6 +310,8 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
} }
buf.ReleaseMulti(writeBuffers) 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) { func (s *System) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
destination := ipHdr.DestinationAddr() destination := ipHdr.DestinationAddr()
if destination == s.broadcastAddr || !destination.IsGlobalUnicast() { if destination == s.broadcastAddr || !destination.IsGlobalUnicast() {
return return
} }
if s.dispatchIPv4(ipHdr, destination) {
return false, nil
}
writeBack = true writeBack = true
switch ipHdr.TransportProtocol() { switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber: 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) { func (s *System) processIPv6(ipHdr header.IPv6) (writeBack bool, err error) {
if !ipHdr.DestinationAddr().IsGlobalUnicast() { destination := ipHdr.DestinationAddr()
if !destination.IsGlobalUnicast() {
return return
} }
if s.dispatchIPv6(ipHdr, destination) {
return false, nil
}
writeBack = true writeBack = true
switch ipHdr.TransportProtocol() { switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber: case header.TCPProtocolNumber:
@ -404,14 +463,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
} }
} }
if !loopback { if !loopback {
natPort, err := s.tcpNat.Lookup(source, destination, s.handler) natPort := s.tcpNat.Lookup(source, destination)
if err != nil {
if errors.Is(err, ErrDrop) {
return false, nil
} else {
return false, s.resetIPv4TCP(ipHdr, tcpHdr)
}
}
ipHdr.SetSourceAddr(s.inet4NextAddress) ipHdr.SetSourceAddr(s.inet4NextAddress)
tcpHdr.SetSourcePort(natPort) tcpHdr.SetSourcePort(natPort)
ipHdr.SetDestinationAddr(s.inet4Address) ipHdr.SetDestinationAddr(s.inet4Address)
@ -429,51 +481,6 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
return true, nil 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) { func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, error) {
source := netip.AddrPortFrom(ipHdr.SourceAddr(), tcpHdr.SourcePort()) source := netip.AddrPortFrom(ipHdr.SourceAddr(), tcpHdr.SourcePort())
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort()) 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 { if !loopback {
natPort, err := s.tcpNat.Lookup(source, destination, s.handler) natPort := s.tcpNat.Lookup(source, destination)
if err != nil {
if errors.Is(err, ErrDrop) {
return false, nil
} else {
return false, s.resetIPv6TCP(ipHdr, tcpHdr)
}
}
ipHdr.SetSourceAddr(s.inet6NextAddress) ipHdr.SetSourceAddr(s.inet6NextAddress)
tcpHdr.SetSourcePort(natPort) tcpHdr.SetSourcePort(natPort)
ipHdr.SetDestinationAddr(s.inet6Address) ipHdr.SetDestinationAddr(s.inet6Address)
@ -523,50 +523,6 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
return true, nil 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 { func (s *System) processIPv4UDP(ipHdr header.IPv4, udpHdr header.UDP) error {
if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 { if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 {
return E.New("ipv4: fragment dropped") 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) { 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 var writer N.PacketWriter
if source.IsIPv4() { if source.IsIPv4() {
packet := userData.(header.IPv4) 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 { if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
return false, nil 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) icmpHdr.SetType(header.ICMPv4EchoReply)
sourceAddress := ipHdr.SourceAddr() sourceAddress := ipHdr.SourceAddr()
ipHdr.SetSourceAddr(ipHdr.DestinationAddr()) ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
@ -672,70 +592,10 @@ func (s *System) processIPv4ICMP(ipHdr header.IPv4, icmpHdr header.ICMPv4) (bool
return true, nil 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) { func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool, error) {
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 { if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return false, nil 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) icmpHdr.SetType(header.ICMPv6EchoReply)
sourceAddress := ipHdr.SourceAddr() sourceAddress := ipHdr.SourceAddr()
ipHdr.SetSourceAddr(ipHdr.DestinationAddr()) ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
@ -748,50 +608,6 @@ func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool
return true, nil 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 { type systemUDPPacketWriter4 struct {
tun Tun tun Tun
frontHeadroom int frontHeadroom int
@ -868,45 +684,37 @@ func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.S
return common.Error(w.tun.Write(newPacket.Bytes())) return common.Error(w.tun.Write(newPacket.Bytes()))
} }
type systemICMPDirectPacketWriter4 struct { type systemWriteback struct {
tun Tun tun Tun
linuxTUN LinuxTUN
frontHeadroom int frontHeadroom int
source netip.Addr
} }
func (w *systemICMPDirectPacketWriter4) WritePacket(p []byte) error { func newSystemWriteback(tunInterface Tun, frontHeadroom int) *systemWriteback {
newPacket := buf.NewSize(w.frontHeadroom + len(p)) writeback := &systemWriteback{tun: tunInterface, frontHeadroom: frontHeadroom}
defer newPacket.Release() if linuxTUN, isLinuxTUN := tunInterface.(LinuxTUN); isLinuxTUN {
newPacket.Resize(w.frontHeadroom, 0) writeback.linuxTUN = linuxTUN
newPacket.Write(p) }
ipHdr := header.IPv4(newPacket.Bytes()) return writeback
ipHdr.SetDestinationAddr(w.source) }
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
func (w *systemWriteback) ReturnHeadroom() int {
return w.frontHeadroom + PacketOffset
}
func (w *systemWriteback) WriteReturnPackets(packets [][]byte) error {
if w.linuxTUN != nil {
return common.Error(w.linuxTUN.BatchWrite(packets, w.frontHeadroom))
}
var writeErrors []error
for _, packet := range packets {
if PacketOffset > 0 { if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version) PacketFillHeader(packet, header.IPVersion(packet[PacketOffset:]))
} else {
newPacket.Advance(-w.frontHeadroom)
} }
return common.Error(w.tun.Write(newPacket.Bytes())) _, err := w.tun.Write(packet)
if err != nil {
writeErrors = append(writeErrors, err)
} }
type systemICMPDirectPacketWriter6 struct {
tun Tun
frontHeadroom int
source netip.Addr
} }
return E.Errors(writeErrors...)
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)
}
return common.Error(w.tun.Write(newPacket.Bytes()))
} }

View file

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

14
tun.go
View file

@ -7,7 +7,6 @@ import (
"runtime" "runtime"
"strconv" "strconv"
"strings" "strings"
"time"
"github.com/sagernet/sing/common" "github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf" "github.com/sagernet/sing/common/buf"
@ -15,27 +14,16 @@ import (
E "github.com/sagernet/sing/common/exceptions" E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format" F "github.com/sagernet/sing/common/format"
"github.com/sagernet/sing/common/logger" "github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network" N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/ranges" "github.com/sagernet/sing/common/ranges"
) )
type Handler interface { type Handler interface {
PrepareConnection( JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort) FlowVerdict
network string,
source M.Socksaddr,
destination M.Socksaddr,
routeContext DirectRouteContext,
timeout time.Duration,
) (DirectRouteDestination, error)
N.TCPConnectionHandlerEx N.TCPConnectionHandlerEx
N.UDPConnectionHandlerEx N.UDPConnectionHandlerEx
} }
type DirectRouteContext interface {
WritePacket(packet []byte) error
}
type Tun interface { type Tun interface {
io.ReadWriter io.ReadWriter
Name() (string, error) 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) { func (t *NativeTun) BatchWrite(buffers [][]byte, offset int) (int, error) {
t.writeAccess.Lock() t.writeAccess.Lock()
defer func() { defer func() {
if t.vnetHdr {
t.tcpGROTable.reset() t.tcpGROTable.reset()
t.udpGROTable.reset() t.udpGROTable.reset()
}
t.writeAccess.Unlock() t.writeAccess.Unlock()
}() }()
var ( var (