Add flow tracking
This commit is contained in:
parent
ed63adda33
commit
14c8f75f7a
7 changed files with 240 additions and 39 deletions
43
flow.go
43
flow.go
|
|
@ -6,6 +6,7 @@ type FlowVerdict struct {
|
||||||
Action FlowAction
|
Action FlowAction
|
||||||
Port Port
|
Port Port
|
||||||
Destination netip.AddrPort
|
Destination netip.AddrPort
|
||||||
|
NewTracker func() FlowTracker
|
||||||
}
|
}
|
||||||
|
|
||||||
type FlowAction uint8
|
type FlowAction uint8
|
||||||
|
|
@ -18,6 +19,48 @@ const (
|
||||||
ActionBypass
|
ActionBypass
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type FlowTracker interface {
|
||||||
|
AttachFlow(handle FlowHandle)
|
||||||
|
CountForward(n int)
|
||||||
|
CountReverse(n int)
|
||||||
|
FlowEstablished()
|
||||||
|
CloseFlow(reason FlowCloseReason)
|
||||||
|
}
|
||||||
|
|
||||||
|
type FlowHandle interface {
|
||||||
|
CloseFlow()
|
||||||
|
}
|
||||||
|
|
||||||
|
type FlowCloseReason uint8
|
||||||
|
|
||||||
|
const (
|
||||||
|
FlowCloseReset FlowCloseReason = iota
|
||||||
|
FlowCloseFinished
|
||||||
|
FlowCloseTimeout
|
||||||
|
FlowCloseEvicted
|
||||||
|
FlowCloseShutdown
|
||||||
|
FlowCloseInterrupted
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r FlowCloseReason) String() string {
|
||||||
|
switch r {
|
||||||
|
case FlowCloseReset:
|
||||||
|
return "connection reset"
|
||||||
|
case FlowCloseFinished:
|
||||||
|
return "finished"
|
||||||
|
case FlowCloseTimeout:
|
||||||
|
return "idle timeout"
|
||||||
|
case FlowCloseEvicted:
|
||||||
|
return "evicted"
|
||||||
|
case FlowCloseShutdown:
|
||||||
|
return "stack closed"
|
||||||
|
case FlowCloseInterrupted:
|
||||||
|
return "interrupted"
|
||||||
|
default:
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
type Port interface {
|
type Port interface {
|
||||||
PortAddresses() (v4 netip.Addr, v6 netip.Addr)
|
PortAddresses() (v4 netip.Addr, v6 netip.Addr)
|
||||||
PortMTU() uint32
|
PortMTU() uint32
|
||||||
|
|
|
||||||
|
|
@ -14,6 +14,7 @@ import (
|
||||||
const (
|
const (
|
||||||
tcpEstablishedTimeout = 2*time.Hour + 4*time.Minute
|
tcpEstablishedTimeout = 2*time.Hour + 4*time.Minute
|
||||||
tcpTransitoryTimeout = 4 * time.Minute
|
tcpTransitoryTimeout = 4 * time.Minute
|
||||||
|
tcpClosingTimeout = 10 * time.Second
|
||||||
|
|
||||||
defaultUDPTimeout = 5 * time.Minute
|
defaultUDPTimeout = 5 * time.Minute
|
||||||
|
|
||||||
|
|
@ -46,6 +47,7 @@ type forwardFlow struct {
|
||||||
reverseRule rewriteRule
|
reverseRule rewriteRule
|
||||||
effectiveMTU uint32
|
effectiveMTU uint32
|
||||||
protocol uint8
|
protocol uint8
|
||||||
|
tracker FlowTracker
|
||||||
|
|
||||||
clientAddress netip.Addr
|
clientAddress netip.Addr
|
||||||
clientSelector uint16
|
clientSelector uint16
|
||||||
|
|
@ -62,14 +64,29 @@ type forwardFlow struct {
|
||||||
lastReverse atomic.Int64
|
lastReverse atomic.Int64
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *forwardFlow) close(reason FlowCloseReason) {
|
||||||
|
if !f.closed.CompareAndSwap(false, true) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if f.tracker != nil {
|
||||||
|
f.tracker.CloseFlow(reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *forwardFlow) CloseFlow() {
|
||||||
|
f.close(FlowCloseInterrupted)
|
||||||
|
}
|
||||||
|
|
||||||
func (f *forwardFlow) observeReverse(packet *forwardPacket, now int64) {
|
func (f *forwardFlow) observeReverse(packet *forwardPacket, now int64) {
|
||||||
f.lastReverse.Store(now)
|
f.lastReverse.Store(now)
|
||||||
if packet.protocol != uint8(header.TCPProtocolNumber) {
|
if packet.protocol != uint8(header.TCPProtocolNumber) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
f.established.Store(true)
|
if f.established.CompareAndSwap(false, true) && f.tracker != nil {
|
||||||
|
f.tracker.FlowEstablished()
|
||||||
|
}
|
||||||
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||||
f.closed.Store(true)
|
f.close(FlowCloseReset)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if packet.tcpFlags&header.TCPFlagFin != 0 {
|
if packet.tcpFlags&header.TCPFlagFin != 0 {
|
||||||
|
|
@ -128,6 +145,11 @@ func (d *ForwardDispatcher) Close() {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
d.returnPath.closed.Store(true)
|
d.returnPath.closed.Store(true)
|
||||||
|
for _, entry := range d.table {
|
||||||
|
if entry.flow != nil {
|
||||||
|
entry.flow.close(FlowCloseShutdown)
|
||||||
|
}
|
||||||
|
}
|
||||||
for port, nat := range d.ports {
|
for port, nat := range d.ports {
|
||||||
if nat != nil {
|
if nat != nil {
|
||||||
port.DetachReturn(&d.returnPath)
|
port.DetachReturn(&d.returnPath)
|
||||||
|
|
@ -147,7 +169,7 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
|
||||||
now := d.now()
|
now := d.now()
|
||||||
entry, loaded := d.table[key]
|
entry, loaded := d.table[key]
|
||||||
if loaded && d.entryExpired(entry, now) {
|
if loaded && d.entryExpired(entry, now) {
|
||||||
d.removeEntry(key, entry)
|
d.removeEntry(key, entry, FlowCloseTimeout)
|
||||||
loaded = false
|
loaded = false
|
||||||
}
|
}
|
||||||
if loaded {
|
if loaded {
|
||||||
|
|
@ -171,7 +193,7 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
|
||||||
if packet.protocol == uint8(header.TCPProtocolNumber) {
|
if packet.protocol == uint8(header.TCPProtocolNumber) {
|
||||||
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||||
d.forwardToPort(flow, packet, raw)
|
d.forwardToPort(flow, packet, raw)
|
||||||
flow.closed.Store(true)
|
flow.close(FlowCloseReset)
|
||||||
d.tombstoneEntry(entry, now)
|
d.tombstoneEntry(entry, now)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
@ -186,7 +208,7 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
|
||||||
case ActionAccept:
|
case ActionAccept:
|
||||||
entry.deadline = now + int64(entry.idle)
|
entry.deadline = now + int64(entry.idle)
|
||||||
if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 {
|
if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||||
d.removeEntry(key, entry)
|
d.removeEntry(key, entry, FlowCloseReset)
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
case ActionReject:
|
case ActionReject:
|
||||||
|
|
@ -200,7 +222,11 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool {
|
func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool {
|
||||||
verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination)
|
var firstPacket []byte
|
||||||
|
if packet.protocol == uint8(header.UDPProtocolNumber) {
|
||||||
|
firstPacket = header.UDP(packet.transport).Payload()
|
||||||
|
}
|
||||||
|
verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination, firstPacket)
|
||||||
switch verdict.Action {
|
switch verdict.Action {
|
||||||
case ActionFlow:
|
case ActionFlow:
|
||||||
if verdict.Port != nil {
|
if verdict.Port != nil {
|
||||||
|
|
@ -249,6 +275,9 @@ func (d *ForwardDispatcher) idleTimeout(protocol uint8, established bool) time.D
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *ForwardDispatcher) flowIdle(flow *forwardFlow) time.Duration {
|
func (d *ForwardDispatcher) flowIdle(flow *forwardFlow) time.Duration {
|
||||||
|
if flow.protocol == uint8(header.TCPProtocolNumber) && flow.finForward && flow.finReverse.Load() {
|
||||||
|
return tcpClosingTimeout
|
||||||
|
}
|
||||||
established := flow.established.Load() && !flow.finForward && !flow.finReverse.Load()
|
established := flow.established.Load() && !flow.finForward && !flow.finReverse.Load()
|
||||||
return d.idleTimeout(flow.protocol, established)
|
return d.idleTimeout(flow.protocol, established)
|
||||||
}
|
}
|
||||||
|
|
@ -324,6 +353,12 @@ func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdic
|
||||||
flow.reverseRule.sourcePort = clientDestinationPort
|
flow.reverseRule.sourcePort = clientDestinationPort
|
||||||
flow.reverseRule.rewriteSourcePort = true
|
flow.reverseRule.rewriteSourcePort = true
|
||||||
}
|
}
|
||||||
|
if verdict.NewTracker != nil {
|
||||||
|
flow.tracker = verdict.NewTracker()
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.AttachFlow(flow)
|
||||||
|
}
|
||||||
|
}
|
||||||
nat.insert(reverseKey, flow)
|
nat.insert(reverseKey, flow)
|
||||||
return flow, true
|
return flow, true
|
||||||
}
|
}
|
||||||
|
|
@ -354,6 +389,9 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
|
||||||
func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPacket, raw []byte) {
|
func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPacket, raw []byte) {
|
||||||
if flow.effectiveMTU != 0 && uint32(len(raw)) > flow.effectiveMTU {
|
if flow.effectiveMTU != 0 && uint32(len(raw)) > flow.effectiveMTU {
|
||||||
if packet.protocol == uint8(header.TCPProtocolNumber) {
|
if packet.protocol == uint8(header.TCPProtocolNumber) {
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.CountForward(len(raw))
|
||||||
|
}
|
||||||
d.rewriteForward(flow, packet)
|
d.rewriteForward(flow, packet)
|
||||||
d.resegmentTCP(flow, packet, raw)
|
d.resegmentTCP(flow, packet, raw)
|
||||||
return
|
return
|
||||||
|
|
@ -361,6 +399,9 @@ func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPack
|
||||||
if packet.ipVersion == 4 {
|
if packet.ipVersion == 4 {
|
||||||
ipHdr := packet.network.(header.IPv4)
|
ipHdr := packet.network.(header.IPv4)
|
||||||
if ipHdr.Flags()&header.IPv4FlagDontFragment == 0 {
|
if ipHdr.Flags()&header.IPv4FlagDontFragment == 0 {
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.CountForward(len(raw))
|
||||||
|
}
|
||||||
d.rewriteForward(flow, packet)
|
d.rewriteForward(flow, packet)
|
||||||
fragments, ok := fragmentIPv4Packet(ipHdr, flow.effectiveMTU)
|
fragments, ok := fragmentIPv4Packet(ipHdr, flow.effectiveMTU)
|
||||||
if ok {
|
if ok {
|
||||||
|
|
@ -382,6 +423,9 @@ func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPack
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.CountForward(len(raw))
|
||||||
|
}
|
||||||
d.rewriteForward(flow, packet)
|
d.rewriteForward(flow, packet)
|
||||||
d.stagePort(flow.nat, raw)
|
d.stagePort(flow.nat, raw)
|
||||||
}
|
}
|
||||||
|
|
@ -460,10 +504,13 @@ func (d *ForwardDispatcher) tombstoneEntry(entry *flowEntry, now int64) {
|
||||||
entry.deadline = now + int64(entry.idle)
|
entry.deadline = now + int64(entry.idle)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *ForwardDispatcher) removeEntry(key flowKey, entry *flowEntry) {
|
func (d *ForwardDispatcher) removeEntry(key flowKey, entry *flowEntry, reason FlowCloseReason) {
|
||||||
delete(d.table, key)
|
delete(d.table, key)
|
||||||
if entry.flow != nil {
|
if entry.flow != nil {
|
||||||
entry.flow.closed.Store(true)
|
if reason == FlowCloseTimeout && entry.flow.finForward && entry.flow.finReverse.Load() {
|
||||||
|
reason = FlowCloseFinished
|
||||||
|
}
|
||||||
|
entry.flow.close(reason)
|
||||||
entry.flow.nat.delete(entry.flow.reverseKey)
|
entry.flow.nat.delete(entry.flow.reverseKey)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -484,7 +531,7 @@ func (d *ForwardDispatcher) evictEntries(now int64) {
|
||||||
)
|
)
|
||||||
for key, entry := range d.table {
|
for key, entry := range d.table {
|
||||||
if d.entryExpired(entry, now) {
|
if d.entryExpired(entry, now) {
|
||||||
d.removeEntry(key, entry)
|
d.removeEntry(key, entry, FlowCloseTimeout)
|
||||||
freed++
|
freed++
|
||||||
} else if oldest == nil || entry.deadline < oldest.deadline {
|
} else if oldest == nil || entry.deadline < oldest.deadline {
|
||||||
oldestKey = key
|
oldestKey = key
|
||||||
|
|
@ -496,7 +543,7 @@ func (d *ForwardDispatcher) evictEntries(now int64) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if freed == 0 && oldest != nil {
|
if freed == 0 && oldest != nil {
|
||||||
d.removeEntry(oldestKey, oldest)
|
d.removeEntry(oldestKey, oldest, FlowCloseEvicted)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -510,7 +557,7 @@ func (d *ForwardDispatcher) maybeSweep(now int64) {
|
||||||
if entry.action == ActionFlow && entry.flow.closed.Load() {
|
if entry.action == ActionFlow && entry.flow.closed.Load() {
|
||||||
d.tombstoneEntry(entry, now)
|
d.tombstoneEntry(entry, now)
|
||||||
} else if d.entryExpired(entry, now) {
|
} else if d.entryExpired(entry, now) {
|
||||||
d.removeEntry(key, entry)
|
d.removeEntry(key, entry, FlowCloseTimeout)
|
||||||
}
|
}
|
||||||
visited++
|
visited++
|
||||||
if visited >= flowSweepLimit {
|
if visited >= flowSweepLimit {
|
||||||
|
|
@ -579,6 +626,9 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
||||||
if flow.closed.Load() {
|
if flow.closed.Load() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.CountReverse(len(raw) - headroom)
|
||||||
|
}
|
||||||
flow.observeReverse(&parsed, now)
|
flow.observeReverse(&parsed, now)
|
||||||
if parsed.isTCPSyn() {
|
if parsed.isTCPSyn() {
|
||||||
applyRewriteRaw(&parsed, &flow.reverseRule)
|
applyRewriteRaw(&parsed, &flow.reverseRule)
|
||||||
|
|
|
||||||
|
|
@ -269,6 +269,7 @@ func (h *nfqueueHandler) handlePacket(attr nfqueue.Attribute) int {
|
||||||
packet.protocol,
|
packet.protocol,
|
||||||
packet.source,
|
packet.source,
|
||||||
packet.destination,
|
packet.destination,
|
||||||
|
packet.firstPacket,
|
||||||
)
|
)
|
||||||
|
|
||||||
// Use NfRepeat for bypass/reset so the packet re-enters the chain
|
// Use NfRepeat for bypass/reset so the packet re-enters the chain
|
||||||
|
|
|
||||||
|
|
@ -28,7 +28,7 @@ type ICMPForwarder struct {
|
||||||
returnPath icmpForwarderReturn
|
returnPath icmpForwarderReturn
|
||||||
|
|
||||||
flowAccess sync.Mutex
|
flowAccess sync.Mutex
|
||||||
flows map[icmpFlowKey]time.Time
|
flows map[icmpFlowKey]*icmpFlow
|
||||||
lastSweep time.Time
|
lastSweep time.Time
|
||||||
attachedPorts map[Port]bool
|
attachedPorts map[Port]bool
|
||||||
}
|
}
|
||||||
|
|
@ -40,12 +40,32 @@ type icmpFlowKey struct {
|
||||||
identifier uint16
|
identifier uint16
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type icmpFlow struct {
|
||||||
|
port Port
|
||||||
|
tracker FlowTracker
|
||||||
|
deadline time.Time
|
||||||
|
closed atomic.Bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *icmpFlow) close(reason FlowCloseReason) {
|
||||||
|
if !f.closed.CompareAndSwap(false, true) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if f.tracker != nil {
|
||||||
|
f.tracker.CloseFlow(reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *icmpFlow) CloseFlow() {
|
||||||
|
f.close(FlowCloseInterrupted)
|
||||||
|
}
|
||||||
|
|
||||||
func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) *ICMPForwarder {
|
func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) *ICMPForwarder {
|
||||||
forwarder := &ICMPForwarder{
|
forwarder := &ICMPForwarder{
|
||||||
stack: stack,
|
stack: stack,
|
||||||
handler: handler,
|
handler: handler,
|
||||||
logger: logger,
|
logger: logger,
|
||||||
flows: make(map[icmpFlowKey]time.Time),
|
flows: make(map[icmpFlowKey]*icmpFlow),
|
||||||
attachedPorts: make(map[Port]bool),
|
attachedPorts: make(map[Port]bool),
|
||||||
}
|
}
|
||||||
forwarder.returnPath.forwarder = forwarder
|
forwarder.returnPath.forwarder = forwarder
|
||||||
|
|
@ -56,6 +76,10 @@ func (f *ICMPForwarder) Close() error {
|
||||||
f.returnPath.closed.Store(true)
|
f.returnPath.closed.Store(true)
|
||||||
f.flowAccess.Lock()
|
f.flowAccess.Lock()
|
||||||
defer f.flowAccess.Unlock()
|
defer f.flowAccess.Unlock()
|
||||||
|
for key, flow := range f.flows {
|
||||||
|
flow.close(FlowCloseShutdown)
|
||||||
|
delete(f.flows, key)
|
||||||
|
}
|
||||||
for port := range f.attachedPorts {
|
for port := range f.attachedPorts {
|
||||||
port.DetachReturn(&f.returnPath)
|
port.DetachReturn(&f.returnPath)
|
||||||
delete(f.attachedPorts, port)
|
delete(f.attachedPorts, port)
|
||||||
|
|
@ -71,16 +95,25 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
identifier := icmpHdr.Ident()
|
identifier := icmpHdr.Ident()
|
||||||
|
key := icmpFlowKey{
|
||||||
|
source: AddrFromAddress(ipHdr.SourceAddress()),
|
||||||
|
destination: AddrFromAddress(ipHdr.DestinationAddress()),
|
||||||
|
identifier: identifier,
|
||||||
|
}
|
||||||
|
if f.forwardCached(key, pkt) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
verdict := f.handler.JudgeFlow(
|
verdict := f.handler.JudgeFlow(
|
||||||
uint8(header.ICMPv4ProtocolNumber),
|
uint8(header.ICMPv4ProtocolNumber),
|
||||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
|
netip.AddrPortFrom(key.source, identifier),
|
||||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
|
netip.AddrPortFrom(key.destination, identifier),
|
||||||
|
nil,
|
||||||
)
|
)
|
||||||
switch verdict.Action {
|
switch verdict.Action {
|
||||||
case ActionReject, ActionDrop:
|
case ActionReject, ActionDrop:
|
||||||
return true
|
return true
|
||||||
case ActionFlow:
|
case ActionFlow:
|
||||||
if f.forwardFlow(verdict.Port, false, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
|
if f.installFlow(key, verdict, pkt) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -117,16 +150,26 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
identifier := icmpHdr.Ident()
|
identifier := icmpHdr.Ident()
|
||||||
|
key := icmpFlowKey{
|
||||||
|
v6: true,
|
||||||
|
source: AddrFromAddress(ipHdr.SourceAddress()),
|
||||||
|
destination: AddrFromAddress(ipHdr.DestinationAddress()),
|
||||||
|
identifier: identifier,
|
||||||
|
}
|
||||||
|
if f.forwardCached(key, pkt) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
verdict := f.handler.JudgeFlow(
|
verdict := f.handler.JudgeFlow(
|
||||||
uint8(header.ICMPv6ProtocolNumber),
|
uint8(header.ICMPv6ProtocolNumber),
|
||||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
|
netip.AddrPortFrom(key.source, identifier),
|
||||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
|
netip.AddrPortFrom(key.destination, identifier),
|
||||||
|
nil,
|
||||||
)
|
)
|
||||||
switch verdict.Action {
|
switch verdict.Action {
|
||||||
case ActionReject, ActionDrop:
|
case ActionReject, ActionDrop:
|
||||||
return true
|
return true
|
||||||
case ActionFlow:
|
case ActionFlow:
|
||||||
if f.forwardFlow(verdict.Port, true, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
|
if f.installFlow(key, verdict, pkt) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -163,13 +206,38 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, destination netip.Addr, identifier uint16, pkt *stack.PacketBuffer) bool {
|
func (f *ICMPForwarder) forwardCached(key icmpFlowKey, pkt *stack.PacketBuffer) bool {
|
||||||
|
now := time.Now()
|
||||||
|
f.flowAccess.Lock()
|
||||||
|
flow, loaded := f.flows[key]
|
||||||
|
if loaded {
|
||||||
|
if flow.closed.Load() {
|
||||||
|
delete(f.flows, key)
|
||||||
|
loaded = false
|
||||||
|
} else if now.After(flow.deadline) {
|
||||||
|
delete(f.flows, key)
|
||||||
|
flow.close(FlowCloseTimeout)
|
||||||
|
loaded = false
|
||||||
|
} else {
|
||||||
|
flow.deadline = now.Add(defaultICMPTimeout)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
f.flowAccess.Unlock()
|
||||||
|
if !loaded {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
f.writeToPort(flow, pkt)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *ICMPForwarder) installFlow(key icmpFlowKey, verdict FlowVerdict, pkt *stack.PacketBuffer) bool {
|
||||||
|
port := verdict.Port
|
||||||
if port == nil {
|
if port == nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
inet4Address, inet6Address := port.PortAddresses()
|
inet4Address, inet6Address := port.PortAddresses()
|
||||||
portAddress := inet4Address
|
portAddress := inet4Address
|
||||||
if v6 {
|
if key.v6 {
|
||||||
portAddress = inet6Address
|
portAddress = inet6Address
|
||||||
}
|
}
|
||||||
if !portAddress.IsValid() || !portAddress.IsUnspecified() {
|
if !portAddress.IsValid() || !portAddress.IsUnspecified() {
|
||||||
|
|
@ -188,14 +256,29 @@ func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, desti
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
if now.Sub(f.lastSweep) >= defaultICMPTimeout {
|
if now.Sub(f.lastSweep) >= defaultICMPTimeout {
|
||||||
f.lastSweep = now
|
f.lastSweep = now
|
||||||
for key, deadline := range f.flows {
|
for flowKey, cachedFlow := range f.flows {
|
||||||
if now.After(deadline) {
|
if cachedFlow.closed.Load() {
|
||||||
delete(f.flows, key)
|
delete(f.flows, flowKey)
|
||||||
|
} else if now.After(cachedFlow.deadline) {
|
||||||
|
delete(f.flows, flowKey)
|
||||||
|
cachedFlow.close(FlowCloseTimeout)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
f.flows[icmpFlowKey{v6: v6, source: source, destination: destination, identifier: identifier}] = now.Add(defaultICMPTimeout)
|
flow := &icmpFlow{port: port, deadline: now.Add(defaultICMPTimeout)}
|
||||||
|
if verdict.NewTracker != nil {
|
||||||
|
flow.tracker = verdict.NewTracker()
|
||||||
|
}
|
||||||
|
f.flows[key] = flow
|
||||||
f.flowAccess.Unlock()
|
f.flowAccess.Unlock()
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.AttachFlow(flow)
|
||||||
|
}
|
||||||
|
f.writeToPort(flow, pkt)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *ICMPForwarder) writeToPort(flow *icmpFlow, pkt *stack.PacketBuffer) {
|
||||||
networkSlice := pkt.NetworkHeader().Slice()
|
networkSlice := pkt.NetworkHeader().Slice()
|
||||||
transportSlice := pkt.TransportHeader().Slice()
|
transportSlice := pkt.TransportHeader().Slice()
|
||||||
dataSlice := pkt.Data().AsRange().ToSlice()
|
dataSlice := pkt.Data().AsRange().ToSlice()
|
||||||
|
|
@ -203,27 +286,34 @@ func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, desti
|
||||||
packetSlice = append(packetSlice, networkSlice...)
|
packetSlice = append(packetSlice, networkSlice...)
|
||||||
packetSlice = append(packetSlice, transportSlice...)
|
packetSlice = append(packetSlice, transportSlice...)
|
||||||
packetSlice = append(packetSlice, dataSlice...)
|
packetSlice = append(packetSlice, dataSlice...)
|
||||||
err := port.WritePackets([][]byte{packetSlice})
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.CountForward(len(packetSlice))
|
||||||
|
}
|
||||||
|
err := flow.port.WritePackets([][]byte{packetSlice})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.logger.Trace(E.Cause(err, "forward ICMP packet"))
|
f.logger.Trace(E.Cause(err, "forward ICMP packet"))
|
||||||
}
|
}
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *ICMPForwarder) lookupFlow(key icmpFlowKey) bool {
|
func (f *ICMPForwarder) lookupFlow(key icmpFlowKey) *icmpFlow {
|
||||||
f.flowAccess.Lock()
|
f.flowAccess.Lock()
|
||||||
defer f.flowAccess.Unlock()
|
defer f.flowAccess.Unlock()
|
||||||
deadline, loaded := f.flows[key]
|
flow, loaded := f.flows[key]
|
||||||
if !loaded {
|
if !loaded {
|
||||||
return false
|
return nil
|
||||||
|
}
|
||||||
|
if flow.closed.Load() {
|
||||||
|
delete(f.flows, key)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
if now.After(deadline) {
|
if now.After(flow.deadline) {
|
||||||
delete(f.flows, key)
|
delete(f.flows, key)
|
||||||
return false
|
flow.close(FlowCloseTimeout)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
f.flows[key] = now.Add(defaultICMPTimeout)
|
flow.deadline = now.Add(defaultICMPTimeout)
|
||||||
return true
|
return flow
|
||||||
}
|
}
|
||||||
|
|
||||||
type icmpForwarderReturn struct {
|
type icmpForwarderReturn struct {
|
||||||
|
|
@ -289,9 +379,13 @@ func (f *ICMPForwarder) returnPacket(packet []byte) bool {
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
if !f.lookupFlow(key) {
|
flow := f.lookupFlow(key)
|
||||||
|
if flow == nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.CountReverse(len(packet))
|
||||||
|
}
|
||||||
return f.writeBack(packet, header.IPv4ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
return f.writeBack(packet, header.IPv4ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
||||||
case header.IPv6Version:
|
case header.IPv6Version:
|
||||||
ipHdr := header.IPv6(packet)
|
ipHdr := header.IPv6(packet)
|
||||||
|
|
@ -308,9 +402,13 @@ func (f *ICMPForwarder) returnPacket(packet []byte) bool {
|
||||||
destination: AddrFromAddress(ipHdr.SourceAddress()),
|
destination: AddrFromAddress(ipHdr.SourceAddress()),
|
||||||
identifier: icmpHdr.Ident(),
|
identifier: icmpHdr.Ident(),
|
||||||
}
|
}
|
||||||
if !f.lookupFlow(key) {
|
flow := f.lookupFlow(key)
|
||||||
|
if flow == nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
if flow.tracker != nil {
|
||||||
|
flow.tracker.CountReverse(len(packet))
|
||||||
|
}
|
||||||
return f.writeBack(packet, header.IPv6ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
return f.writeBack(packet, header.IPv6ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
|
|
|
||||||
|
|
@ -77,7 +77,7 @@ func (f *TCPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac
|
||||||
func (f *TCPForwarder) Forward(r *tcp.ForwarderRequest) {
|
func (f *TCPForwarder) Forward(r *tcp.ForwarderRequest) {
|
||||||
source := M.SocksaddrFrom(AddrFromAddress(r.ID().RemoteAddress), r.ID().RemotePort)
|
source := M.SocksaddrFrom(AddrFromAddress(r.ID().RemoteAddress), r.ID().RemotePort)
|
||||||
destination := M.SocksaddrFrom(AddrFromAddress(r.ID().LocalAddress), r.ID().LocalPort)
|
destination := M.SocksaddrFrom(AddrFromAddress(r.ID().LocalAddress), r.ID().LocalPort)
|
||||||
switch f.handler.JudgeFlow(uint8(header.TCPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action {
|
switch f.handler.JudgeFlow(uint8(header.TCPProtocolNumber), source.AddrPort(), destination.AddrPort(), nil).Action {
|
||||||
case ActionReject:
|
case ActionReject:
|
||||||
r.Complete(true)
|
r.Complete(true)
|
||||||
return
|
return
|
||||||
|
|
|
||||||
|
|
@ -57,7 +57,16 @@ func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac
|
||||||
func rangeIterate(r stack.Range, fn func(*buffer.View))
|
func rangeIterate(r stack.Range, fn func(*buffer.View))
|
||||||
|
|
||||||
func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) {
|
func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) {
|
||||||
switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action {
|
firstPacketBuffer := userData.(*stack.PacketBuffer)
|
||||||
|
var firstPacket []byte
|
||||||
|
rangeIterate(firstPacketBuffer.Data().AsRange(), func(view *buffer.View) {
|
||||||
|
if firstPacket == nil {
|
||||||
|
firstPacket = view.AsSlice()
|
||||||
|
} else {
|
||||||
|
firstPacket = append(firstPacket[:len(firstPacket):len(firstPacket)], view.AsSlice()...)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort(), firstPacket).Action {
|
||||||
case ActionReject:
|
case ActionReject:
|
||||||
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
|
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
|
||||||
return false, nil, nil, nil
|
return false, nil, nil, nil
|
||||||
|
|
|
||||||
2
tun.go
2
tun.go
|
|
@ -19,7 +19,7 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
type Handler interface {
|
type Handler interface {
|
||||||
JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort) FlowVerdict
|
JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort, firstPacket []byte) FlowVerdict
|
||||||
N.TCPConnectionHandlerEx
|
N.TCPConnectionHandlerEx
|
||||||
N.UDPConnectionHandlerEx
|
N.UDPConnectionHandlerEx
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue