Minor fixes
This commit is contained in:
parent
14c8f75f7a
commit
c17af6ee8c
13 changed files with 357 additions and 181 deletions
192
flow_dispatch.go
192
flow_dispatch.go
|
|
@ -106,6 +106,7 @@ type ForwardDispatcher struct {
|
|||
lastSweep int64
|
||||
ports map[Port]*portNAT
|
||||
natList atomic.Pointer[[]*portNAT]
|
||||
revNAT atomic.Pointer[map[netip.Addr]*portNAT]
|
||||
|
||||
activeNATs []*portNAT
|
||||
writebackBatch [][]byte
|
||||
|
|
@ -113,6 +114,14 @@ type ForwardDispatcher struct {
|
|||
|
||||
segmentBuffers [][]byte
|
||||
segmentSizes []int
|
||||
segmentUsed int
|
||||
}
|
||||
|
||||
func addrToTCPIP(addr netip.Addr) tcpip.Address {
|
||||
if addr.Is4() {
|
||||
return tcpip.AddrFrom4(addr.As4())
|
||||
}
|
||||
return tcpip.AddrFrom16(addr.As16())
|
||||
}
|
||||
|
||||
func NewForwardDispatcher(handler Handler, writeback ForwardWriteback, logger logger.Logger, udpTimeout time.Duration, icmpTimeout time.Duration) *ForwardDispatcher {
|
||||
|
|
@ -206,10 +215,16 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for
|
|||
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, FlowCloseReset)
|
||||
if packet.protocol == uint8(header.TCPProtocolNumber) {
|
||||
if packet.tcpFlags&header.TCPFlagRst != 0 {
|
||||
d.removeEntry(key, entry, FlowCloseReset)
|
||||
return false
|
||||
}
|
||||
if packet.tcpFlags&header.TCPFlagSyn == 0 {
|
||||
entry.idle = tcpEstablishedTimeout
|
||||
}
|
||||
}
|
||||
entry.deadline = now + int64(entry.idle)
|
||||
return false
|
||||
case ActionReject:
|
||||
entry.deadline = now + int64(entry.idle)
|
||||
|
|
@ -330,24 +345,24 @@ func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdic
|
|||
dnatPort: serverPort != clientDestinationPort && !isICMP,
|
||||
}
|
||||
flow.forwardRule = rewriteRule{
|
||||
sourceAddress: tcpip.AddrFromSlice(portAddress.AsSlice()),
|
||||
sourceAddress: addrToTCPIP(portAddress),
|
||||
sourcePort: selector,
|
||||
rewriteSourcePort: true,
|
||||
}
|
||||
if flow.dnatAddress {
|
||||
flow.forwardRule.destinationAddress = tcpip.AddrFromSlice(serverAddress.AsSlice())
|
||||
flow.forwardRule.destinationAddress = addrToTCPIP(serverAddress)
|
||||
}
|
||||
if flow.dnatPort {
|
||||
flow.forwardRule.destinationPort = serverPort
|
||||
flow.forwardRule.rewriteDestinationPort = true
|
||||
}
|
||||
flow.reverseRule = rewriteRule{
|
||||
destinationAddress: tcpip.AddrFromSlice(flow.clientAddress.AsSlice()),
|
||||
destinationAddress: addrToTCPIP(flow.clientAddress),
|
||||
destinationPort: flow.clientSelector,
|
||||
rewriteDestinationPort: true,
|
||||
}
|
||||
if flow.dnatAddress {
|
||||
flow.reverseRule.sourceAddress = tcpip.AddrFromSlice(clientDestinationAddress.AsSlice())
|
||||
flow.reverseRule.sourceAddress = addrToTCPIP(clientDestinationAddress)
|
||||
}
|
||||
if flow.dnatPort {
|
||||
flow.reverseRule.sourcePort = clientDestinationPort
|
||||
|
|
@ -371,7 +386,6 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
|
|||
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)
|
||||
|
|
@ -383,6 +397,20 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT {
|
|||
}
|
||||
natList = append(natList, nat)
|
||||
d.natList.Store(&natList)
|
||||
revMap := make(map[netip.Addr]*portNAT)
|
||||
if currentRev := d.revNAT.Load(); currentRev != nil {
|
||||
for addr, existing := range *currentRev {
|
||||
revMap[addr] = existing
|
||||
}
|
||||
}
|
||||
v4Address, v6Address := port.PortAddresses()
|
||||
if v4Address.IsValid() {
|
||||
revMap[v4Address] = nat
|
||||
}
|
||||
if v6Address.IsValid() {
|
||||
revMap[v6Address] = nat
|
||||
}
|
||||
d.revNAT.Store(&revMap)
|
||||
return nat
|
||||
}
|
||||
|
||||
|
|
@ -473,6 +501,12 @@ func (d *ForwardDispatcher) Flush() {
|
|||
d.flushPort(nat)
|
||||
}
|
||||
d.activeNATs = d.activeNATs[:0]
|
||||
if retain := max(d.segmentUsed, segmentRetainCount); len(d.segmentBuffers) > retain {
|
||||
clear(d.segmentBuffers[retain:])
|
||||
d.segmentBuffers = d.segmentBuffers[:retain]
|
||||
d.segmentSizes = d.segmentSizes[:retain]
|
||||
}
|
||||
d.segmentUsed = 0
|
||||
if len(d.writebackBatch) > 0 {
|
||||
err := d.writeback.WriteReturnPackets(d.writebackBatch)
|
||||
if err != nil {
|
||||
|
|
@ -581,6 +615,14 @@ func (r *forwardReturn) ReturnHeadroom() int {
|
|||
return r.dispatcher.writeback.ReturnHeadroom()
|
||||
}
|
||||
|
||||
type returnDecision uint8
|
||||
|
||||
const (
|
||||
returnPass returnDecision = iota
|
||||
returnWrite
|
||||
returnDrop
|
||||
)
|
||||
|
||||
func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
||||
if r.closed.Load() {
|
||||
return packets
|
||||
|
|
@ -590,65 +632,96 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
|
|||
return packets
|
||||
}
|
||||
natList := *natListPtr
|
||||
var revMap map[netip.Addr]*portNAT
|
||||
if revPtr := r.dispatcher.revNAT.Load(); revPtr != nil {
|
||||
revMap = *revPtr
|
||||
}
|
||||
headroom := r.dispatcher.writeback.ReturnHeadroom()
|
||||
now := r.dispatcher.now()
|
||||
|
||||
if len(packets) == 1 {
|
||||
switch r.classifyReturn(packets[0], natList, revMap, headroom, now) {
|
||||
case returnWrite:
|
||||
if err := r.dispatcher.writeback.WriteReturnPackets(packets[:1]); err != nil {
|
||||
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
|
||||
}
|
||||
return packets[:0]
|
||||
case returnDrop:
|
||||
return packets[:0]
|
||||
default:
|
||||
return packets
|
||||
}
|
||||
}
|
||||
|
||||
unconsumed := packets[:0]
|
||||
var writeBatch [][]byte
|
||||
now := r.dispatcher.now()
|
||||
for _, raw := range packets {
|
||||
if len(raw) < headroom+header.IPv4MinimumSize {
|
||||
switch r.classifyReturn(raw, natList, revMap, headroom, now) {
|
||||
case returnWrite:
|
||||
writeBatch = append(writeBatch, raw)
|
||||
case returnDrop:
|
||||
default:
|
||||
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
|
||||
}
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountReverse(len(raw) - headroom)
|
||||
}
|
||||
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 {
|
||||
if err := r.dispatcher.writeback.WriteReturnPackets(writeBatch); err != nil {
|
||||
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
|
||||
}
|
||||
}
|
||||
return unconsumed
|
||||
}
|
||||
|
||||
func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
|
||||
func (r *forwardReturn) classifyReturn(raw []byte, natList []*portNAT, revMap map[netip.Addr]*portNAT, headroom int, now int64) returnDecision {
|
||||
if len(raw) < headroom+header.IPv4MinimumSize {
|
||||
return returnPass
|
||||
}
|
||||
parsed, ok := parseForwardPacket(raw[headroom:])
|
||||
if !ok || parsed.fragment {
|
||||
return returnPass
|
||||
}
|
||||
if !parsed.hasFlow {
|
||||
if parsed.isICMPError() && returnICMPError(natList, revMap, &parsed) {
|
||||
return returnWrite
|
||||
}
|
||||
return returnPass
|
||||
}
|
||||
flow := findReverseFlow(natList, revMap, parsed.flowKey())
|
||||
if flow == nil {
|
||||
return returnPass
|
||||
}
|
||||
if flow.closed.Load() {
|
||||
return returnDrop
|
||||
}
|
||||
if flow.tracker != nil {
|
||||
flow.tracker.CountReverse(len(raw) - headroom)
|
||||
}
|
||||
flow.observeReverse(&parsed, now)
|
||||
if parsed.isTCPSyn() {
|
||||
applyRewriteRaw(&parsed, &flow.reverseRule)
|
||||
clampTCPMSS(&parsed, flow.effectiveMTU)
|
||||
recomputeChecksums(&parsed)
|
||||
} else {
|
||||
applyRewrite(&parsed, &flow.reverseRule)
|
||||
}
|
||||
return returnWrite
|
||||
}
|
||||
|
||||
func findReverseFlow(natList []*portNAT, revMap map[netip.Addr]*portNAT, key flowKey) *forwardFlow {
|
||||
if nat, ok := revMap[key.destination.Addr()]; ok {
|
||||
if flow := nat.lookup(key); flow != nil {
|
||||
return flow
|
||||
}
|
||||
}
|
||||
for _, nat := range natList {
|
||||
if flow := nat.lookup(key); flow != nil {
|
||||
return flow
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func returnICMPError(natList []*portNAT, revMap map[netip.Addr]*portNAT, parsed *forwardPacket) bool {
|
||||
inner, ok := parsed.icmpErrorInner()
|
||||
if !ok {
|
||||
return false
|
||||
|
|
@ -657,20 +730,13 @@ func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
|
|||
if !parsedInner {
|
||||
return false
|
||||
}
|
||||
key := embedded.flowKey().reversed()
|
||||
var flow *forwardFlow
|
||||
for _, nat := range natList {
|
||||
flow = nat.lookup(key)
|
||||
if flow != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
flow := findReverseFlow(natList, revMap, embedded.flowKey().reversed())
|
||||
if flow == nil || flow.closed.Load() {
|
||||
return false
|
||||
}
|
||||
rewriteEmbeddedSource(&embedded, tcpip.AddrFromSlice(flow.clientAddress.AsSlice()), flow.clientSelector, true)
|
||||
rewriteEmbeddedSource(&embedded, addrToTCPIP(flow.clientAddress), flow.clientSelector, true)
|
||||
if flow.dnatAddress || flow.dnatPort {
|
||||
rewriteEmbeddedDestination(&embedded, tcpip.AddrFromSlice(flow.clientDestinationAddress.AsSlice()), flow.clientDestinationPort, flow.dnatPort)
|
||||
rewriteEmbeddedDestination(&embedded, addrToTCPIP(flow.clientDestinationAddress), flow.clientDestinationPort, flow.dnatPort)
|
||||
}
|
||||
parsed.network.SetDestinationAddr(flow.clientAddress)
|
||||
if parsed.network.SourceAddr() == flow.serverAddress {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue