From ed63adda337c00f01579cbffd7f0d7937a19f5a6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Mon, 6 Jul 2026 11:49:18 +0800 Subject: [PATCH] Add flow dispatcher --- flow.go | 32 ++ flow_dispatch.go | 631 +++++++++++++++++++++++++++++++++++ flow_mtu.go | 159 +++++++++ flow_nat.go | 106 ++++++ flow_parse.go | 256 ++++++++++++++ flow_reject.go | 210 ++++++++++++ flow_rewrite.go | 350 +++++++++++++++++++ nfqueue_linux.go | 52 +-- ping/destination.go | 27 +- ping/destination_gvisor.go | 143 -------- ping/destination_rewriter.go | 79 ----- ping/port.go | 194 +++++++++++ ping/source_rewriter.go | 150 --------- redirect_linux.go | 34 +- route_direct.go | 61 ---- stack.go | 6 - stack_gvisor.go | 57 +++- stack_gvisor_filter.go | 69 +++- stack_gvisor_icmp.go | 366 ++++++++++++-------- stack_gvisor_lazy.go | 3 +- stack_gvisor_tcp.go | 11 +- stack_gvisor_udp.go | 11 +- stack_mixed.go | 20 +- stack_system.go | 372 +++++---------------- stack_system_nat.go | 13 +- tun.go | 14 +- tun_linux.go | 6 +- 27 files changed, 2469 insertions(+), 963 deletions(-) create mode 100644 flow.go create mode 100644 flow_dispatch.go create mode 100644 flow_mtu.go create mode 100644 flow_nat.go create mode 100644 flow_parse.go create mode 100644 flow_reject.go create mode 100644 flow_rewrite.go delete mode 100644 ping/destination_gvisor.go delete mode 100644 ping/destination_rewriter.go create mode 100644 ping/port.go delete mode 100644 ping/source_rewriter.go delete mode 100644 route_direct.go diff --git a/flow.go b/flow.go new file mode 100644 index 0000000..a920e93 --- /dev/null +++ b/flow.go @@ -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 +} diff --git a/flow_dispatch.go b/flow_dispatch.go new file mode 100644 index 0000000..d9ec9a3 --- /dev/null +++ b/flow_dispatch.go @@ -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 +} diff --git a/flow_mtu.go b/flow_mtu.go new file mode 100644 index 0000000..26c675b --- /dev/null +++ b/flow_mtu.go @@ -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 +} diff --git a/flow_nat.go b/flow_nat.go new file mode 100644 index 0000000..fe4ea1a --- /dev/null +++ b/flow_nat.go @@ -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 +} diff --git a/flow_parse.go b/flow_parse.go new file mode 100644 index 0000000..4588799 --- /dev/null +++ b/flow_parse.go @@ -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 +} diff --git a/flow_reject.go b/flow_reject.go new file mode 100644 index 0000000..af82eb5 --- /dev/null +++ b/flow_reject.go @@ -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 + } +} diff --git a/flow_rewrite.go b/flow_rewrite.go new file mode 100644 index 0000000..e4d7967 --- /dev/null +++ b/flow_rewrite.go @@ -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) + } + } +} diff --git a/nfqueue_linux.go b/nfqueue_linux.go index 7c2114c..8513aa6 100644 --- a/nfqueue_linux.go +++ b/nfqueue_linux.go @@ -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) diff --git a/ping/destination.go b/ping/destination.go index ece45b9..1e5d533 100644 --- a/ping/destination.go +++ b/ping/destination.go @@ -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")) } diff --git a/ping/destination_gvisor.go b/ping/destination_gvisor.go deleted file mode 100644 index 0c508ee..0000000 --- a/ping/destination_gvisor.go +++ /dev/null @@ -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 -} diff --git a/ping/destination_rewriter.go b/ping/destination_rewriter.go deleted file mode 100644 index 26bb355..0000000 --- a/ping/destination_rewriter.go +++ /dev/null @@ -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) -} diff --git a/ping/port.go b/ping/port.go new file mode 100644 index 0000000..e0cc2a1 --- /dev/null +++ b/ping/port.go @@ -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...) +} diff --git a/ping/source_rewriter.go b/ping/source_rewriter.go deleted file mode 100644 index 545560d..0000000 --- a/ping/source_rewriter.go +++ /dev/null @@ -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) -} diff --git a/redirect_linux.go b/redirect_linux.go index ff89719..04a1fee 100644 --- a/redirect_linux.go +++ b/redirect_linux.go @@ -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() diff --git a/route_direct.go b/route_direct.go deleted file mode 100644 index 444eb5e..0000000 --- a/route_direct.go +++ /dev/null @@ -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 -} diff --git a/stack.go b/stack.go index d128feb..eaf2405 100644 --- a/stack.go +++ b/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 diff --git a/stack_gvisor.go b/stack_gvisor.go index 92e84ba..03b2873 100644 --- a/stack_gvisor.go +++ b/stack_gvisor.go @@ -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()) diff --git a/stack_gvisor_filter.go b/stack_gvisor_filter.go index 18e46e8..abdbcaa 100644 --- a/stack_gvisor_filter.go +++ b/stack_gvisor_filter.go @@ -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) +} diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 1f11bbd..979f4ad 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -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 } diff --git a/stack_gvisor_lazy.go b/stack_gvisor_lazy.go index dcbcafb..258392d 100644 --- a/stack_gvisor_lazy.go +++ b/stack_gvisor_lazy.go @@ -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 diff --git a/stack_gvisor_tcp.go b/stack_gvisor_tcp.go index ba8af6d..371c480 100644 --- a/stack_gvisor_tcp.go +++ b/stack_gvisor_tcp.go @@ -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{ diff --git a/stack_gvisor_udp.go b/stack_gvisor_udp.go index 0c96213..fb31ba9 100644 --- a/stack_gvisor_udp.go +++ b/stack_gvisor_udp.go @@ -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 diff --git a/stack_mixed.go b/stack_mixed.go index 8f38c4b..4680380 100644 --- a/stack_mixed.go +++ b/stack_mixed.go @@ -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()) diff --git a/stack_system.go b/stack_system.go index 181db37..8ab9be6 100644 --- a/stack_system.go +++ b/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...) } diff --git a/stack_system_nat.go b/stack_system_nat.go index cc46017..fd0e382 100644 --- a/stack_system_nat.go +++ b/stack_system_nat.go @@ -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 } diff --git a/tun.go b/tun.go index 4a01ee5..7d10c75 100644 --- a/tun.go +++ b/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) diff --git a/tun_linux.go b/tun_linux.go index a41cbff..487051d 100644 --- a/tun_linux.go +++ b/tun_linux.go @@ -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 (