diff --git a/gvisor_udp.go b/gvisor_udp.go index fff72e0..0c02823 100644 --- a/gvisor_udp.go +++ b/gvisor_udp.go @@ -23,6 +23,11 @@ type UDPForwarder struct { ctx context.Context stack *stack.Stack udpNat *udpnat.Service[netip.AddrPort] + + // cache + cacheProto tcpip.NetworkProtocolNumber + cacheID stack.TransportEndpointID + cachePacket stack.PacketBufferPtr } func NewUDPForwarder(ctx context.Context, stack *stack.Stack, handler Handler, udpTimeout int64) *UDPForwarder { @@ -37,24 +42,37 @@ func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt stack.Pack var upstreamMetadata M.Metadata upstreamMetadata.Source = M.SocksaddrFrom(AddrFromAddress(id.RemoteAddress), id.RemotePort) upstreamMetadata.Destination = M.SocksaddrFrom(AddrFromAddress(id.LocalAddress), id.LocalPort) - var netProto tcpip.NetworkProtocolNumber if upstreamMetadata.Source.IsIPv4() { - netProto = header.IPv4ProtocolNumber + f.cacheProto = header.IPv4ProtocolNumber } else { - netProto = header.IPv6ProtocolNumber + f.cacheProto = header.IPv6ProtocolNumber } + gBuffer := pkt.ToBuffer() + sBuffer := buf.NewSize(int(gBuffer.Size())) + gBuffer.Apply(func(view *buffer.View) { + sBuffer.Write(view.AsSlice()) + }) + f.cacheID = id + f.cachePacket = pkt f.udpNat.NewPacket( f.ctx, upstreamMetadata.Source.AddrPort(), - buf.As(pkt.Data().AsRange().ToSlice()), + sBuffer, upstreamMetadata, - func(natConn N.PacketConn) N.PacketWriter { - return &UDPBackWriter{f.stack, id.RemoteAddress, id.RemotePort, netProto} - }, + f.newUDPConn, ) return true } +func (f *UDPForwarder) newUDPConn(natConn N.PacketConn) N.PacketWriter { + return &UDPBackWriter{ + stack: f.stack, + source: f.cacheID.RemoteAddress, + sourcePort: f.cacheID.RemotePort, + sourceNetwork: f.cacheProto, + } +} + type UDPBackWriter struct { stack *stack.Stack source tcpip.Address diff --git a/lwip.go b/lwip.go index 1b52bed..42cb651 100644 --- a/lwip.go +++ b/lwip.go @@ -52,7 +52,7 @@ func (l *LWIP) loopIn() { l.loopInWintun(winTun) return } - buffer := make([]byte, int(l.tunMtu) + PacketOffset) + buffer := make([]byte, int(l.tunMtu)+PacketOffset) for { n, err := l.tun.Read(buffer) if err != nil { diff --git a/system.go b/system.go index 48699a6..a454dd5 100644 --- a/system.go +++ b/system.go @@ -134,7 +134,7 @@ func (s *System) tunLoop() { s.wintunLoop(winTun) return } - packetBuffer := make([]byte, s.mtu + PacketOffset) + packetBuffer := make([]byte, s.mtu+PacketOffset) for { n, err := s.tun.Read(packetBuffer) if err != nil { diff --git a/tun_darwin_gvisor.go b/tun_darwin_gvisor.go index 6a9c784..85e3a62 100644 --- a/tun_darwin_gvisor.go +++ b/tun_darwin_gvisor.go @@ -51,7 +51,7 @@ func (e *DarwinEndpoint) Attach(dispatcher stack.NetworkDispatcher) { } func (e *DarwinEndpoint) dispatchLoop() { - packetBuffer := make([]byte, e.tun.mtu + 4) + packetBuffer := make([]byte, e.tun.mtu+4) for { n, err := e.tun.tunFile.Read(packetBuffer) if err != nil {