Improve wintun read

This commit is contained in:
世界 2022-08-08 21:34:32 +08:00
parent 0fd822f913
commit d378b6ca53
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
2 changed files with 27 additions and 20 deletions

View file

@ -136,26 +136,33 @@ func (t *NativeTun) configure() error {
} }
func (t *NativeTun) Read(p []byte) (n int, err error) { func (t *NativeTun) Read(p []byte) (n int, err error) {
err = t.ReadFunc(func(b []byte) {
n = copy(p, b)
})
return
}
func (t *NativeTun) ReadFunc(block func(b []byte)) error {
t.running.Add(1) t.running.Add(1)
defer t.running.Done() defer t.running.Done()
retry: retry:
if atomic.LoadInt32(&t.close) == 1 { if atomic.LoadInt32(&t.close) == 1 {
return 0, os.ErrClosed return os.ErrClosed
} }
start := nanotime() start := nanotime()
shouldSpin := atomic.LoadUint64(&t.rate.current) >= spinloopRateThreshold && uint64(start-atomic.LoadInt64(&t.rate.nextStartTime)) <= rateMeasurementGranularity*2 shouldSpin := atomic.LoadUint64(&t.rate.current) >= spinloopRateThreshold && uint64(start-atomic.LoadInt64(&t.rate.nextStartTime)) <= rateMeasurementGranularity*2
for { for {
if atomic.LoadInt32(&t.close) == 1 { if atomic.LoadInt32(&t.close) == 1 {
return 0, os.ErrClosed return os.ErrClosed
} }
packet, err := t.session.ReceivePacket() packet, err := t.session.ReceivePacket()
switch err { switch err {
case nil: case nil:
packetSize := len(packet) packetSize := len(packet)
n = copy(p, packet) block(packet)
t.session.ReleaseReceivePacket(packet) t.session.ReleaseReceivePacket(packet)
t.rate.update(uint64(packetSize)) t.rate.update(uint64(packetSize))
return n, nil return nil
case windows.ERROR_NO_MORE_ITEMS: case windows.ERROR_NO_MORE_ITEMS:
if !shouldSpin || uint64(nanotime()-start) >= spinloopDuration { if !shouldSpin || uint64(nanotime()-start) >= spinloopDuration {
windows.WaitForSingleObject(t.readWait, windows.INFINITE) windows.WaitForSingleObject(t.readWait, windows.INFINITE)
@ -164,11 +171,11 @@ retry:
procyield(1) procyield(1)
continue continue
case windows.ERROR_HANDLE_EOF: case windows.ERROR_HANDLE_EOF:
return 0, os.ErrClosed return os.ErrClosed
case windows.ERROR_INVALID_DATA: case windows.ERROR_INVALID_DATA:
return 0, errors.New("send ring corrupt") return errors.New("send ring corrupt")
} }
return 0, fmt.Errorf("read failed: %w", err) return fmt.Errorf("read failed: %w", err)
} }
} }

View file

@ -3,9 +3,6 @@
package tun package tun
import ( import (
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"gvisor.dev/gvisor/pkg/bufferv2" "gvisor.dev/gvisor/pkg/bufferv2"
"gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header" "gvisor.dev/gvisor/pkg/tcpip/header"
@ -59,29 +56,32 @@ func (e *WintunEndpoint) Attach(dispatcher stack.NetworkDispatcher) {
} }
func (e *WintunEndpoint) dispatchLoop() { func (e *WintunEndpoint) dispatchLoop() {
_buffer := buf.StackNewSize(int(e.tun.mtu))
defer common.KeepAlive(_buffer)
buffer := common.Dup(_buffer)
defer buffer.Release()
data := buffer.FreeBytes()
for { for {
n, err := e.tun.Read(data) var buffer bufferv2.Buffer
err := e.tun.ReadFunc(func(b []byte) {
buffer = bufferv2.MakeWithData(b)
})
if err != nil { if err != nil {
break break
} }
packet := data[:n] ihl, ok := buffer.PullUp(0, 1)
if !ok {
buffer.Release()
continue
}
var networkProtocol tcpip.NetworkProtocolNumber var networkProtocol tcpip.NetworkProtocolNumber
switch header.IPVersion(packet) { switch header.IPVersion(ihl.AsSlice()) {
case header.IPv4Version: case header.IPv4Version:
networkProtocol = header.IPv4ProtocolNumber networkProtocol = header.IPv4ProtocolNumber
case header.IPv6Version: case header.IPv6Version:
networkProtocol = header.IPv6ProtocolNumber networkProtocol = header.IPv6ProtocolNumber
default: default:
e.tun.Write(packet) e.tun.Write(buffer.Flatten())
buffer.Release()
continue continue
} }
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: bufferv2.MakeWithData(packet), Payload: buffer,
IsForwardedPacket: true, IsForwardedPacket: true,
}) })
dispatcher := e.dispatcher dispatcher := e.dispatcher