Add flow dispatcher
This commit is contained in:
parent
47bdde06c3
commit
ed63adda33
27 changed files with 2469 additions and 963 deletions
32
flow.go
Normal file
32
flow.go
Normal 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
631
flow_dispatch.go
Normal 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
159
flow_mtu.go
Normal 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
106
flow_nat.go
Normal 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
256
flow_parse.go
Normal 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
210
flow_reject.go
Normal 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
350
flow_rewrite.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -4,15 +4,12 @@ package tun
|
|||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/netip"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/sagernet/sing-tun/gtcpip/header"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
|
||||
"github.com/florianl/go-nfqueue/v2"
|
||||
"github.com/mdlayher/netlink"
|
||||
|
|
@ -105,9 +102,9 @@ const ipv6AuthenticationHeaderIdentifier header.IPv6ExtensionHeaderIdentifier =
|
|||
|
||||
type preMatchPacket struct {
|
||||
protocol uint8
|
||||
network string
|
||||
source M.Socksaddr
|
||||
destination M.Socksaddr
|
||||
source netip.AddrPort
|
||||
destination netip.AddrPort
|
||||
firstPacket []byte
|
||||
}
|
||||
|
||||
func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
|
||||
|
|
@ -161,20 +158,23 @@ func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
|
|||
if !flags.Contains(header.TCPFlagSyn) || flags.Contains(header.TCPFlagAck) {
|
||||
return preMatchPacket{}, false
|
||||
}
|
||||
parsed.network = N.NetworkTCP
|
||||
parsed.source = M.SocksaddrFrom(source, tcpHdr.SourcePort())
|
||||
parsed.destination = M.SocksaddrFrom(destination, tcpHdr.DestinationPort())
|
||||
parsed.source = netip.AddrPortFrom(source, tcpHdr.SourcePort())
|
||||
parsed.destination = netip.AddrPortFrom(destination, tcpHdr.DestinationPort())
|
||||
case uint8(header.UDPProtocolNumber):
|
||||
if len(transport) < header.UDPMinimumSize {
|
||||
return preMatchPacket{}, false
|
||||
}
|
||||
udpHdr := header.UDP(transport)
|
||||
if int(udpHdr.Length()) < header.UDPMinimumSize {
|
||||
udpLength := int(udpHdr.Length())
|
||||
if udpLength < header.UDPMinimumSize {
|
||||
return preMatchPacket{}, false
|
||||
}
|
||||
parsed.network = N.NetworkUDP
|
||||
parsed.source = M.SocksaddrFrom(source, udpHdr.SourcePort())
|
||||
parsed.destination = M.SocksaddrFrom(destination, udpHdr.DestinationPort())
|
||||
if udpLength < len(transport) {
|
||||
transport = transport[:udpLength]
|
||||
}
|
||||
parsed.source = netip.AddrPortFrom(source, udpHdr.SourcePort())
|
||||
parsed.destination = netip.AddrPortFrom(destination, udpHdr.DestinationPort())
|
||||
parsed.firstPacket = header.UDP(transport).Payload()
|
||||
case uint8(header.ICMPv4ProtocolNumber):
|
||||
if !source.Is4() || len(transport) < header.ICMPv4MinimumSize {
|
||||
return preMatchPacket{}, false
|
||||
|
|
@ -183,9 +183,9 @@ func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
|
|||
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
|
||||
return preMatchPacket{}, false
|
||||
}
|
||||
parsed.network = N.NetworkICMP
|
||||
parsed.source = M.SocksaddrFrom(source, 0)
|
||||
parsed.destination = M.SocksaddrFrom(destination, 0)
|
||||
identifier := icmpHdr.Ident()
|
||||
parsed.source = netip.AddrPortFrom(source, identifier)
|
||||
parsed.destination = netip.AddrPortFrom(destination, identifier)
|
||||
case uint8(header.ICMPv6ProtocolNumber):
|
||||
if !source.Is6() || len(transport) < header.ICMPv6MinimumSize {
|
||||
return preMatchPacket{}, false
|
||||
|
|
@ -194,9 +194,9 @@ func parsePreMatchPacket(packet []byte) (preMatchPacket, bool) {
|
|||
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
|
||||
return preMatchPacket{}, false
|
||||
}
|
||||
parsed.network = N.NetworkICMP
|
||||
parsed.source = M.SocksaddrFrom(source, 0)
|
||||
parsed.destination = M.SocksaddrFrom(destination, 0)
|
||||
identifier := icmpHdr.Ident()
|
||||
parsed.source = netip.AddrPortFrom(source, identifier)
|
||||
parsed.destination = netip.AddrPortFrom(destination, identifier)
|
||||
default:
|
||||
return preMatchPacket{}, false
|
||||
}
|
||||
|
|
@ -265,22 +265,26 @@ func (h *nfqueueHandler) handlePacket(attr nfqueue.Attribute) int {
|
|||
return 0
|
||||
}
|
||||
|
||||
_, pErr := h.handler.PrepareConnection(packet.network, packet.source, packet.destination, nil, 0)
|
||||
verdict := h.handler.JudgeFlow(
|
||||
packet.protocol,
|
||||
packet.source,
|
||||
packet.destination,
|
||||
)
|
||||
|
||||
// Use NfRepeat for bypass/reset so the packet re-enters the chain
|
||||
// from the beginning, allowing mark-checking rules to save the mark
|
||||
// to conntrack. NfAccept is a terminal verdict in nftables — it exits
|
||||
// the chain immediately, skipping any rules after the queue statement.
|
||||
switch {
|
||||
case errors.Is(pErr, ErrBypass):
|
||||
switch verdict.Action {
|
||||
case ActionBypass:
|
||||
h.setVerdict(packetID, nfqueue.NfRepeat, h.outputMark)
|
||||
case errors.Is(pErr, ErrReset):
|
||||
case ActionReject:
|
||||
if packet.protocol == uint8(unix.IPPROTO_TCP) {
|
||||
h.setVerdict(packetID, nfqueue.NfRepeat, h.resetMark)
|
||||
} else {
|
||||
h.setVerdict(packetID, nfqueue.NfAccept, 0)
|
||||
}
|
||||
case errors.Is(pErr, ErrDrop):
|
||||
case ActionDrop:
|
||||
h.setVerdict(packetID, nfqueue.NfDrop, 0)
|
||||
default:
|
||||
h.setVerdict(packetID, nfqueue.NfAccept, 0)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import (
|
|||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing-tun"
|
||||
"github.com/sagernet/sing-tun/gtcpip/header"
|
||||
"github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/control"
|
||||
|
|
@ -20,14 +19,16 @@ import (
|
|||
// Although its theoretical maximum may be 64k, I don’t yet know of any practical use case for that. For memory-usage reasons, I’m just using a 2k buffer.
|
||||
const maxICMPPacketSize = 2048
|
||||
|
||||
var _ tun.DirectRouteDestination = (*Destination)(nil)
|
||||
type PacketWriter interface {
|
||||
WritePacket(packet []byte) error
|
||||
}
|
||||
|
||||
type Destination struct {
|
||||
conn *Conn
|
||||
ctx context.Context
|
||||
logger logger.ContextLogger
|
||||
destination netip.Addr
|
||||
routeContext tun.DirectRouteContext
|
||||
writer PacketWriter
|
||||
timeout time.Duration
|
||||
requestAccess sync.Mutex
|
||||
requests map[pingRequest]time.Time
|
||||
|
|
@ -45,9 +46,9 @@ func ConnectDestination(
|
|||
logger logger.ContextLogger,
|
||||
controlFunc control.Func,
|
||||
destination netip.Addr,
|
||||
routeContext tun.DirectRouteContext,
|
||||
writer PacketWriter,
|
||||
timeout time.Duration,
|
||||
) (tun.DirectRouteDestination, error) {
|
||||
) (*Destination, error) {
|
||||
var (
|
||||
conn *Conn
|
||||
err error
|
||||
|
|
@ -65,13 +66,13 @@ func ConnectDestination(
|
|||
return nil, err
|
||||
}
|
||||
d := &Destination{
|
||||
conn: conn,
|
||||
ctx: ctx,
|
||||
logger: logger,
|
||||
destination: destination,
|
||||
routeContext: routeContext,
|
||||
timeout: timeout,
|
||||
requests: make(map[pingRequest]time.Time),
|
||||
conn: conn,
|
||||
ctx: ctx,
|
||||
logger: logger,
|
||||
destination: destination,
|
||||
writer: writer,
|
||||
timeout: timeout,
|
||||
requests: make(map[pingRequest]time.Time),
|
||||
}
|
||||
go d.loopRead()
|
||||
return d, nil
|
||||
|
|
@ -158,7 +159,7 @@ func (d *Destination) loopRead() {
|
|||
}
|
||||
d.logger.TraceContext(d.ctx, "read ICMPv6 echo reply from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence())
|
||||
}
|
||||
err = d.routeContext.WritePacket(buffer.Bytes())
|
||||
err = d.writer.WritePacket(buffer.Bytes())
|
||||
if err != nil {
|
||||
d.logger.ErrorContext(d.ctx, E.Cause(err, "write ICMP echo reply"))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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
194
ping/port.go
Normal 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...)
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -139,22 +139,24 @@ func (r *autoRedirect) Start() error {
|
|||
r.redirectServer = server
|
||||
}
|
||||
if r.useNFTables {
|
||||
var handler *nfqueueHandler
|
||||
handler, err = newNFQueueHandler(nfqueueOptions{
|
||||
Context: r.ctx,
|
||||
Handler: r.handler,
|
||||
Logger: r.logger,
|
||||
Queue: r.effectiveNFQueue(),
|
||||
OutputMark: r.effectiveOutputMark(),
|
||||
ResetMark: r.effectiveResetMark(),
|
||||
})
|
||||
if err != nil {
|
||||
r.logger.Warn("nfqueue not available, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
|
||||
} else if err = handler.Start(); err != nil {
|
||||
r.logger.Warn("nfqueue start failed, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
|
||||
} else {
|
||||
r.nfqueueHandler = handler
|
||||
r.nfqueueEnabled = true
|
||||
if r.handler != nil {
|
||||
var handler *nfqueueHandler
|
||||
handler, err = newNFQueueHandler(nfqueueOptions{
|
||||
Context: r.ctx,
|
||||
Handler: r.handler,
|
||||
Logger: r.logger,
|
||||
Queue: r.effectiveNFQueue(),
|
||||
OutputMark: r.effectiveOutputMark(),
|
||||
ResetMark: r.effectiveResetMark(),
|
||||
})
|
||||
if err != nil {
|
||||
r.logger.Warn("nfqueue not available, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
|
||||
} else if err = handler.Start(); err != nil {
|
||||
r.logger.Warn("nfqueue start failed, pre-match disabled (missing nfnetlink_queue and nft_queue kernel module?): ", err)
|
||||
} else {
|
||||
r.nfqueueHandler = handler
|
||||
r.nfqueueEnabled = true
|
||||
}
|
||||
}
|
||||
r.cleanupNFTables()
|
||||
err = r.setupNFTables()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
6
stack.go
6
stack.go
|
|
@ -12,12 +12,6 @@ import (
|
|||
"github.com/sagernet/sing/common/logger"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrDrop = E.New("drop by rule")
|
||||
ErrReset = E.New("reset by rule")
|
||||
ErrBypass = E.New("bypass by rule")
|
||||
)
|
||||
|
||||
type Stack interface {
|
||||
Start() error
|
||||
Close() error
|
||||
|
|
|
|||
|
|
@ -6,8 +6,10 @@ import (
|
|||
"context"
|
||||
"net/netip"
|
||||
"runtime"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/gvisor/pkg/buffer"
|
||||
"github.com/sagernet/gvisor/pkg/tcpip"
|
||||
"github.com/sagernet/gvisor/pkg/tcpip/adapters/gonet"
|
||||
"github.com/sagernet/gvisor/pkg/tcpip/header"
|
||||
|
|
@ -40,6 +42,8 @@ type GVisor struct {
|
|||
logger logger.Logger
|
||||
stack *stack.Stack
|
||||
endpoint stack.LinkEndpoint
|
||||
dispatcher *ForwardDispatcher
|
||||
icmpForwarder *ICMPForwarder
|
||||
}
|
||||
|
||||
type GVisorTun interface {
|
||||
|
|
@ -88,23 +92,39 @@ func (t *GVisor) Start() error {
|
|||
if err != nil {
|
||||
return err
|
||||
}
|
||||
linkEndpoint = &LinkEndpointFilter{linkEndpoint, t.broadcastAddr, t.tun}
|
||||
if t.handler != nil {
|
||||
t.dispatcher = NewForwardDispatcher(t.handler, &gvisorWriteback{tun: t.tun}, t.logger, t.udpTimeout, t.icmpTimeout)
|
||||
}
|
||||
linkEndpoint = &LinkEndpointFilter{
|
||||
LinkEndpoint: linkEndpoint,
|
||||
BroadcastAddress: t.broadcastAddr,
|
||||
Writer: t.tun,
|
||||
Dispatcher: t.dispatcher,
|
||||
Inet4Address: t.inet4Address,
|
||||
Inet6Address: t.inet6Address,
|
||||
Inet4LoopbackAddress: t.inet4LoopbackAddress,
|
||||
Inet6LoopbackAddress: t.inet6LoopbackAddress,
|
||||
}
|
||||
ipStack, err := newGVisorStack(linkEndpoint, nicOptions, false, true)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ipStack.SetTransportProtocolHandler(tcp.ProtocolNumber, NewTCPForwarderWithLoopback(t.ctx, ipStack, t.handler, t.inet4LoopbackAddress, t.inet6LoopbackAddress, t.tun).HandlePacket)
|
||||
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpTimeout).HandlePacket)
|
||||
icmpForwarder := NewICMPForwarder(t.ctx, ipStack, t.logger, t.handler, t.icmpTimeout)
|
||||
icmpForwarder.SetLocalAddresses(t.inet4Address, t.inet6Address)
|
||||
icmpForwarder := NewICMPForwarder(ipStack, t.handler, t.logger)
|
||||
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket)
|
||||
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket)
|
||||
t.icmpForwarder = icmpForwarder
|
||||
t.stack = ipStack
|
||||
t.endpoint = linkEndpoint
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *GVisor) Close() error {
|
||||
t.dispatcher.Close()
|
||||
if t.icmpForwarder != nil {
|
||||
t.icmpForwarder.Close()
|
||||
}
|
||||
if t.stack == nil {
|
||||
return nil
|
||||
}
|
||||
|
|
@ -116,6 +136,37 @@ func (t *GVisor) Close() error {
|
|||
return nil
|
||||
}
|
||||
|
||||
type gvisorWriteback struct {
|
||||
tun GVisorTun
|
||||
access sync.Mutex
|
||||
}
|
||||
|
||||
func (w *gvisorWriteback) ReturnHeadroom() int {
|
||||
return 0
|
||||
}
|
||||
|
||||
func (w *gvisorWriteback) WriteReturnPackets(packets [][]byte) error {
|
||||
w.access.Lock()
|
||||
defer w.access.Unlock()
|
||||
var writeErrors []error
|
||||
for _, packet := range packets {
|
||||
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
Payload: buffer.MakeWithData(packet),
|
||||
})
|
||||
if header.IPVersion(packet) == header.IPv6Version {
|
||||
pkt.NetworkProtocolNumber = header.IPv6ProtocolNumber
|
||||
} else {
|
||||
pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber
|
||||
}
|
||||
_, err := w.tun.WritePacket(pkt)
|
||||
pkt.DecRef()
|
||||
if err != nil {
|
||||
writeErrors = append(writeErrors, err)
|
||||
}
|
||||
}
|
||||
return E.Errors(writeErrors...)
|
||||
}
|
||||
|
||||
func AddressFromAddr(destination netip.Addr) tcpip.Address {
|
||||
if destination.Is6() {
|
||||
return tcpip.AddrFrom16(destination.As16())
|
||||
|
|
|
|||
|
|
@ -14,20 +14,39 @@ var _ stack.LinkEndpoint = (*LinkEndpointFilter)(nil)
|
|||
|
||||
type LinkEndpointFilter struct {
|
||||
stack.LinkEndpoint
|
||||
BroadcastAddress netip.Addr
|
||||
Writer GVisorTun
|
||||
BroadcastAddress netip.Addr
|
||||
Writer GVisorTun
|
||||
Dispatcher *ForwardDispatcher
|
||||
Inet4Address netip.Addr
|
||||
Inet6Address netip.Addr
|
||||
Inet4LoopbackAddress []netip.Addr
|
||||
Inet6LoopbackAddress []netip.Addr
|
||||
}
|
||||
|
||||
func (w *LinkEndpointFilter) Attach(dispatcher stack.NetworkDispatcher) {
|
||||
w.LinkEndpoint.Attach(&networkDispatcherFilter{dispatcher, w.BroadcastAddress, w.Writer})
|
||||
w.LinkEndpoint.Attach(&networkDispatcherFilter{
|
||||
NetworkDispatcher: dispatcher,
|
||||
broadcastAddress: w.BroadcastAddress,
|
||||
writer: w.Writer,
|
||||
dispatcher: w.Dispatcher,
|
||||
inet4Address: w.Inet4Address,
|
||||
inet6Address: w.Inet6Address,
|
||||
inet4LoopbackAddress: w.Inet4LoopbackAddress,
|
||||
inet6LoopbackAddress: w.Inet6LoopbackAddress,
|
||||
})
|
||||
}
|
||||
|
||||
var _ stack.NetworkDispatcher = (*networkDispatcherFilter)(nil)
|
||||
|
||||
type networkDispatcherFilter struct {
|
||||
stack.NetworkDispatcher
|
||||
broadcastAddress netip.Addr
|
||||
writer GVisorTun
|
||||
broadcastAddress netip.Addr
|
||||
writer GVisorTun
|
||||
dispatcher *ForwardDispatcher
|
||||
inet4Address netip.Addr
|
||||
inet6Address netip.Addr
|
||||
inet4LoopbackAddress []netip.Addr
|
||||
inet6LoopbackAddress []netip.Addr
|
||||
}
|
||||
|
||||
func (w *networkDispatcherFilter) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) {
|
||||
|
|
@ -50,5 +69,45 @@ func (w *networkDispatcherFilter) DeliverNetworkPacket(protocol tcpip.NetworkPro
|
|||
w.writer.WritePacket(pkt)
|
||||
return
|
||||
}
|
||||
if w.dispatcher != nil && pkt.GSOOptions.Type == stack.GSONone && !pkt.GSOOptions.NeedsCsum {
|
||||
if view, loaded := pkt.Data().PullUp(pkt.Data().Size()); loaded {
|
||||
consumed := w.dispatch(protocol, destination, view)
|
||||
w.dispatcher.Flush()
|
||||
if consumed {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
w.NetworkDispatcher.DeliverNetworkPacket(protocol, pkt)
|
||||
}
|
||||
|
||||
func (w *networkDispatcherFilter) dispatch(protocol tcpip.NetworkProtocolNumber, destination netip.Addr, view []byte) bool {
|
||||
if protocol == header.IPv4ProtocolNumber {
|
||||
switch header.IPv4(view).TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
for _, inet4LoopbackAddress := range w.inet4LoopbackAddress {
|
||||
if destination == inet4LoopbackAddress {
|
||||
return false
|
||||
}
|
||||
}
|
||||
case header.ICMPv4ProtocolNumber:
|
||||
if destination == w.inet4Address {
|
||||
return false
|
||||
}
|
||||
}
|
||||
} else {
|
||||
switch header.IPv6(view).TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
for _, inet6LoopbackAddress := range w.inet6LoopbackAddress {
|
||||
if destination == inet6LoopbackAddress {
|
||||
return false
|
||||
}
|
||||
}
|
||||
case header.ICMPv6ProtocolNumber:
|
||||
if destination == w.inet6Address {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return w.dispatcher.Dispatch(view)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,10 +3,9 @@
|
|||
package tun
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/gvisor/pkg/buffer"
|
||||
|
|
@ -17,42 +16,51 @@ import (
|
|||
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv4"
|
||||
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv6"
|
||||
"github.com/sagernet/gvisor/pkg/tcpip/stack"
|
||||
"github.com/sagernet/sing/common/buf"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
)
|
||||
|
||||
type ICMPForwarder struct {
|
||||
ctx context.Context
|
||||
stack *stack.Stack
|
||||
logger logger.Logger
|
||||
inet4Address netip.Addr
|
||||
inet6Address netip.Addr
|
||||
handler Handler
|
||||
mapping *DirectRouteMapping
|
||||
stack *stack.Stack
|
||||
handler Handler
|
||||
logger logger.Logger
|
||||
|
||||
returnPath icmpForwarderReturn
|
||||
|
||||
flowAccess sync.Mutex
|
||||
flows map[icmpFlowKey]time.Time
|
||||
lastSweep time.Time
|
||||
attachedPorts map[Port]bool
|
||||
}
|
||||
|
||||
func NewICMPForwarder(
|
||||
ctx context.Context,
|
||||
stack *stack.Stack,
|
||||
logger logger.Logger,
|
||||
handler Handler,
|
||||
timeout time.Duration,
|
||||
) *ICMPForwarder {
|
||||
return &ICMPForwarder{
|
||||
ctx: ctx,
|
||||
stack: stack,
|
||||
logger: logger,
|
||||
handler: handler,
|
||||
mapping: NewDirectRouteMapping(timeout),
|
||||
type icmpFlowKey struct {
|
||||
v6 bool
|
||||
source netip.Addr
|
||||
destination netip.Addr
|
||||
identifier uint16
|
||||
}
|
||||
|
||||
func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) *ICMPForwarder {
|
||||
forwarder := &ICMPForwarder{
|
||||
stack: stack,
|
||||
handler: handler,
|
||||
logger: logger,
|
||||
flows: make(map[icmpFlowKey]time.Time),
|
||||
attachedPorts: make(map[Port]bool),
|
||||
}
|
||||
forwarder.returnPath.forwarder = forwarder
|
||||
return forwarder
|
||||
}
|
||||
|
||||
func (f *ICMPForwarder) SetLocalAddresses(inet4Address, inet6Address netip.Addr) {
|
||||
f.inet4Address = inet4Address
|
||||
f.inet6Address = inet6Address
|
||||
func (f *ICMPForwarder) Close() error {
|
||||
f.returnPath.closed.Store(true)
|
||||
f.flowAccess.Lock()
|
||||
defer f.flowAccess.Unlock()
|
||||
for port := range f.attachedPorts {
|
||||
port.DetachReturn(&f.returnPath)
|
||||
delete(f.attachedPorts, port)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
|
|
@ -62,34 +70,17 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
|||
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
|
||||
return false
|
||||
}
|
||||
sourceAddr := M.AddrFromIP(ipHdr.SourceAddressSlice())
|
||||
destinationAddr := M.AddrFromIP(ipHdr.DestinationAddressSlice())
|
||||
if destinationAddr != f.inet4Address {
|
||||
action, err := f.mapping.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
|
||||
return f.handler.PrepareConnection(
|
||||
N.NetworkICMP,
|
||||
M.SocksaddrFrom(sourceAddr, 0),
|
||||
M.SocksaddrFrom(destinationAddr, 0),
|
||||
&ICMPBackWriter{
|
||||
stack: f.stack,
|
||||
packet: pkt,
|
||||
source: ipHdr.SourceAddress(),
|
||||
sourceNetwork: header.IPv4ProtocolNumber,
|
||||
},
|
||||
timeout,
|
||||
)
|
||||
})
|
||||
if errors.Is(err, ErrReset) {
|
||||
gWriteUnreachable(f.stack, pkt)
|
||||
return true
|
||||
} else if errors.Is(err, ErrDrop) {
|
||||
return true
|
||||
}
|
||||
if action != nil {
|
||||
err = icmpWritePacketBuffer(action, pkt)
|
||||
if err != nil {
|
||||
f.logger.Error(E.Cause(err, "write ICMPv4 echo request"))
|
||||
}
|
||||
identifier := icmpHdr.Ident()
|
||||
verdict := f.handler.JudgeFlow(
|
||||
uint8(header.ICMPv4ProtocolNumber),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
|
||||
)
|
||||
switch verdict.Action {
|
||||
case ActionReject, ActionDrop:
|
||||
return true
|
||||
case ActionFlow:
|
||||
if f.forwardFlow(verdict.Port, false, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
|
@ -125,35 +116,17 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
|||
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
|
||||
return false
|
||||
}
|
||||
sourceAddr := M.AddrFromIP(ipHdr.SourceAddressSlice())
|
||||
destinationAddr := M.AddrFromIP(ipHdr.DestinationAddressSlice())
|
||||
if destinationAddr != f.inet6Address {
|
||||
action, err := f.mapping.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
|
||||
return f.handler.PrepareConnection(
|
||||
N.NetworkICMP,
|
||||
M.SocksaddrFrom(sourceAddr, 0),
|
||||
M.SocksaddrFrom(destinationAddr, 0),
|
||||
&ICMPBackWriter{
|
||||
stack: f.stack,
|
||||
packet: pkt,
|
||||
source: ipHdr.SourceAddress(),
|
||||
sourceNetwork: header.IPv6ProtocolNumber,
|
||||
},
|
||||
timeout,
|
||||
)
|
||||
})
|
||||
if errors.Is(err, ErrReset) {
|
||||
gWriteUnreachable(f.stack, pkt)
|
||||
return true
|
||||
} else if errors.Is(err, ErrDrop) {
|
||||
return true
|
||||
}
|
||||
if action != nil {
|
||||
pkt.IncRef()
|
||||
err = icmpWritePacketBuffer(action, pkt)
|
||||
if err != nil {
|
||||
f.logger.Error(E.Cause(err, "write ICMPv6 echo request"))
|
||||
}
|
||||
identifier := icmpHdr.Ident()
|
||||
verdict := f.handler.JudgeFlow(
|
||||
uint8(header.ICMPv6ProtocolNumber),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
|
||||
)
|
||||
switch verdict.Action {
|
||||
case ActionReject, ActionDrop:
|
||||
return true
|
||||
case ActionFlow:
|
||||
if f.forwardFlow(verdict.Port, true, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
|
@ -190,64 +163,179 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
|||
}
|
||||
}
|
||||
|
||||
type ICMPBackWriter struct {
|
||||
access sync.Mutex
|
||||
stack *stack.Stack
|
||||
packet *stack.PacketBuffer
|
||||
source tcpip.Address
|
||||
sourceNetwork tcpip.NetworkProtocolNumber
|
||||
}
|
||||
|
||||
func (w *ICMPBackWriter) WritePacket(p []byte) error {
|
||||
if w.sourceNetwork == header.IPv4ProtocolNumber {
|
||||
route, err := w.stack.FindRoute(
|
||||
DefaultNIC,
|
||||
header.IPv4(p).SourceAddress(),
|
||||
w.source,
|
||||
w.sourceNetwork,
|
||||
false,
|
||||
)
|
||||
func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, destination netip.Addr, identifier uint16, pkt *stack.PacketBuffer) bool {
|
||||
if port == nil {
|
||||
return false
|
||||
}
|
||||
inet4Address, inet6Address := port.PortAddresses()
|
||||
portAddress := inet4Address
|
||||
if v6 {
|
||||
portAddress = inet6Address
|
||||
}
|
||||
if !portAddress.IsValid() || !portAddress.IsUnspecified() {
|
||||
return false
|
||||
}
|
||||
f.flowAccess.Lock()
|
||||
if !f.attachedPorts[port] {
|
||||
err := port.AttachReturn(&f.returnPath)
|
||||
if err != nil {
|
||||
return gonet.TranslateNetstackError(err)
|
||||
f.flowAccess.Unlock()
|
||||
f.logger.Trace(E.Cause(err, "attach ICMP return path"))
|
||||
return false
|
||||
}
|
||||
defer route.Release()
|
||||
packet := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
Payload: buffer.MakeWithData(p),
|
||||
})
|
||||
defer packet.DecRef()
|
||||
parse.IPv4(packet)
|
||||
err = route.WritePacketDirect(packet)
|
||||
if err != nil {
|
||||
return gonet.TranslateNetstackError(err)
|
||||
}
|
||||
} else {
|
||||
route, err := w.stack.FindRoute(
|
||||
DefaultNIC,
|
||||
header.IPv6(p).SourceAddress(),
|
||||
w.source,
|
||||
w.sourceNetwork,
|
||||
false,
|
||||
)
|
||||
if err != nil {
|
||||
return gonet.TranslateNetstackError(err)
|
||||
}
|
||||
defer route.Release()
|
||||
packet := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
Payload: buffer.MakeWithData(p),
|
||||
})
|
||||
parse.IPv6(packet)
|
||||
defer packet.DecRef()
|
||||
err = route.WritePacketDirect(packet)
|
||||
if err != nil {
|
||||
return gonet.TranslateNetstackError(err)
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
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 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())
|
||||
func (f *ICMPForwarder) lookupFlow(key icmpFlowKey) bool {
|
||||
f.flowAccess.Lock()
|
||||
defer f.flowAccess.Unlock()
|
||||
deadline, loaded := f.flows[key]
|
||||
if !loaded {
|
||||
return false
|
||||
}
|
||||
now := time.Now()
|
||||
if now.After(deadline) {
|
||||
delete(f.flows, key)
|
||||
return false
|
||||
}
|
||||
f.flows[key] = now.Add(defaultICMPTimeout)
|
||||
return true
|
||||
}
|
||||
|
||||
type icmpForwarderReturn struct {
|
||||
forwarder *ICMPForwarder
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func (r *icmpForwarderReturn) ReturnHeadroom() int {
|
||||
return 0
|
||||
}
|
||||
|
||||
func (r *icmpForwarderReturn) ReturnPackets(packets [][]byte) [][]byte {
|
||||
if r.closed.Load() {
|
||||
return packets
|
||||
}
|
||||
unconsumed := packets[:0]
|
||||
for _, packet := range packets {
|
||||
if !r.forwarder.returnPacket(packet) {
|
||||
unconsumed = append(unconsumed, packet)
|
||||
}
|
||||
}
|
||||
return unconsumed
|
||||
}
|
||||
|
||||
func (f *ICMPForwarder) returnPacket(packet []byte) bool {
|
||||
if len(packet) == 0 {
|
||||
return false
|
||||
}
|
||||
switch header.IPVersion(packet) {
|
||||
case header.IPv4Version:
|
||||
ipHdr := header.IPv4(packet)
|
||||
if !ipHdr.IsValid(len(packet)) || ipHdr.TransportProtocol() != header.ICMPv4ProtocolNumber || len(ipHdr.Payload()) < header.ICMPv4MinimumSize {
|
||||
return false
|
||||
}
|
||||
icmpHdr := header.ICMPv4(ipHdr.Payload())
|
||||
var key icmpFlowKey
|
||||
switch icmpHdr.Type() {
|
||||
case header.ICMPv4EchoReply:
|
||||
key = icmpFlowKey{
|
||||
source: AddrFromAddress(ipHdr.DestinationAddress()),
|
||||
destination: AddrFromAddress(ipHdr.SourceAddress()),
|
||||
identifier: icmpHdr.Ident(),
|
||||
}
|
||||
case header.ICMPv4TimeExceeded, header.ICMPv4DstUnreachable:
|
||||
inner := icmpHdr.Payload()
|
||||
if len(inner) < header.IPv4MinimumSize {
|
||||
return false
|
||||
}
|
||||
innerIPHdr := header.IPv4(inner)
|
||||
innerHeaderLength := int(innerIPHdr.HeaderLength())
|
||||
if innerHeaderLength < header.IPv4MinimumSize || len(inner) < innerHeaderLength+header.ICMPv4MinimumSize {
|
||||
return false
|
||||
}
|
||||
if innerIPHdr.TransportProtocol() != header.ICMPv4ProtocolNumber {
|
||||
return false
|
||||
}
|
||||
innerICMPHdr := header.ICMPv4(inner[innerHeaderLength:])
|
||||
key = icmpFlowKey{
|
||||
source: AddrFromAddress(innerIPHdr.SourceAddress()),
|
||||
destination: AddrFromAddress(innerIPHdr.DestinationAddress()),
|
||||
identifier: innerICMPHdr.Ident(),
|
||||
}
|
||||
default:
|
||||
return false
|
||||
}
|
||||
if !f.lookupFlow(key) {
|
||||
return false
|
||||
}
|
||||
return f.writeBack(packet, header.IPv4ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
||||
case header.IPv6Version:
|
||||
ipHdr := header.IPv6(packet)
|
||||
if !ipHdr.IsValid(len(packet)) || ipHdr.TransportProtocol() != header.ICMPv6ProtocolNumber || len(ipHdr.Payload()) < header.ICMPv6MinimumSize {
|
||||
return false
|
||||
}
|
||||
icmpHdr := header.ICMPv6(ipHdr.Payload())
|
||||
if icmpHdr.Type() != header.ICMPv6EchoReply {
|
||||
return false
|
||||
}
|
||||
key := icmpFlowKey{
|
||||
v6: true,
|
||||
source: AddrFromAddress(ipHdr.DestinationAddress()),
|
||||
destination: AddrFromAddress(ipHdr.SourceAddress()),
|
||||
identifier: icmpHdr.Ident(),
|
||||
}
|
||||
if !f.lookupFlow(key) {
|
||||
return false
|
||||
}
|
||||
return f.writeBack(packet, header.IPv6ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (f *ICMPForwarder) writeBack(packet []byte, protocol tcpip.NetworkProtocolNumber, localAddress tcpip.Address, remoteAddress tcpip.Address) bool {
|
||||
route, gErr := f.stack.FindRoute(DefaultNIC, localAddress, remoteAddress, protocol, false)
|
||||
if gErr != nil {
|
||||
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "find route for ICMP reply"))
|
||||
return true
|
||||
}
|
||||
defer route.Release()
|
||||
packetBuffer := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
Payload: buffer.MakeWithData(packet),
|
||||
})
|
||||
defer packetBuffer.DecRef()
|
||||
if protocol == header.IPv4ProtocolNumber {
|
||||
parse.IPv4(packetBuffer)
|
||||
} else {
|
||||
parse.IPv6(packetBuffer)
|
||||
}
|
||||
gErr = route.WritePacketDirect(packetBuffer)
|
||||
if gErr != nil {
|
||||
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "write ICMP reply"))
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ package tun
|
|||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"os"
|
||||
"sync"
|
||||
|
|
@ -74,7 +73,7 @@ func (c *gLazyConn) HandshakeFailure(err error) error {
|
|||
if c.handshakeDone {
|
||||
return os.ErrInvalid
|
||||
}
|
||||
c.request.Complete(!errors.Is(err, ErrDrop))
|
||||
c.request.Complete(true)
|
||||
c.handshakeDone = true
|
||||
c.handshakeErr = err
|
||||
return nil
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ package tun
|
|||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/netip"
|
||||
|
||||
"github.com/sagernet/gvisor/pkg/tcpip"
|
||||
|
|
@ -14,7 +13,6 @@ import (
|
|||
"github.com/sagernet/sing-tun/gtcpip/checksum"
|
||||
"github.com/sagernet/sing/common"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
)
|
||||
|
||||
type TCPForwarder struct {
|
||||
|
|
@ -79,9 +77,12 @@ func (f *TCPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac
|
|||
func (f *TCPForwarder) Forward(r *tcp.ForwarderRequest) {
|
||||
source := M.SocksaddrFrom(AddrFromAddress(r.ID().RemoteAddress), r.ID().RemotePort)
|
||||
destination := M.SocksaddrFrom(AddrFromAddress(r.ID().LocalAddress), r.ID().LocalPort)
|
||||
_, pErr := f.handler.PrepareConnection(N.NetworkTCP, source, destination, nil, 0)
|
||||
if pErr != nil {
|
||||
r.Complete(!errors.Is(pErr, ErrDrop))
|
||||
switch f.handler.JudgeFlow(uint8(header.TCPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action {
|
||||
case ActionReject:
|
||||
r.Complete(true)
|
||||
return
|
||||
case ActionDrop:
|
||||
r.Complete(false)
|
||||
return
|
||||
}
|
||||
conn := &gLazyConn{
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ package tun
|
|||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"math"
|
||||
"net/netip"
|
||||
"os"
|
||||
|
|
@ -58,11 +57,11 @@ func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac
|
|||
func rangeIterate(r stack.Range, fn func(*buffer.View))
|
||||
|
||||
func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) {
|
||||
_, pErr := f.handler.PrepareConnection(N.NetworkUDP, source, destination, nil, 0)
|
||||
if pErr != nil {
|
||||
if !errors.Is(pErr, ErrDrop) {
|
||||
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
|
||||
}
|
||||
switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action {
|
||||
case ActionReject:
|
||||
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
|
||||
return false, nil, nil, nil
|
||||
case ActionDrop:
|
||||
return false, nil, nil, nil
|
||||
}
|
||||
var sourceNetwork tcpip.NetworkProtocolNumber
|
||||
|
|
|
|||
|
|
@ -73,8 +73,6 @@ func (m *Mixed) tunLoop() {
|
|||
return
|
||||
}
|
||||
if linuxTUN, isLinuxTUN := m.tun.(LinuxTUN); isLinuxTUN {
|
||||
m.frontHeadroom = linuxTUN.FrontHeadroom()
|
||||
m.txChecksumOffload = linuxTUN.TXChecksumOffload()
|
||||
batchSize := linuxTUN.BatchSize()
|
||||
if batchSize > 1 {
|
||||
m.batchLoopLinux(linuxTUN, batchSize)
|
||||
|
|
@ -105,6 +103,7 @@ func (m *Mixed) tunLoop() {
|
|||
m.logger.Trace(E.Cause(err, "write packet"))
|
||||
}
|
||||
}
|
||||
m.dispatcher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -124,6 +123,7 @@ func (m *Mixed) wintunLoop(winTun WinTun) {
|
|||
m.logger.Trace(E.Cause(err, "write packet"))
|
||||
}
|
||||
}
|
||||
m.dispatcher.Flush()
|
||||
release()
|
||||
}
|
||||
}
|
||||
|
|
@ -164,11 +164,13 @@ func (m *Mixed) batchLoopLinux(linuxTUN LinuxTUN, batchSize int) {
|
|||
}
|
||||
writeBuffers = writeBuffers[:0]
|
||||
}
|
||||
m.dispatcher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
|
||||
var writeBuffers []*buf.Buffer
|
||||
var releaseBuffers []*buf.Buffer
|
||||
for {
|
||||
buffers, err := darwinTUN.BatchRead()
|
||||
if err != nil {
|
||||
|
|
@ -181,6 +183,7 @@ func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
|
|||
continue
|
||||
}
|
||||
writeBuffers = writeBuffers[:0]
|
||||
releaseBuffers = releaseBuffers[:0]
|
||||
for _, buffer := range buffers {
|
||||
packetSize := buffer.Len()
|
||||
if packetSize < header.IPv4MinimumSize {
|
||||
|
|
@ -190,7 +193,7 @@ func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
|
|||
if m.processPacket(buffer.Bytes()) {
|
||||
writeBuffers = append(writeBuffers, buffer)
|
||||
} else {
|
||||
buffer.Release()
|
||||
releaseBuffers = append(releaseBuffers, buffer)
|
||||
}
|
||||
}
|
||||
if len(writeBuffers) > 0 {
|
||||
|
|
@ -200,6 +203,8 @@ func (m *Mixed) batchLoopDarwin(darwinTUN DarwinTUN) {
|
|||
}
|
||||
buf.ReleaseMulti(writeBuffers)
|
||||
}
|
||||
m.dispatcher.Flush()
|
||||
buf.ReleaseMulti(releaseBuffers)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -229,6 +234,9 @@ func (m *Mixed) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
|
|||
if destination == m.broadcastAddr || !destination.IsGlobalUnicast() {
|
||||
return
|
||||
}
|
||||
if m.dispatchIPv4(ipHdr, destination) {
|
||||
return false, nil
|
||||
}
|
||||
switch ipHdr.TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
writeBack, err = m.processIPv4TCP(ipHdr, ipHdr.Payload())
|
||||
|
|
@ -249,9 +257,13 @@ func (m *Mixed) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
|
|||
|
||||
func (m *Mixed) processIPv6(ipHdr header.IPv6) (writeBack bool, err error) {
|
||||
writeBack = true
|
||||
if !ipHdr.DestinationAddr().IsGlobalUnicast() {
|
||||
destination := ipHdr.DestinationAddr()
|
||||
if !destination.IsGlobalUnicast() {
|
||||
return
|
||||
}
|
||||
if m.dispatchIPv6(ipHdr, destination) {
|
||||
return false, nil
|
||||
}
|
||||
switch ipHdr.TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
writeBack, err = m.processIPv6TCP(ipHdr, ipHdr.Payload())
|
||||
|
|
|
|||
372
stack_system.go
372
stack_system.go
|
|
@ -5,6 +5,7 @@ import (
|
|||
"errors"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
|
|
@ -46,7 +47,7 @@ type System struct {
|
|||
tcpPort6 uint16
|
||||
tcpNat *TCPNat
|
||||
udpNat *udpnat.Service
|
||||
directNat *DirectRouteMapping
|
||||
dispatcher *ForwardDispatcher
|
||||
bindInterface bool
|
||||
interfaceFinder control.InterfaceFinder
|
||||
frontHeadroom int
|
||||
|
|
@ -101,6 +102,7 @@ func NewSystem(options StackOptions) (Stack, error) {
|
|||
}
|
||||
|
||||
func (s *System) Close() error {
|
||||
s.dispatcher.Close()
|
||||
return common.Close(
|
||||
s.tcpListener,
|
||||
s.tcpListener6,
|
||||
|
|
@ -162,7 +164,13 @@ func (s *System) start() error {
|
|||
}
|
||||
s.tcpNat = NewNat(s.ctx, s.udpTimeout)
|
||||
s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false)
|
||||
s.directNat = NewDirectRouteMapping(s.icmpTimeout)
|
||||
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
|
||||
s.frontHeadroom = linuxTUN.FrontHeadroom()
|
||||
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
|
||||
}
|
||||
if s.handler != nil {
|
||||
s.dispatcher = NewForwardDispatcher(s.handler, newSystemWriteback(s.tun, s.frontHeadroom), s.logger, s.udpTimeout, s.icmpTimeout)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
|
@ -172,8 +180,6 @@ func (s *System) tunLoop() {
|
|||
return
|
||||
}
|
||||
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
|
||||
s.frontHeadroom = linuxTUN.FrontHeadroom()
|
||||
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
|
||||
batchSize := linuxTUN.BatchSize()
|
||||
if batchSize > 1 {
|
||||
s.batchLoopLinux(linuxTUN, batchSize)
|
||||
|
|
@ -204,6 +210,7 @@ func (s *System) tunLoop() {
|
|||
s.logger.Trace(E.Cause(err, "write packet"))
|
||||
}
|
||||
}
|
||||
s.dispatcher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -223,6 +230,7 @@ func (s *System) wintunLoop(winTun WinTun) {
|
|||
s.logger.Trace(E.Cause(err, "write packet"))
|
||||
}
|
||||
}
|
||||
s.dispatcher.Flush()
|
||||
release()
|
||||
}
|
||||
}
|
||||
|
|
@ -263,11 +271,13 @@ func (s *System) batchLoopLinux(linuxTUN LinuxTUN, batchSize int) {
|
|||
}
|
||||
writeBuffers = writeBuffers[:0]
|
||||
}
|
||||
s.dispatcher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
|
||||
var writeBuffers []*buf.Buffer
|
||||
var releaseBuffers []*buf.Buffer
|
||||
for {
|
||||
buffers, err := darwinTUN.BatchRead()
|
||||
if err != nil {
|
||||
|
|
@ -280,6 +290,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
|
|||
continue
|
||||
}
|
||||
writeBuffers = writeBuffers[:0]
|
||||
releaseBuffers = releaseBuffers[:0]
|
||||
for _, buffer := range buffers {
|
||||
packetSize := buffer.Len()
|
||||
if packetSize < header.IPv4MinimumSize {
|
||||
|
|
@ -289,7 +300,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
|
|||
if s.processPacket(buffer.Bytes()) {
|
||||
writeBuffers = append(writeBuffers, buffer)
|
||||
} else {
|
||||
buffer.Release()
|
||||
releaseBuffers = append(releaseBuffers, buffer)
|
||||
}
|
||||
}
|
||||
if len(writeBuffers) > 0 {
|
||||
|
|
@ -299,6 +310,8 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
|
|||
}
|
||||
buf.ReleaseMulti(writeBuffers)
|
||||
}
|
||||
s.dispatcher.Flush()
|
||||
buf.ReleaseMulti(releaseBuffers)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -338,11 +351,53 @@ func (s *System) acceptLoop(listener net.Listener) {
|
|||
}
|
||||
}
|
||||
|
||||
func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool {
|
||||
switch ipHdr.TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
if slices.Contains(s.inet4LoopbackAddress, destination) {
|
||||
return false
|
||||
}
|
||||
if ipHdr.SourceAddr() == s.inet4Address &&
|
||||
ipHdr.FragmentOffset() == 0 &&
|
||||
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
|
||||
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort {
|
||||
return false
|
||||
}
|
||||
case header.ICMPv4ProtocolNumber:
|
||||
if destination == s.inet4Address {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return s.dispatcher.Dispatch(ipHdr)
|
||||
}
|
||||
|
||||
func (s *System) dispatchIPv6(ipHdr header.IPv6, destination netip.Addr) bool {
|
||||
switch ipHdr.TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
if slices.Contains(s.inet6LoopbackAddress, destination) {
|
||||
return false
|
||||
}
|
||||
if ipHdr.SourceAddr() == s.inet6Address &&
|
||||
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
|
||||
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort6 {
|
||||
return false
|
||||
}
|
||||
case header.ICMPv6ProtocolNumber:
|
||||
if destination == s.inet6Address {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return s.dispatcher.Dispatch(ipHdr)
|
||||
}
|
||||
|
||||
func (s *System) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
|
||||
destination := ipHdr.DestinationAddr()
|
||||
if destination == s.broadcastAddr || !destination.IsGlobalUnicast() {
|
||||
return
|
||||
}
|
||||
if s.dispatchIPv4(ipHdr, destination) {
|
||||
return false, nil
|
||||
}
|
||||
writeBack = true
|
||||
switch ipHdr.TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
|
|
@ -360,9 +415,13 @@ func (s *System) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
|
|||
}
|
||||
|
||||
func (s *System) processIPv6(ipHdr header.IPv6) (writeBack bool, err error) {
|
||||
if !ipHdr.DestinationAddr().IsGlobalUnicast() {
|
||||
destination := ipHdr.DestinationAddr()
|
||||
if !destination.IsGlobalUnicast() {
|
||||
return
|
||||
}
|
||||
if s.dispatchIPv6(ipHdr, destination) {
|
||||
return false, nil
|
||||
}
|
||||
writeBack = true
|
||||
switch ipHdr.TransportProtocol() {
|
||||
case header.TCPProtocolNumber:
|
||||
|
|
@ -404,14 +463,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
|
|||
}
|
||||
}
|
||||
if !loopback {
|
||||
natPort, err := s.tcpNat.Lookup(source, destination, s.handler)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrDrop) {
|
||||
return false, nil
|
||||
} else {
|
||||
return false, s.resetIPv4TCP(ipHdr, tcpHdr)
|
||||
}
|
||||
}
|
||||
natPort := s.tcpNat.Lookup(source, destination)
|
||||
ipHdr.SetSourceAddr(s.inet4NextAddress)
|
||||
tcpHdr.SetSourcePort(natPort)
|
||||
ipHdr.SetDestinationAddr(s.inet4Address)
|
||||
|
|
@ -429,51 +481,6 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
|
|||
return true, nil
|
||||
}
|
||||
|
||||
func (s *System) resetIPv4TCP(origIPHdr header.IPv4, origTCPHdr header.TCP) error {
|
||||
frontHeadroom := s.frontHeadroom + PacketOffset
|
||||
newPacket := buf.NewSize(frontHeadroom + header.IPv4MinimumSize + header.TCPMinimumSize)
|
||||
defer newPacket.Release()
|
||||
newPacket.Resize(frontHeadroom, header.IPv4MinimumSize+header.TCPMinimumSize)
|
||||
ipHdr := header.IPv4(newPacket.Bytes())
|
||||
ipHdr.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(newPacket.Len()),
|
||||
Protocol: uint8(header.TCPProtocolNumber),
|
||||
SrcAddr: origIPHdr.DestinationAddr(),
|
||||
DstAddr: origIPHdr.SourceAddr(),
|
||||
})
|
||||
tcpHdr := header.TCP(ipHdr.Payload())
|
||||
fields := header.TCPFields{
|
||||
SrcPort: origTCPHdr.DestinationPort(),
|
||||
DstPort: origTCPHdr.SourcePort(),
|
||||
DataOffset: header.TCPMinimumSize,
|
||||
Flags: header.TCPFlagRst,
|
||||
}
|
||||
if origTCPHdr.Flags()&header.TCPFlagAck != 0 {
|
||||
fields.SeqNum = origTCPHdr.AckNumber()
|
||||
} else {
|
||||
fields.Flags |= header.TCPFlagAck
|
||||
ackNum := origTCPHdr.SequenceNumber() + uint32(len(origTCPHdr.Payload()))
|
||||
if origTCPHdr.Flags()&header.TCPFlagSyn != 0 {
|
||||
ackNum++
|
||||
}
|
||||
if origTCPHdr.Flags()&header.TCPFlagFin != 0 {
|
||||
ackNum++
|
||||
}
|
||||
fields.AckNum = ackNum
|
||||
}
|
||||
tcpHdr.Encode(&fields)
|
||||
if !s.txChecksumOffload {
|
||||
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
|
||||
}
|
||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||
if PacketOffset > 0 {
|
||||
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
|
||||
} else {
|
||||
newPacket.Advance(-s.frontHeadroom)
|
||||
}
|
||||
return common.Error(s.tun.Write(newPacket.Bytes()))
|
||||
}
|
||||
|
||||
func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, error) {
|
||||
source := netip.AddrPortFrom(ipHdr.SourceAddr(), tcpHdr.SourcePort())
|
||||
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort())
|
||||
|
|
@ -499,14 +506,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
|
|||
}
|
||||
}
|
||||
if !loopback {
|
||||
natPort, err := s.tcpNat.Lookup(source, destination, s.handler)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrDrop) {
|
||||
return false, nil
|
||||
} else {
|
||||
return false, s.resetIPv6TCP(ipHdr, tcpHdr)
|
||||
}
|
||||
}
|
||||
natPort := s.tcpNat.Lookup(source, destination)
|
||||
ipHdr.SetSourceAddr(s.inet6NextAddress)
|
||||
tcpHdr.SetSourcePort(natPort)
|
||||
ipHdr.SetDestinationAddr(s.inet6Address)
|
||||
|
|
@ -523,50 +523,6 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
|
|||
return true, nil
|
||||
}
|
||||
|
||||
func (s *System) resetIPv6TCP(origIPHdr header.IPv6, origTCPHdr header.TCP) error {
|
||||
frontHeadroom := s.frontHeadroom + PacketOffset
|
||||
newPacket := buf.NewSize(frontHeadroom + header.IPv6MinimumSize + header.TCPMinimumSize)
|
||||
defer newPacket.Release()
|
||||
newPacket.Resize(frontHeadroom, header.IPv6MinimumSize+header.TCPMinimumSize)
|
||||
ipHdr := header.IPv6(newPacket.Bytes())
|
||||
ipHdr.Encode(&header.IPv6Fields{
|
||||
PayloadLength: uint16(header.TCPMinimumSize),
|
||||
TransportProtocol: header.TCPProtocolNumber,
|
||||
SrcAddr: origIPHdr.DestinationAddr(),
|
||||
DstAddr: origIPHdr.SourceAddr(),
|
||||
})
|
||||
tcpHdr := header.TCP(ipHdr.Payload())
|
||||
fields := header.TCPFields{
|
||||
SrcPort: origTCPHdr.DestinationPort(),
|
||||
DstPort: origTCPHdr.SourcePort(),
|
||||
DataOffset: header.TCPMinimumSize,
|
||||
Flags: header.TCPFlagRst,
|
||||
}
|
||||
if origTCPHdr.Flags()&header.TCPFlagAck != 0 {
|
||||
fields.SeqNum = origTCPHdr.AckNumber()
|
||||
} else {
|
||||
fields.Flags |= header.TCPFlagAck
|
||||
ackNum := origTCPHdr.SequenceNumber() + uint32(len(origTCPHdr.Payload()))
|
||||
if origTCPHdr.Flags()&header.TCPFlagSyn != 0 {
|
||||
ackNum++
|
||||
}
|
||||
if origTCPHdr.Flags()&header.TCPFlagFin != 0 {
|
||||
ackNum++
|
||||
}
|
||||
fields.AckNum = ackNum
|
||||
}
|
||||
tcpHdr.Encode(&fields)
|
||||
if !s.txChecksumOffload {
|
||||
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
|
||||
}
|
||||
if PacketOffset > 0 {
|
||||
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
|
||||
} else {
|
||||
newPacket.Advance(-s.frontHeadroom)
|
||||
}
|
||||
return common.Error(s.tun.Write(newPacket.Bytes()))
|
||||
}
|
||||
|
||||
func (s *System) processIPv4UDP(ipHdr header.IPv4, udpHdr header.UDP) error {
|
||||
if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 {
|
||||
return E.New("ipv4: fragment dropped")
|
||||
|
|
@ -594,19 +550,6 @@ func (s *System) processIPv6UDP(ipHdr header.IPv6, udpHdr header.UDP) error {
|
|||
}
|
||||
|
||||
func (s *System) preparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) {
|
||||
_, pErr := s.handler.PrepareConnection(N.NetworkUDP, source, destination, nil, 0)
|
||||
if pErr != nil {
|
||||
if !errors.Is(pErr, ErrDrop) {
|
||||
if source.IsIPv4() {
|
||||
ipHdr := userData.(header.IPv4)
|
||||
s.rejectIPv4WithICMP(ipHdr, header.ICMPv4PortUnreachable)
|
||||
} else {
|
||||
ipHdr := userData.(header.IPv6)
|
||||
s.rejectIPv6WithICMP(ipHdr, header.ICMPv6PortUnreachable)
|
||||
}
|
||||
}
|
||||
return false, nil, nil, nil
|
||||
}
|
||||
var writer N.PacketWriter
|
||||
if source.IsIPv4() {
|
||||
packet := userData.(header.IPv4)
|
||||
|
|
@ -640,29 +583,6 @@ func (s *System) processIPv4ICMP(ipHdr header.IPv4, icmpHdr header.ICMPv4) (bool
|
|||
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
|
||||
return false, nil
|
||||
}
|
||||
sourceAddr := ipHdr.SourceAddr()
|
||||
destinationAddr := ipHdr.DestinationAddr()
|
||||
if destinationAddr != s.inet4Address {
|
||||
action, err := s.directNat.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
|
||||
return s.handler.PrepareConnection(
|
||||
N.NetworkICMP,
|
||||
M.SocksaddrFrom(sourceAddr, 0),
|
||||
M.SocksaddrFrom(destinationAddr, 0),
|
||||
&systemICMPDirectPacketWriter4{s.tun, s.frontHeadroom + PacketOffset, sourceAddr},
|
||||
timeout,
|
||||
)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrReset) {
|
||||
return false, s.rejectIPv4WithICMP(ipHdr, header.ICMPv4HostUnreachable)
|
||||
} else if errors.Is(err, ErrDrop) {
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
if action != nil {
|
||||
return false, action.WritePacket(buf.As(ipHdr).ToOwned())
|
||||
}
|
||||
}
|
||||
icmpHdr.SetType(header.ICMPv4EchoReply)
|
||||
sourceAddress := ipHdr.SourceAddr()
|
||||
ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
|
||||
|
|
@ -672,70 +592,10 @@ func (s *System) processIPv4ICMP(ipHdr header.IPv4, icmpHdr header.ICMPv4) (bool
|
|||
return true, nil
|
||||
}
|
||||
|
||||
func (s *System) rejectIPv4WithICMP(ipHdr header.IPv4, code header.ICMPv4Code) error {
|
||||
frontHeadroom := s.frontHeadroom + PacketOffset
|
||||
mtu := s.mtu
|
||||
const maxIPData = header.IPv4MinimumProcessableDatagramSize - header.IPv4MinimumSize
|
||||
if mtu > maxIPData {
|
||||
mtu = maxIPData
|
||||
}
|
||||
available := mtu - header.ICMPv4MinimumSize
|
||||
if available < len(ipHdr)+header.ICMPv4MinimumErrorPayloadSize {
|
||||
return nil
|
||||
}
|
||||
payload := ipHdr
|
||||
if len(payload) > available {
|
||||
payload = payload[:available]
|
||||
}
|
||||
newPacket := buf.NewSize(frontHeadroom + header.IPv4MinimumSize + header.ICMPv4MinimumSize + len(payload))
|
||||
defer newPacket.Release()
|
||||
newPacket.Resize(frontHeadroom, header.IPv4MinimumSize+header.ICMPv4MinimumSize+len(payload))
|
||||
newIPHdr := header.IPv4(newPacket.Bytes())
|
||||
newIPHdr.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(newPacket.Len()),
|
||||
Protocol: uint8(header.ICMPv4ProtocolNumber),
|
||||
SrcAddr: ipHdr.DestinationAddr(),
|
||||
DstAddr: ipHdr.SourceAddr(),
|
||||
})
|
||||
newIPHdr.SetChecksum(^newIPHdr.CalculateChecksum())
|
||||
icmpHdr := header.ICMPv4(newIPHdr.Payload())
|
||||
icmpHdr.SetType(header.ICMPv4DstUnreachable)
|
||||
icmpHdr.SetCode(code)
|
||||
icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr[:header.ICMPv4MinimumSize], checksum.Checksum(ipHdr.Payload(), 0)))
|
||||
copy(icmpHdr.Payload(), payload)
|
||||
if PacketOffset > 0 {
|
||||
newPacket.ExtendHeader(PacketOffset)[3] = syscall.AF_INET
|
||||
} else {
|
||||
newPacket.Advance(-s.frontHeadroom)
|
||||
}
|
||||
return common.Error(s.tun.Write(newPacket.Bytes()))
|
||||
}
|
||||
|
||||
func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool, error) {
|
||||
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
|
||||
return false, nil
|
||||
}
|
||||
sourceAddr := ipHdr.SourceAddr()
|
||||
destinationAddr := ipHdr.DestinationAddr()
|
||||
if destinationAddr != s.inet6Address {
|
||||
action, err := s.directNat.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
|
||||
return s.handler.PrepareConnection(
|
||||
N.NetworkICMP,
|
||||
M.SocksaddrFrom(sourceAddr, 0),
|
||||
M.SocksaddrFrom(destinationAddr, 0),
|
||||
&systemICMPDirectPacketWriter6{s.tun, s.frontHeadroom + PacketOffset, sourceAddr},
|
||||
timeout,
|
||||
)
|
||||
})
|
||||
if errors.Is(err, ErrReset) {
|
||||
return false, s.rejectIPv6WithICMP(ipHdr, header.ICMPv6AddressUnreachable)
|
||||
} else if errors.Is(err, ErrDrop) {
|
||||
return false, nil
|
||||
}
|
||||
if action != nil {
|
||||
return false, action.WritePacket(buf.As(ipHdr).ToOwned())
|
||||
}
|
||||
}
|
||||
icmpHdr.SetType(header.ICMPv6EchoReply)
|
||||
sourceAddress := ipHdr.SourceAddr()
|
||||
ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
|
||||
|
|
@ -748,50 +608,6 @@ func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool
|
|||
return true, nil
|
||||
}
|
||||
|
||||
func (s *System) rejectIPv6WithICMP(ipHdr header.IPv6, code header.ICMPv6Code) error {
|
||||
frontHeadroom := s.frontHeadroom + PacketOffset
|
||||
mtu := s.mtu
|
||||
const maxIPv6Data = header.IPv6MinimumMTU - header.IPv6FixedHeaderSize
|
||||
if mtu > maxIPv6Data {
|
||||
mtu = maxIPv6Data
|
||||
}
|
||||
available := mtu - header.ICMPv6ErrorHeaderSize
|
||||
if available < header.IPv6MinimumSize {
|
||||
return nil
|
||||
}
|
||||
payload := ipHdr
|
||||
if len(payload) > available {
|
||||
payload = payload[:available]
|
||||
}
|
||||
newPacket := buf.NewSize(frontHeadroom + header.IPv6MinimumSize + header.ICMPv6DstUnreachableMinimumSize + len(payload))
|
||||
defer newPacket.Release()
|
||||
newPacket.Resize(frontHeadroom, header.IPv6MinimumSize+header.ICMPv6DstUnreachableMinimumSize+len(payload))
|
||||
newIPHdr := header.IPv6(newPacket.Bytes())
|
||||
newIPHdr.Encode(&header.IPv6Fields{
|
||||
PayloadLength: uint16(header.ICMPv6DstUnreachableMinimumSize + len(payload)),
|
||||
TransportProtocol: header.ICMPv6ProtocolNumber,
|
||||
SrcAddr: ipHdr.DestinationAddr(),
|
||||
DstAddr: ipHdr.SourceAddr(),
|
||||
})
|
||||
icmpHdr := header.ICMPv6(newIPHdr.Payload())
|
||||
icmpHdr.SetType(header.ICMPv6DstUnreachable)
|
||||
icmpHdr.SetCode(code)
|
||||
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
|
||||
Header: icmpHdr[:header.ICMPv6DstUnreachableMinimumSize],
|
||||
Src: newIPHdr.SourceAddressSlice(),
|
||||
Dst: newIPHdr.DestinationAddressSlice(),
|
||||
PayloadCsum: checksum.Checksum(payload, 0),
|
||||
PayloadLen: len(payload),
|
||||
}))
|
||||
copy(icmpHdr.Payload(), payload)
|
||||
if PacketOffset > 0 {
|
||||
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
|
||||
} else {
|
||||
newPacket.Advance(-s.frontHeadroom)
|
||||
}
|
||||
return common.Error(s.tun.Write(newPacket.Bytes()))
|
||||
}
|
||||
|
||||
type systemUDPPacketWriter4 struct {
|
||||
tun Tun
|
||||
frontHeadroom int
|
||||
|
|
@ -868,45 +684,37 @@ func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.S
|
|||
return common.Error(w.tun.Write(newPacket.Bytes()))
|
||||
}
|
||||
|
||||
type systemICMPDirectPacketWriter4 struct {
|
||||
type systemWriteback struct {
|
||||
tun Tun
|
||||
linuxTUN LinuxTUN
|
||||
frontHeadroom int
|
||||
source netip.Addr
|
||||
}
|
||||
|
||||
func (w *systemICMPDirectPacketWriter4) WritePacket(p []byte) error {
|
||||
newPacket := buf.NewSize(w.frontHeadroom + len(p))
|
||||
defer newPacket.Release()
|
||||
newPacket.Resize(w.frontHeadroom, 0)
|
||||
newPacket.Write(p)
|
||||
ipHdr := header.IPv4(newPacket.Bytes())
|
||||
ipHdr.SetDestinationAddr(w.source)
|
||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||
if PacketOffset > 0 {
|
||||
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
|
||||
} else {
|
||||
newPacket.Advance(-w.frontHeadroom)
|
||||
func newSystemWriteback(tunInterface Tun, frontHeadroom int) *systemWriteback {
|
||||
writeback := &systemWriteback{tun: tunInterface, frontHeadroom: frontHeadroom}
|
||||
if linuxTUN, isLinuxTUN := tunInterface.(LinuxTUN); isLinuxTUN {
|
||||
writeback.linuxTUN = linuxTUN
|
||||
}
|
||||
return common.Error(w.tun.Write(newPacket.Bytes()))
|
||||
return writeback
|
||||
}
|
||||
|
||||
type systemICMPDirectPacketWriter6 struct {
|
||||
tun Tun
|
||||
frontHeadroom int
|
||||
source netip.Addr
|
||||
func (w *systemWriteback) ReturnHeadroom() int {
|
||||
return w.frontHeadroom + PacketOffset
|
||||
}
|
||||
|
||||
func (w *systemICMPDirectPacketWriter6) WritePacket(p []byte) error {
|
||||
newPacket := buf.NewSize(w.frontHeadroom + len(p))
|
||||
defer newPacket.Release()
|
||||
newPacket.Resize(w.frontHeadroom, 0)
|
||||
newPacket.Write(p)
|
||||
ipHdr := header.IPv6(newPacket.Bytes())
|
||||
ipHdr.SetDestinationAddr(w.source)
|
||||
if PacketOffset > 0 {
|
||||
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
|
||||
} else {
|
||||
newPacket.Advance(-w.frontHeadroom)
|
||||
func (w *systemWriteback) WriteReturnPackets(packets [][]byte) error {
|
||||
if w.linuxTUN != nil {
|
||||
return common.Error(w.linuxTUN.BatchWrite(packets, w.frontHeadroom))
|
||||
}
|
||||
return common.Error(w.tun.Write(newPacket.Bytes()))
|
||||
var writeErrors []error
|
||||
for _, packet := range packets {
|
||||
if PacketOffset > 0 {
|
||||
PacketFillHeader(packet, header.IPVersion(packet[PacketOffset:]))
|
||||
}
|
||||
_, err := w.tun.Write(packet)
|
||||
if err != nil {
|
||||
writeErrors = append(writeErrors, err)
|
||||
}
|
||||
}
|
||||
return E.Errors(writeErrors...)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,9 +5,6 @@ import (
|
|||
"net/netip"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
)
|
||||
|
||||
type TCPNat struct {
|
||||
|
|
@ -85,17 +82,13 @@ func (n *TCPNat) LookupBack(port uint16) *TCPSession {
|
|||
return session
|
||||
}
|
||||
|
||||
func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort, handler Handler) (uint16, error) {
|
||||
func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort) uint16 {
|
||||
key := tcpNatKey{Source: source, Destination: destination}
|
||||
n.addrAccess.RLock()
|
||||
port, loaded := n.addrMap[key]
|
||||
n.addrAccess.RUnlock()
|
||||
if loaded {
|
||||
return port, nil
|
||||
}
|
||||
_, pErr := handler.PrepareConnection(N.NetworkTCP, M.SocksaddrFromNetIP(source), M.SocksaddrFromNetIP(destination), nil, 0)
|
||||
if pErr != nil {
|
||||
return 0, pErr
|
||||
return port
|
||||
}
|
||||
n.addrAccess.Lock()
|
||||
nextPort := n.portIndex
|
||||
|
|
@ -114,5 +107,5 @@ func (n *TCPNat) Lookup(source netip.AddrPort, destination netip.AddrPort, handl
|
|||
LastActive: time.Now(),
|
||||
}
|
||||
n.portAccess.Unlock()
|
||||
return nextPort, nil
|
||||
return nextPort
|
||||
}
|
||||
|
|
|
|||
14
tun.go
14
tun.go
|
|
@ -7,7 +7,6 @@ import (
|
|||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing/common"
|
||||
"github.com/sagernet/sing/common/buf"
|
||||
|
|
@ -15,27 +14,16 @@ import (
|
|||
E "github.com/sagernet/sing/common/exceptions"
|
||||
F "github.com/sagernet/sing/common/format"
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/sagernet/sing/common/ranges"
|
||||
)
|
||||
|
||||
type Handler interface {
|
||||
PrepareConnection(
|
||||
network string,
|
||||
source M.Socksaddr,
|
||||
destination M.Socksaddr,
|
||||
routeContext DirectRouteContext,
|
||||
timeout time.Duration,
|
||||
) (DirectRouteDestination, error)
|
||||
JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort) FlowVerdict
|
||||
N.TCPConnectionHandlerEx
|
||||
N.UDPConnectionHandlerEx
|
||||
}
|
||||
|
||||
type DirectRouteContext interface {
|
||||
WritePacket(packet []byte) error
|
||||
}
|
||||
|
||||
type Tun interface {
|
||||
io.ReadWriter
|
||||
Name() (string, error)
|
||||
|
|
|
|||
|
|
@ -568,8 +568,10 @@ func (t *NativeTun) readNonblocking(buffer []byte) (int, error) {
|
|||
func (t *NativeTun) BatchWrite(buffers [][]byte, offset int) (int, error) {
|
||||
t.writeAccess.Lock()
|
||||
defer func() {
|
||||
t.tcpGROTable.reset()
|
||||
t.udpGROTable.reset()
|
||||
if t.vnetHdr {
|
||||
t.tcpGROTable.reset()
|
||||
t.udpGROTable.reset()
|
||||
}
|
||||
t.writeAccess.Unlock()
|
||||
}()
|
||||
var (
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue