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
|
||||
Port Port
|
||||
Destination netip.AddrPort
|
||||
NewTracker func() FlowTracker
|
||||
}
|
||||
|
||||
type FlowAction uint8
|
||||
|
|
@ -18,6 +19,48 @@ const (
|
|||
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 {
|
||||
PortAddresses() (v4 netip.Addr, v6 netip.Addr)
|
||||
PortMTU() uint32
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import (
|
|||
const (
|
||||
tcpEstablishedTimeout = 2*time.Hour + 4*time.Minute
|
||||
tcpTransitoryTimeout = 4 * time.Minute
|
||||
tcpClosingTimeout = 10 * time.Second
|
||||
|
||||
defaultUDPTimeout = 5 * time.Minute
|
||||
|
||||
|
|
@ -46,6 +47,7 @@ type forwardFlow struct {
|
|||
reverseRule rewriteRule
|
||||
effectiveMTU uint32
|
||||
protocol uint8
|
||||
tracker FlowTracker
|
||||
|
||||
clientAddress netip.Addr
|
||||
clientSelector uint16
|
||||
|
|
@ -62,14 +64,29 @@ type forwardFlow struct {
|
|||
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) {
|
||||
f.lastReverse.Store(now)
|
||||
if packet.protocol != uint8(header.TCPProtocolNumber) {
|
||||
return
|
||||
}
|
||||
f.established.Store(true)
|
||||
if f.established.CompareAndSwap(false, true) && f.tracker != nil {
|
||||
f.tracker.FlowEstablished()
|
||||
}
|
||||
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||
f.closed.Store(true)
|
||||
f.close(FlowCloseReset)
|
||||
return
|
||||
}
|
||||
if packet.tcpFlags&header.TCPFlagFin != 0 {
|
||||
|
|
@ -128,6 +145,11 @@ func (d *ForwardDispatcher) Close() {
|
|||
return
|
||||
}
|
||||
d.returnPath.closed.Store(true)
|
||||
for _, entry := range d.table {
|
||||
if entry.flow != nil {
|
||||
entry.flow.close(FlowCloseShutdown)
|
||||
}
|
||||
}
|
||||
for port, nat := range d.ports {
|
||||
if nat != nil {
|
||||
port.DetachReturn(&d.returnPath)
|
||||
|
|
@ -147,7 +169,7 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
|
|||
now := d.now()
|
||||
entry, loaded := d.table[key]
|
||||
if loaded && d.entryExpired(entry, now) {
|
||||
d.removeEntry(key, entry)
|
||||
d.removeEntry(key, entry, FlowCloseTimeout)
|
||||
loaded = false
|
||||
}
|
||||
if loaded {
|
||||
|
|
@ -171,7 +193,7 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
|
|||
if packet.protocol == uint8(header.TCPProtocolNumber) {
|
||||
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||
d.forwardToPort(flow, packet, raw)
|
||||
flow.closed.Store(true)
|
||||
flow.close(FlowCloseReset)
|
||||
d.tombstoneEntry(entry, now)
|
||||
return true
|
||||
}
|
||||
|
|
@ -186,7 +208,7 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
|
|||
case ActionAccept:
|
||||
entry.deadline = now + int64(entry.idle)
|
||||
if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||
d.removeEntry(key, entry)
|
||||
d.removeEntry(key, entry, FlowCloseReset)
|
||||
}
|
||||
return false
|
||||
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 {
|
||||
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 {
|
||||
case ActionFlow:
|
||||
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 {
|
||||
if flow.protocol == uint8(header.TCPProtocolNumber) && flow.finForward && flow.finReverse.Load() {
|
||||
return tcpClosingTimeout
|
||||
}
|
||||
established := flow.established.Load() && !flow.finForward && !flow.finReverse.Load()
|
||||
return d.idleTimeout(flow.protocol, established)
|
||||
}
|
||||
|
|
@ -324,6 +353,12 @@ func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdic
|
|||
flow.reverseRule.sourcePort = clientDestinationPort
|
||||
flow.reverseRule.rewriteSourcePort = true
|
||||
}
|
||||
if verdict.NewTracker != nil {
|
||||
flow.tracker = verdict.NewTracker()
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.AttachFlow(flow)
|
||||
}
|
||||
}
|
||||
nat.insert(reverseKey, flow)
|
||||
return flow, true
|
||||
}
|
||||
|
|
@ -354,6 +389,9 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
|
|||
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) {
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountForward(len(raw))
|
||||
}
|
||||
d.rewriteForward(flow, packet)
|
||||
d.resegmentTCP(flow, packet, raw)
|
||||
return
|
||||
|
|
@ -361,6 +399,9 @@ func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPack
|
|||
if packet.ipVersion == 4 {
|
||||
ipHdr := packet.network.(header.IPv4)
|
||||
if ipHdr.Flags()&header.IPv4FlagDontFragment == 0 {
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountForward(len(raw))
|
||||
}
|
||||
d.rewriteForward(flow, packet)
|
||||
fragments, ok := fragmentIPv4Packet(ipHdr, flow.effectiveMTU)
|
||||
if ok {
|
||||
|
|
@ -382,6 +423,9 @@ func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPack
|
|||
}
|
||||
return
|
||||
}
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountForward(len(raw))
|
||||
}
|
||||
d.rewriteForward(flow, packet)
|
||||
d.stagePort(flow.nat, raw)
|
||||
}
|
||||
|
|
@ -460,10 +504,13 @@ func (d *ForwardDispatcher) tombstoneEntry(entry *flowEntry, now int64) {
|
|||
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)
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
|
@ -484,7 +531,7 @@ func (d *ForwardDispatcher) evictEntries(now int64) {
|
|||
)
|
||||
for key, entry := range d.table {
|
||||
if d.entryExpired(entry, now) {
|
||||
d.removeEntry(key, entry)
|
||||
d.removeEntry(key, entry, FlowCloseTimeout)
|
||||
freed++
|
||||
} else if oldest == nil || entry.deadline < oldest.deadline {
|
||||
oldestKey = key
|
||||
|
|
@ -496,7 +543,7 @@ func (d *ForwardDispatcher) evictEntries(now int64) {
|
|||
}
|
||||
}
|
||||
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() {
|
||||
d.tombstoneEntry(entry, now)
|
||||
} else if d.entryExpired(entry, now) {
|
||||
d.removeEntry(key, entry)
|
||||
d.removeEntry(key, entry, FlowCloseTimeout)
|
||||
}
|
||||
visited++
|
||||
if visited >= flowSweepLimit {
|
||||
|
|
@ -579,6 +626,9 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
|||
if flow.closed.Load() {
|
||||
continue
|
||||
}
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountReverse(len(raw) - headroom)
|
||||
}
|
||||
flow.observeReverse(&parsed, now)
|
||||
if parsed.isTCPSyn() {
|
||||
applyRewriteRaw(&parsed, &flow.reverseRule)
|
||||
|
|
|
|||
|
|
@ -269,6 +269,7 @@ func (h *nfqueueHandler) handlePacket(attr nfqueue.Attribute) int {
|
|||
packet.protocol,
|
||||
packet.source,
|
||||
packet.destination,
|
||||
packet.firstPacket,
|
||||
)
|
||||
|
||||
// Use NfRepeat for bypass/reset so the packet re-enters the chain
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ type ICMPForwarder struct {
|
|||
returnPath icmpForwarderReturn
|
||||
|
||||
flowAccess sync.Mutex
|
||||
flows map[icmpFlowKey]time.Time
|
||||
flows map[icmpFlowKey]*icmpFlow
|
||||
lastSweep time.Time
|
||||
attachedPorts map[Port]bool
|
||||
}
|
||||
|
|
@ -40,12 +40,32 @@ type icmpFlowKey struct {
|
|||
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 {
|
||||
forwarder := &ICMPForwarder{
|
||||
stack: stack,
|
||||
handler: handler,
|
||||
logger: logger,
|
||||
flows: make(map[icmpFlowKey]time.Time),
|
||||
flows: make(map[icmpFlowKey]*icmpFlow),
|
||||
attachedPorts: make(map[Port]bool),
|
||||
}
|
||||
forwarder.returnPath.forwarder = forwarder
|
||||
|
|
@ -56,6 +76,10 @@ func (f *ICMPForwarder) Close() error {
|
|||
f.returnPath.closed.Store(true)
|
||||
f.flowAccess.Lock()
|
||||
defer f.flowAccess.Unlock()
|
||||
for key, flow := range f.flows {
|
||||
flow.close(FlowCloseShutdown)
|
||||
delete(f.flows, key)
|
||||
}
|
||||
for port := range f.attachedPorts {
|
||||
port.DetachReturn(&f.returnPath)
|
||||
delete(f.attachedPorts, port)
|
||||
|
|
@ -71,16 +95,25 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
|||
return false
|
||||
}
|
||||
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(
|
||||
uint8(header.ICMPv4ProtocolNumber),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
|
||||
netip.AddrPortFrom(key.source, identifier),
|
||||
netip.AddrPortFrom(key.destination, identifier),
|
||||
nil,
|
||||
)
|
||||
switch verdict.Action {
|
||||
case ActionReject, ActionDrop:
|
||||
return true
|
||||
case ActionFlow:
|
||||
if f.forwardFlow(verdict.Port, false, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
|
||||
if f.installFlow(key, verdict, pkt) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
|
@ -117,16 +150,26 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
|||
return false
|
||||
}
|
||||
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(
|
||||
uint8(header.ICMPv6ProtocolNumber),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier),
|
||||
netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier),
|
||||
netip.AddrPortFrom(key.source, identifier),
|
||||
netip.AddrPortFrom(key.destination, identifier),
|
||||
nil,
|
||||
)
|
||||
switch verdict.Action {
|
||||
case ActionReject, ActionDrop:
|
||||
return true
|
||||
case ActionFlow:
|
||||
if f.forwardFlow(verdict.Port, true, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) {
|
||||
if f.installFlow(key, verdict, pkt) {
|
||||
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 {
|
||||
return false
|
||||
}
|
||||
inet4Address, inet6Address := port.PortAddresses()
|
||||
portAddress := inet4Address
|
||||
if v6 {
|
||||
if key.v6 {
|
||||
portAddress = inet6Address
|
||||
}
|
||||
if !portAddress.IsValid() || !portAddress.IsUnspecified() {
|
||||
|
|
@ -188,14 +256,29 @@ func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, desti
|
|||
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)
|
||||
for flowKey, cachedFlow := range f.flows {
|
||||
if cachedFlow.closed.Load() {
|
||||
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()
|
||||
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()
|
||||
transportSlice := pkt.TransportHeader().Slice()
|
||||
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, transportSlice...)
|
||||
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 {
|
||||
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()
|
||||
defer f.flowAccess.Unlock()
|
||||
deadline, loaded := f.flows[key]
|
||||
flow, loaded := f.flows[key]
|
||||
if !loaded {
|
||||
return false
|
||||
return nil
|
||||
}
|
||||
if flow.closed.Load() {
|
||||
delete(f.flows, key)
|
||||
return nil
|
||||
}
|
||||
now := time.Now()
|
||||
if now.After(deadline) {
|
||||
if now.After(flow.deadline) {
|
||||
delete(f.flows, key)
|
||||
return false
|
||||
flow.close(FlowCloseTimeout)
|
||||
return nil
|
||||
}
|
||||
f.flows[key] = now.Add(defaultICMPTimeout)
|
||||
return true
|
||||
flow.deadline = now.Add(defaultICMPTimeout)
|
||||
return flow
|
||||
}
|
||||
|
||||
type icmpForwarderReturn struct {
|
||||
|
|
@ -289,9 +379,13 @@ func (f *ICMPForwarder) returnPacket(packet []byte) bool {
|
|||
default:
|
||||
return false
|
||||
}
|
||||
if !f.lookupFlow(key) {
|
||||
flow := f.lookupFlow(key)
|
||||
if flow == nil {
|
||||
return false
|
||||
}
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountReverse(len(packet))
|
||||
}
|
||||
return f.writeBack(packet, header.IPv4ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
||||
case header.IPv6Version:
|
||||
ipHdr := header.IPv6(packet)
|
||||
|
|
@ -308,9 +402,13 @@ func (f *ICMPForwarder) returnPacket(packet []byte) bool {
|
|||
destination: AddrFromAddress(ipHdr.SourceAddress()),
|
||||
identifier: icmpHdr.Ident(),
|
||||
}
|
||||
if !f.lookupFlow(key) {
|
||||
flow := f.lookupFlow(key)
|
||||
if flow == nil {
|
||||
return false
|
||||
}
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountReverse(len(packet))
|
||||
}
|
||||
return f.writeBack(packet, header.IPv6ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress())
|
||||
default:
|
||||
return false
|
||||
|
|
|
|||
|
|
@ -77,7 +77,7 @@ 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)
|
||||
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:
|
||||
r.Complete(true)
|
||||
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 (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:
|
||||
gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer))
|
||||
return false, nil, nil, nil
|
||||
|
|
|
|||
2
tun.go
2
tun.go
|
|
@ -19,7 +19,7 @@ import (
|
|||
)
|
||||
|
||||
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.UDPConnectionHandlerEx
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue