diff --git a/go.mod b/go.mod index fb9cd6e2..5cd712af 100644 --- a/go.mod +++ b/go.mod @@ -7,6 +7,7 @@ require ( github.com/google/btree v1.1.3 github.com/metacubex/fswatch v0.1.1 github.com/metacubex/gvisor v0.0.0-20260807021258-5683e078dbc4 + github.com/metacubex/mipstack v0.0.0-20260910230046-ba762df4c91d github.com/metacubex/nftables v0.0.0-20260426003805-208c2c1ba2cb github.com/metacubex/sing v0.5.7 github.com/sagernet/netlink v0.0.0-20240612041022-b9a21c07ac6a @@ -15,6 +16,7 @@ require ( golang.org/x/exp v0.0.0-20240904232852-e7e105dedf7e golang.org/x/net v0.35.0 golang.org/x/sys v0.30.0 + golang.org/x/time v0.10.0 ) require ( @@ -27,6 +29,5 @@ require ( github.com/pmezard/go-difflib v1.0.0 // indirect github.com/vishvananda/netns v0.0.5 // indirect golang.org/x/sync v0.11.0 // indirect - golang.org/x/time v0.10.0 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index e97607b3..ea0785f0 100644 --- a/go.sum +++ b/go.sum @@ -18,6 +18,8 @@ github.com/metacubex/fswatch v0.1.1 h1:jqU7C/v+g0qc2RUFgmAOPoVvfl2BXXUXEumn6oQux github.com/metacubex/fswatch v0.1.1/go.mod h1:czrTT7Zlbz7vWft8RQu9Qqh+JoX+Nnb+UabuyN1YsgI= github.com/metacubex/gvisor v0.0.0-20260807021258-5683e078dbc4 h1:NqW3qka+vHQWb59Z5ZcXPWEGuBV8qhWEW/QHxN80qOM= github.com/metacubex/gvisor v0.0.0-20260807021258-5683e078dbc4/go.mod h1:mBJW3UXUusd8ZYO5M6mig0uScIey9iIubgdIiQAk2ug= +github.com/metacubex/mipstack v0.0.0-20260910230046-ba762df4c91d h1:hc1OeKdo7YIcmTrzwlJLjNZ4EZDKRHi/Ntv7GdYs2YI= +github.com/metacubex/mipstack v0.0.0-20260910230046-ba762df4c91d/go.mod h1:+bbwALZI0pbi2auSG5A3ptdpV2DZ2eLObziXo+P7oj0= github.com/metacubex/nftables v0.0.0-20260426003805-208c2c1ba2cb h1:wk6mHYPURSUvWcUv72gNP79oiylFsscBSDPJ6ieV6Iw= github.com/metacubex/nftables v0.0.0-20260426003805-208c2c1ba2cb/go.mod h1:73ZrCfhdkW4F2E2GAlta3km/S2RHhFNogCMtWZV2anQ= github.com/metacubex/sing v0.5.7 h1:8OC+fhKFSv/l9ehEhJRaZZAOuthfZo68SteBVLe8QqM= diff --git a/stack.go b/stack.go index d8723967..fcc39276 100644 --- a/stack.go +++ b/stack.go @@ -50,6 +50,8 @@ func NewStack( } else { return NewSystem(options) } + case "mips": + return NewMipstack(options) case "gvisor": return NewGVisor(options) case "mixed": diff --git a/stack_mipstack.go b/stack_mipstack.go new file mode 100644 index 00000000..b08f44c6 --- /dev/null +++ b/stack_mipstack.go @@ -0,0 +1,390 @@ +package tun + +import ( + "context" + "net/netip" + "time" + + "github.com/metacubex/mipstack" + "github.com/metacubex/sing-tun/internal/gtcpip/header" + E "github.com/metacubex/sing/common/exceptions" + "github.com/metacubex/sing/common/logger" +) + +type Mipstack struct { + ctx context.Context + tun Tun + mtu uint32 + recvMsgX bool + inet4Address netip.Addr + inet6Address netip.Addr + inet4LoopbackAddress []netip.Addr + inet6LoopbackAddress []netip.Addr + broadcastAddr netip.Addr + icmpMapping *DirectRouteMapping + handler Handler + logger logger.Logger + stack *mipstack.Stack + icmpSlots chan struct{} +} + +func NewMipstack(options StackOptions) (Stack, error) { + var ( + inet4Address netip.Addr + inet6Address netip.Addr + ) + if len(options.TunOptions.Inet4Address) > 0 { + inet4Address = options.TunOptions.Inet4Address[0].Addr() + } + if len(options.TunOptions.Inet6Address) > 0 { + inet6Address = options.TunOptions.Inet6Address[0].Addr() + } + + s := &Mipstack{ + ctx: options.Context, + tun: options.Tun, + mtu: options.TunOptions.MTU, + recvMsgX: options.TunOptions.EXP_RecvMsgX, + inet4Address: inet4Address, + inet6Address: inet6Address, + inet4LoopbackAddress: options.TunOptions.Inet4LoopbackAddress, + inet6LoopbackAddress: options.TunOptions.Inet6LoopbackAddress, + broadcastAddr: BroadcastAddr(options.TunOptions.Inet4Address), + icmpMapping: NewDirectRouteMapping(options.ICMPTimeout), + icmpSlots: make(chan struct{}, 16), + handler: options.Handler, + logger: options.Logger, + } + return s, nil +} + +func (s *Mipstack) config() mipstack.Config { + // With no local addresses, mipstack's default IPv6 route rejects IPv4-only + // links whose MTU is below IPv6's minimum, so install an IPv4 default route. + var routes []mipstack.Route + if s.mtu != 0 && s.mtu < 1280 && !s.inet6Address.IsValid() { + routes = []mipstack.Route{{Destination: netip.MustParsePrefix("0.0.0.0/0")}} + } + return mipstack.Config{ + Routes: routes, + Promiscuous: true, + MTU: s.mtu, + TCP: mipstack.TCPSocketDefaults{ + KeepAlive: true, + KeepAliveConfig: mipstack.KeepAliveConfig{ + Idle: 15 * time.Second, + Interval: 15 * time.Second, + }, + }, + } +} + +func (s *Mipstack) Start() error { + stack, err := mipstack.New(s.config()) + if err != nil { + return err + } + defer func() { + if s.stack == nil { + _ = stack.Close() + } + }() + _, err = mipstack.NewTCPForwarder(stack, mipstack.TCPForwarderOptions{}, s.forwardTCP) + if err != nil { + return err + } + _, err = mipstack.NewUDPForwarder(stack, mipstack.UDPForwarderOptions{}, s.forwardUDP) + if err != nil { + return err + } + _, err = mipstack.NewICMPForwarder(stack, mipstack.ICMPForwarderOptions{}, s.forwardICMP) + if err != nil { + return err + } + _, err = mipstack.NewIPForwarder(stack, mipstack.IPForwarderOptions{}, func(request *mipstack.IPForwarderRequest) { + _ = request.Reject() + }) + if err != nil { + return err + } + err = stack.Start() + if err != nil { + return err + } + s.stack = stack + go s.readLoop() + go s.writeLoop() + return nil +} + +// Start and Close are called sequentially by the owner, like the other stacks. +func (s *Mipstack) Close() error { + // The listener owns TUN.Close. Do not wait for its blocking Read here. + if s.stack != nil { + return s.stack.Close() + } + return nil +} + +func (s *Mipstack) readLoop() { + device := s.tun + if linuxTUN, isLinuxTUN := device.(LinuxTUN); isLinuxTUN && linuxTUN.FrontHeadroom() > 0 { + s.batchLoopLinux(linuxTUN, linuxTUN.BatchSize()) + return + } + if winTun, isWinTun := device.(WinTun); isWinTun { + s.wintunLoop(winTun) + return + } + offset := 0 + if darwinTUN, isDarwinTUN := device.(DarwinTUN); isDarwinTUN { + if s.recvMsgX { + s.batchLoopDarwin(darwinTUN) + return + } + offset = 4 + } + buffer := make([]byte, int(s.mtu)+offset) + for { + n, err := device.Read(buffer) + if n > offset { + s.processPacket(buffer[offset:n]) + } + if err != nil { + if E.IsClosed(err) { + return + } + s.logger.Error(E.Cause(err, "read packet")) + } + } +} + +func (s *Mipstack) wintunLoop(winTun WinTun) { + for { + packet, release, err := winTun.ReadPacket() + if len(packet) > 0 { + s.processPacket(packet) + } + if release != nil { + release() + } + if err != nil { + if !E.IsClosed(err) { + s.logger.Error(E.Cause(err, "read packet")) + } + return + } + } +} + +func (s *Mipstack) batchLoopLinux(linuxTUN LinuxTUN, batchSize int) { + offset := linuxTUN.FrontHeadroom() + buffers := make([][]byte, batchSize) + for i := range buffers { + buffers[i] = make([]byte, int(s.mtu)+offset) + } + sizes := make([]int, len(buffers)) + for { + n, err := linuxTUN.BatchRead(buffers, offset, sizes) + for i := 0; i < n; i++ { + s.processPacket(buffers[i][offset : offset+sizes[i]]) + } + if err != nil { + if E.IsClosed(err) { + return + } + s.logger.Error(E.Cause(err, "batch read packet")) + } + } +} + +func (s *Mipstack) batchLoopDarwin(darwinTUN DarwinTUN) { + for { + buffers, err := darwinTUN.BatchRead() + for _, buffer := range buffers { + s.processPacket(buffer.Bytes()) + buffer.Release() + } + if err != nil { + if E.IsClosed(err) { + return + } + s.logger.Error(E.Cause(err, "batch read packet")) + } + } +} + +func (s *Mipstack) processPacket(packet []byte) { + destination, ok := mipsPacketDestination(packet) + if !ok { + return + } + // Match sing-tun's LinkEndpointFilter, including reflecting non-unicast + // packets unchanged rather than offering them to protocol forwarders. + if destination == s.broadcastAddr || !destination.IsGlobalUnicast() { + if err := s.writePacket(packet); err != nil { + s.logger.Trace(E.Cause(err, "write packet")) + } + return + } + addresses := s.inet4LoopbackAddress + if destination.Is6() { + addresses = s.inet6LoopbackAddress + } + for _, address := range addresses { + if address != destination { + continue + } + parsed, err := mipstack.ParseIPPacket(packet) + if err != nil { + return + } + if !mipsValidLoopbackPacket(parsed) { + break + } + response := append([]byte(nil), packet...) + if destination.Is4() { + ip := header.IPv4(response) + ip.SetSourceAddr(destination) + ip.SetDestinationAddr(parsed.Source) + } else { + ip := header.IPv6(response) + ip.SetSourceAddr(destination) + ip.SetDestinationAddr(parsed.Source) + } + // Swapping the addresses preserves both the IP checksum and the + // TCP pseudo-header sum, also for fragments and extension headers. + if err := s.writePacket(response); err != nil { + s.logger.Trace(E.Cause(err, "write packet")) + } + return + } + // Invalid packet errors are local to the datagram, not fatal device errors. + _, _ = s.stack.Write([][]byte{packet}, 0) +} + +// Inspect only the base header before non-unicast reflection, just as +// LinkEndpointFilter does. Full protocol validation belongs to mipstack. +func mipsPacketDestination(packet []byte) (netip.Addr, bool) { + switch header.IPVersion(packet) { + case header.IPv4Version: + ip := header.IPv4(packet) + if ip.IsValid(len(packet)) { + return ip.DestinationAddr(), true + } + case header.IPv6Version: + ip := header.IPv6(packet) + if ip.IsValid(len(packet)) { + return ip.DestinationAddr(), true + } + } + return netip.Addr{}, false +} + +// Validate complete packets before bypassing the stack. Fragment checksums +// cannot be verified independently; preserve valid IP fragments for the host +// to reassemble after reflection. +func mipsValidLoopbackPacket(parsed mipstack.IPPacket) bool { + if fragment, ok := parsed.Fragment(); ok && !fragment.IsAtomic() { + if fragment.Offset != 0 { + return fragment.Protocol == mipstack.ProtocolTCP + } + // The first fragment can include extension headers before TCP. + // Inspect those with mipstack, without checking a partial TCP checksum. + parsed.Protocol, parsed.Payload = fragment.Protocol, fragment.Payload + parsed.MoreFragments, parsed.FragmentOffset = false, 0 + protocol, _, err := parsed.UpperLayer() + return err == nil && protocol == mipstack.ProtocolTCP + } + _, err := parsed.TCPSegment() + return err == nil +} + +func (s *Mipstack) writeLoop() { + buffers := make([][]byte, s.stack.BatchSize()) + for i := range buffers { + buffers[i] = make([]byte, int(s.mtu)) + } + sizes := make([]int, len(buffers)) + packets := make([][]byte, len(buffers)) + var writeBuffers [][]byte + for { + n, err := s.stack.Read(buffers, sizes, 0) + for i := 0; i < n; i++ { + packets[i] = buffers[i][:sizes[i]] + } + if writeErr := s.writePacketsWithBuffers(packets[:n], &writeBuffers); writeErr != nil { + if E.IsClosed(writeErr) { + return + } + s.logger.Trace(E.Cause(writeErr, "write packet")) + } + if err != nil { + if !E.IsClosed(err) { + s.logger.Error(E.Cause(err, "read stack packet")) + } + return + } + } +} + +func (s *Mipstack) writePacket(packet []byte) error { + return s.writePackets([][]byte{packet}) +} + +// writePackets owns its scratch storage so output and reflection can run concurrently. +func (s *Mipstack) writePackets(packets [][]byte) error { + var writeBuffers [][]byte + return s.writePacketsWithBuffers(packets, &writeBuffers) +} + +// scratch belongs to the caller; the output loop reuses its own GRO storage. +func (s *Mipstack) writePacketsWithBuffers(packets [][]byte, scratch *[][]byte) error { + if len(packets) == 0 { + return nil + } + // Only GSO devices initialize the GRO tables used by BatchWrite. + if linux, ok := s.tun.(LinuxTUN); ok && linux.FrontHeadroom() > 0 { + offset := linux.FrontHeadroom() + // GRO mutates packets and needs tailroom to append adjacent segments. + for len(*scratch) < len(packets) { + *scratch = append(*scratch, make([]byte, offset+65535)) + } + writeBuffers := (*scratch)[:len(packets)] + for i, packet := range packets { + writeBuffers[i] = writeBuffers[i][:offset+len(packet)] + copy(writeBuffers[i][offset:], packet) + } + _, err := linux.BatchWrite(writeBuffers, offset) + return err + } + if darwin, ok := s.tun.(DarwinTUN); ok { + // Use caller-owned storage instead of NativeTun's shared batch descriptors. + if len(*scratch) == 0 { + *scratch = append(*scratch, nil) + } + for _, packet := range packets { + if cap((*scratch)[0]) < 4+len(packet) { + (*scratch)[0] = make([]byte, 4+len(packet)) + } + buffer := (*scratch)[0][:4+len(packet)] + // Darwin utun requires a four-byte, big-endian address family header. + copy(buffer, []byte{0, 0, 0, 2}) // AF_INET on Darwin. + if header.IPVersion(packet) == header.IPv6Version { + buffer[3] = 30 // AF_INET6 on Darwin, even when tested on another OS. + } + copy(buffer[4:], packet) + if _, err := darwin.Write(buffer); err != nil { + return err + } + } + return nil + } + for _, packet := range packets { + _, err := s.tun.Write(packet) + if err != nil { + return err + } + } + return nil +} diff --git a/stack_mipstack_icmp.go b/stack_mipstack_icmp.go new file mode 100644 index 00000000..534a2f50 --- /dev/null +++ b/stack_mipstack_icmp.go @@ -0,0 +1,84 @@ +package tun + +import ( + "errors" + "time" + + mips "github.com/metacubex/mipstack" + "github.com/metacubex/sing/common/buf" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +func (s *Mipstack) forwardICMP(request *mips.ICMPForwarderRequest) { + message := request.Message() + if !message.IsEchoRequest() { + _ = request.Drop() + return + } + if message.Destination == s.inet4Address || message.Destination == s.inet6Address { + _ = request.ReplyEcho() + _ = request.Drop() + return + } + select { + case s.icmpSlots <- struct{}{}: + default: + _ = request.Drop() + return + } + responder, err := request.Detach() + if err != nil { + <-s.icmpSlots + return + } + go func() { + defer func() { <-s.icmpSlots }() + s.forwardDetachedICMP(responder) + }() +} + +func (s *Mipstack) forwardDetachedICMP(responder *mips.ICMPForwarderResponder) { + message := responder.Message() + writer := &mipsICMPWriter{responder: responder} + action, err := s.icmpMapping.Lookup(DirectRouteSession{Source: message.Source, Destination: message.Destination}, func(timeout time.Duration) (DirectRouteDestination, error) { + destination, err := s.handler.PrepareConnection( + N.NetworkICMP, + M.SocksaddrFrom(message.Source, 0), + M.SocksaddrFrom(message.Destination, 0), + writer, + timeout, + ) + if err != nil { + if destination != nil { + _ = destination.Close() + } + return nil, err + } + return destination, nil + }) + if errors.Is(err, ErrReset) { + _ = responder.Reject() + return + } else if errors.Is(err, ErrDrop) { + _ = responder.Drop() + return + } + if action != nil { + packet := responder.IPPacket() + buffer := buf.NewSize(len(packet)) + _, _ = buffer.Write(packet) + _ = action.WritePacket(buffer) + return + } + _ = responder.ReplyEcho() + _ = responder.Drop() +} + +type mipsICMPWriter struct { + responder *mips.ICMPForwarderResponder +} + +func (w *mipsICMPWriter) WritePacket(packet []byte) error { + return w.responder.ReplyIPPacket(packet) +} diff --git a/stack_mipstack_icmp_test.go b/stack_mipstack_icmp_test.go new file mode 100644 index 00000000..07667c5e --- /dev/null +++ b/stack_mipstack_icmp_test.go @@ -0,0 +1,355 @@ +package tun + +import ( + "bytes" + "errors" + "net/netip" + "sync" + "testing" + "time" + + mips "github.com/metacubex/mipstack" + "github.com/metacubex/sing/common/buf" + "github.com/stretchr/testify/require" +) + +type echoDestination struct { + writer DirectRouteContext + closed chan struct{} + once sync.Once +} + +func (d *echoDestination) IsClosed() bool { + select { + case <-d.closed: + return true + default: + return false + } +} +func (d *echoDestination) Close() error { d.once.Do(func() { close(d.closed) }); return nil } +func (d *echoDestination) WritePacket(b *buf.Buffer) error { + defer b.Release() + source, target, protocol, _ := mipsPacketAddresses(b.Bytes()) + offset := 20 + kind := byte(0) + if source.Is6() { + offset = 40 + kind = 129 + } + payload := append([]byte(nil), b.Bytes()[offset:]...) + payload[0] = kind + payload[2], payload[3] = 0, 0 + packet := transportPacket(target, source, protocol, payload) + go func() { _ = d.writer.WritePacket(packet) }() + return nil +} + +func TestICMPDirectSessionLifecycle(t *testing.T) { + for _, mode := range []string{"expire", "closed"} { + t.Run(mode, func(t *testing.T) { + d := newMemoryTun() + created := make(chan *echoDestination, 4) + h := &testHandler{prepare: func(writer DirectRouteContext) (DirectRouteDestination, error) { + destination := &echoDestination{writer: writer, closed: make(chan struct{})} + t.Cleanup(func() { _ = destination.Close() }) + created <- destination + return destination, nil + }} + testStack(t, d, h, func(o *StackOptions) { o.ICMPTimeout = time.Second }) + source, target := netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("8.8.8.8") + echo := func(sequence byte) { + d.in <- transportPacket(source, target, 1, []byte{8, 0, 0, 0, 0, 1, 0, sequence}) + packet := readPacket(t, d) + require.Equal(t, byte(0), packet[20]) + require.Equal(t, sequence, packet[len(packet)-1]) + } + echo(0) + destination := <-created + echo(1) + select { + case <-created: + t.Fatal("ICMP session not reused") + default: + } + if mode == "closed" { + require.NoError(t, destination.Close()) + } else { + // DirectRouteMapping expires entries on lookup, in whole seconds. + time.Sleep(time.Second) + } + echo(2) + select { + case replacement := <-created: + require.NotSame(t, destination, replacement) + default: + t.Fatal("stale ICMP session reused") + } + require.True(t, destination.IsClosed()) + }) + } +} + +func TestICMPPolicy(t *testing.T) { + for _, policy := range []string{"drop", "reset", "fallback"} { + policy := policy + t.Run(policy, func(t *testing.T) { + d := newMemoryTun() + called := make(chan struct{}, 1) + h := &testHandler{prepare: func(DirectRouteContext) (DirectRouteDestination, error) { + called <- struct{}{} + switch policy { + case "drop": + return nil, ErrDrop + case "reset": + return nil, ErrReset + default: + return nil, errors.New("dial failed") + } + }} + testStack(t, d, h, nil) + d.in <- transportPacket(netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("8.8.8.8"), 1, []byte{8, 0, 0, 0, 0, 1, 0, 1}) + select { + case <-called: + case <-time.After(time.Second): + t.Fatal("policy not called") + } + if policy == "drop" { + select { + case p := <-d.out: + t.Fatalf("drop emitted %x", p) + case <-time.After(50 * time.Millisecond): + } + return + } + p := readPacket(t, d) + want := byte(0) + if policy == "reset" { + want = 3 + } + if p[20] != want { + t.Fatalf("wrong ICMP policy response %x", p) + } + }) + } +} + +func TestICMPAdmissionLimitDoesNotBlockInput(t *testing.T) { + d := newMemoryTun() + started := make(chan struct{}) + release := make(chan struct{}) + var startOnce sync.Once + h := &testHandler{prepare: func(DirectRouteContext) (DirectRouteDestination, error) { + startOnce.Do(func() { close(started) }) + <-release + return nil, nil + }} + s := testStack(t, d, h, nil) + source := netip.MustParseAddr("198.18.0.2") + target := netip.MustParseAddr("8.8.8.8") + packet := transportPacket(source, target, 1, []byte{8, 0, 0, 0, 0, 1, 0, 1}) + + // The first request blocks in PrepareConnection. Processing well beyond the + // admission limit proves packet handling keeps admitting and dropping + // packets instead of waiting for preparation. + s.processPacket(packet) + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("ICMP preparation did not start") + } + sent := make(chan struct{}) + go func() { + for i := 0; i < 64; i++ { + s.processPacket(packet) + } + close(sent) + }() + select { + case <-sent: + case <-time.After(time.Second): + t.Fatal("packet input blocked while ICMP preparation was stalled") + } + deadline := time.After(time.Second) + for len(s.icmpSlots) < cap(s.icmpSlots) { + select { + case <-deadline: + t.Fatalf("ICMP admission did not fill: %d/%d", len(s.icmpSlots), cap(s.icmpSlots)) + default: + time.Sleep(time.Millisecond) + } + } + + close(release) + // Exactly the admitted requests produce fallback echo replies. The rest + // were dropped at admission. + for i := 0; i < cap(s.icmpSlots); i++ { + readPacket(t, d) + } + select { + case extra := <-d.out: + t.Fatalf("excess ICMP request was not dropped: %x", extra) + case <-time.After(100 * time.Millisecond): + } + deadline = time.After(time.Second) + for len(s.icmpSlots) != 0 { + select { + case <-deadline: + t.Fatalf("ICMP admission slots did not drain: %d remaining", len(s.icmpSlots)) + default: + time.Sleep(time.Millisecond) + } + } + + // Once preparation has completed and the slots drain, a fresh request is + // admitted again. + s.processPacket(packet) + readPacket(t, d) +} + +func TestICMPInterfaceAddressBypassesPolicy(t *testing.T) { + for _, family := range []string{"ipv4", "ipv6"} { + t.Run(family, func(t *testing.T) { + source := netip.MustParseAddr("198.18.0.2") + local := netip.MustParseAddr("198.18.0.1") + remote := netip.MustParseAddr("8.8.8.8") + protocol, echo, reply, rejected, offset := byte(1), byte(8), byte(0), byte(3), 20 + if family == "ipv6" { + source = netip.MustParseAddr("fd00::2") + local = netip.MustParseAddr("fd00::1") + remote = netip.MustParseAddr("2001:4860:4860::8888") + protocol, echo, reply, rejected, offset = 58, 128, 129, 1, 40 + } + d := newMemoryTun() + called := make(chan struct{}, 2) + testStack(t, d, &testHandler{prepare: func(DirectRouteContext) (DirectRouteDestination, error) { + called <- struct{}{} + return nil, ErrReset + }}, nil) + payload := []byte{echo, 0, 0, 0, 0, 1, 0, 7, 42} + d.in <- transportPacket(source, local, protocol, payload) + response := readPacket(t, d) + src, dst, proto, ok := mipsPacketAddresses(response) + if !ok || src != local || dst != source || proto != protocol || response[offset] != reply || !bytes.Equal(response[offset+4:], payload[4:]) { + t.Fatalf("invalid interface echo reply: %x", response) + } + select { + case <-called: + t.Fatal("interface echo reached routing policy") + default: + } + d.in <- transportPacket(source, remote, protocol, payload) + response = readPacket(t, d) + if response[offset] != rejected { + t.Fatalf("remote echo bypassed reset policy: %x", response) + } + select { + case <-called: + default: + t.Fatal("remote echo did not reach routing policy") + } + }) + } +} + +func TestMipsICMPResetAdministrativelyProhibited(t *testing.T) { + for _, ipv6 := range []bool{false, true} { + t.Run(map[bool]string{false: "ipv4", true: "ipv6"}[ipv6], func(t *testing.T) { + d := newMemoryTun() + s := testStack(t, d, &testHandler{prepare: func(DirectRouteContext) (DirectRouteDestination, error) { return nil, ErrReset }}, nil) + source, target := netip.MustParseAddr("198.18.0.2"), netip.MustParseAddr("8.8.8.8") + protocol, kind, code, offset, limit := byte(1), byte(8), byte(13), 20, 576 + if ipv6 { + source, target = netip.MustParseAddr("fd00::2"), netip.MustParseAddr("2001:4860::8888") + protocol, kind, code, offset, limit = 58, 128, 1, 40, 1280 + } + payload := make([]byte, 1400) + payload[0] = kind + s.processPacket(transportPacket(source, target, protocol, payload)) + response := readPacket(t, d) + require.LessOrEqual(t, len(response), limit) + require.Equal(t, code, response[offset+1]) + if ipv6 { + require.Equal(t, byte(1), response[offset]) + } else { + require.Equal(t, byte(3), response[offset]) + require.Zero(t, mipsTestChecksum(response[offset:])) + } + }) + } +} + +func TestMipsICMPRepliesOnly(t *testing.T) { + for _, ipv6 := range []bool{false, true} { + for _, mode := range []string{"fallback", "error", "direct"} { + t.Run(map[bool]string{false: "ipv4", true: "ipv6"}[ipv6]+"/"+mode, func(t *testing.T) { + source, target := netip.MustParseAddr("198.18.0.2"), netip.MustParseAddr("8.8.8.8") + protocol, echo, replyType := byte(1), byte(8), byte(0) + if ipv6 { + source, target = netip.MustParseAddr("fd00::2"), netip.MustParseAddr("2001:4860::8888") + protocol, echo, replyType = 58, 128, 129 + } + payload := []byte{echo, 0, 0, 0, 0x12, 0x34, 0x56, 0x78, 42, 43} + input := transportPacket(source, target, protocol, payload) + original := append([]byte(nil), input...) + replyPayload := append([]byte(nil), payload...) + replyPayload[0] = replyType + replyPacket := transportPacket(target, source, protocol, replyPayload) + d := newMemoryTun() + var writer *mipsICMPWriter + prepared := make(chan struct{}) + burst := func() { + results := make(chan error, 4) + for i := 0; i < 4; i++ { + go func() { results <- writer.WritePacket(replyPacket) }() + } + for i := 0; i < 4; i++ { + require.NoError(t, <-results) + } + } + s := testStack(t, d, &testHandler{prepare: func(context DirectRouteContext) (DirectRouteDestination, error) { + writer = context.(*mipsICMPWriter) + close(prepared) + require.NotNil(t, writer.responder.IPPacket()) + require.NotNil(t, writer.responder.Message().Payload) + if mode == "error" { + return nil, errors.New("dial failed") + } + if mode == "fallback" { + return nil, nil + } + // Replies may start before PrepareConnection returns. + burst() + destination := &echoDestination{writer: writer, closed: make(chan struct{})} + t.Cleanup(func() { _ = destination.Close() }) + return destination, nil + }}, nil) + s.processPacket(input) + select { + case <-prepared: + case <-time.After(time.Second): + t.Fatal("ICMP preparation not started") + } + require.Equal(t, original, input, "borrowed request was modified") + for i := range input { + input[i] = 0 // Simulate the delivery buffer being reused. + } + count := 1 + if mode == "direct" { + burst() // The same writer remains usable after the callback. + count = 9 + } + for i := 0; i < count; i++ { + packet, err := mips.ParseIPPacket(readPacket(t, d)) + require.NoError(t, err) + message, err := packet.ICMPMessage() // Validates the family-specific checksum. + require.NoError(t, err) + require.Equal(t, target, message.Source) + require.Equal(t, source, message.Destination) + require.Equal(t, replyType, message.Type) + require.Equal(t, payload[4:], message.Body) + } + }) + } + } +} diff --git a/stack_mipstack_io_test.go b/stack_mipstack_io_test.go new file mode 100644 index 00000000..94291037 --- /dev/null +++ b/stack_mipstack_io_test.go @@ -0,0 +1,409 @@ +package tun + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "net/netip" + "os" + "syscall" + "testing" + "time" + + "github.com/metacubex/sing/common/buf" + "github.com/metacubex/sing/common/logger" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +type linuxTun struct { + *memoryTun + headroom int +} + +func (d *linuxTun) FrontHeadroom() int { return d.headroom } +func (d *linuxTun) BatchSize() int { return 4 } +func (d *linuxTun) TXChecksumOffload() bool { return false } +func (d *linuxTun) BatchRead(buffers [][]byte, offset int, sizes []int) (int, error) { + if d.headroom == 0 { + return 0, errors.New("non-GSO device must use Read") + } + if offset != d.headroom { + return 0, errors.New("wrong read headroom") + } + n, err := d.Read(buffers[0][offset:]) + sizes[0] = n + if err != nil { + return 0, err + } + return 1, nil +} +func (d *linuxTun) BatchWrite(buffers [][]byte, offset int) (int, error) { + if d.headroom == 0 { + return 0, errors.New("non-GSO device must use Write") + } + if offset != d.headroom { + return 0, errors.New("wrong write headroom") + } + var total int + for _, p := range buffers { + n, err := d.Write(p[offset:]) + if err != nil { + return total, err + } + total += n + offset + } + return total, nil +} + +type windowsTun struct { + *memoryTun + released chan struct{} +} + +func (d *windowsTun) ReadPacket() ([]byte, func(), error) { + select { + case p := <-d.in: + return p, func() { + for i := range p { + p[i] = 0 + } + d.released <- struct{}{} + }, nil + case <-d.done: + return nil, nil, io.EOF + } +} + +type darwinTun struct{ *memoryTun } + +func (d *darwinTun) Read(p []byte) (int, error) { + n, err := d.memoryTun.Read(p[4:]) + if err != nil { + return 0, err + } + copy(p[:4], []byte{0, 0, 0, 2}) + if p[4]>>4 == 6 { + p[3] = 30 + } + return n + 4, nil +} + +func (d *darwinTun) Write(p []byte) (int, error) { + if len(p) < 5 { + return 0, errors.New("missing Darwin packet header") + } + family := byte(2) + if p[4]>>4 == 6 { + family = 30 + } + if !bytes.Equal(p[:4], []byte{0, 0, 0, family}) { + return 0, errors.New("invalid Darwin packet header") + } + n, err := d.memoryTun.Write(p[4:]) + return n + 4, err +} + +func (d *darwinTun) BatchRead() ([]*buf.Buffer, error) { + select { + case p := <-d.in: + b := buf.NewSize(len(p)) + _, _ = b.Write(p) + return []*buf.Buffer{b}, nil + case <-d.done: + return nil, io.EOF + } +} +func (d *darwinTun) BatchWrite(buffers []*buf.Buffer) error { + for _, b := range buffers { + if _, err := d.memoryTun.Write(b.Bytes()); err != nil { + return err + } + } + return nil +} + +func TestPlatformPacketIO(t *testing.T) { + for _, name := range []string{"linux", "linux-gso", "windows", "darwin", "darwin-raw"} { + t.Run(name, func(t *testing.T) { + memory := newMemoryTun() + var device Tun = memory + var released chan struct{} + switch name { + case "linux": + device = &linuxTun{memory, 0} + case "linux-gso": + device = &linuxTun{memory, 10} + case "windows": + released = make(chan struct{}, 1) + device = &windowsTun{memory, released} + case "darwin", "darwin-raw": + device = &darwinTun{memory} + } + h := &testHandler{udp: func(_ context.Context, _ netip.AddrPort, b *buf.Buffer, m M.Metadata, init func(N.PacketConn) N.PacketWriter) { + w := init(nil) + go func() { + if released != nil { + <-released + } + _ = w.WritePacket(b, m.Destination) + }() + }} + testStack(t, device, h, func(o *StackOptions) { o.TunOptions.EXP_RecvMsgX = name != "darwin-raw" }) + memory.in <- udpPacket(netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("8.8.8.8"), 53, []byte("owned payload")) + p := readPacket(t, memory) + if string(p[28:]) != "owned payload" { + t.Fatalf("buffer ownership violated: %x", p) + } + }) + } +} + +// Fail one operation, then allow subsequent packets through. +type failedTun struct { + *memoryTun + readFailure error + writeFailure error + failed chan struct{} +} + +func (d *failedTun) Read(p []byte) (int, error) { + if err := d.readFailure; err != nil { + d.readFailure = nil + close(d.failed) + return 0, err + } + return d.memoryTun.Read(p) +} + +func (d *failedTun) Write(p []byte) (int, error) { + if err := d.writeFailure; err != nil { + d.writeFailure = nil + close(d.failed) + return 0, err + } + return d.memoryTun.Write(p) +} + +type failedDarwinTun struct { + DarwinTUN + failed chan struct{} + writeFailure error +} + +func (d *failedDarwinTun) Write(packet []byte) (int, error) { + if err := d.writeFailure; err != nil { + d.writeFailure = nil + close(d.failed) + return 0, err + } + return d.DarwinTUN.Write(packet) +} + +func TestMipsIOErrorRecovery(t *testing.T) { + for _, path := range []string{"read", "write", "darwin-write"} { + for _, reflected := range []bool{false, true} { + name := path + "/stack" + if reflected { + name = path + "/reflection" + } + t.Run(name, func(t *testing.T) { + memory := newMemoryTun() + failed := make(chan struct{}) + d := &failedTun{memoryTun: memory, failed: failed} + var device Tun = d + switch path { + case "read": + d.readFailure = syscall.EINTR + case "write": + d.writeFailure = io.ErrShortWrite + case "darwin-write": + device = &failedDarwinTun{DarwinTUN: &darwinTun{memory}, failed: failed, writeFailure: syscall.EAGAIN} + } + h := &testHandler{udp: func(_ context.Context, _ netip.AddrPort, b *buf.Buffer, m M.Metadata, init func(N.PacketConn) N.PacketWriter) { + if err := init(nil).WritePacket(b, m.Destination); err != nil { + t.Error(err) + } + }} + testStack(t, device, h, func(o *StackOptions) { o.TunOptions.EXP_RecvMsgX = true }) + destination := netip.MustParseAddr("8.8.8.8") + if reflected { + destination = netip.MustParseAddr("224.0.0.1") + } + packet := udpPacket(netip.MustParseAddr("198.18.0.1"), destination, 53, []byte("recovered")) + if path != "read" { + memory.in <- packet + } + select { + case <-failed: + case <-time.After(time.Second): + t.Fatal("I/O failure was not exercised") + } + memory.in <- packet + response := readPacket(t, memory) + if !bytes.Equal(response[28:], []byte("recovered")) { + t.Fatalf("unexpected response: %x", response) + } + }) + } + } +} + +func TestMipsDeviceCloseExitsReadLoop(t *testing.T) { + d := newMemoryTun() + _ = d.Close() + s := &Mipstack{tun: d, logger: logger.NOP()} + done := make(chan struct{}) + go func() { s.readLoop(); close(done) }() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("read loop did not exit") + } +} + +type failedWindowsTun struct { + *windowsTun + readFailure error +} + +func (d *failedWindowsTun) ReadPacket() ([]byte, func(), error) { + if err := d.readFailure; err != nil { + d.readFailure = nil + return nil, nil, err + } + return d.windowsTun.ReadPacket() +} + +func TestMipsWindowsReadFailureExitsReadLoop(t *testing.T) { + d := &failedWindowsTun{ + windowsTun: &windowsTun{memoryTun: newMemoryTun()}, + readFailure: errors.New("send ring corrupt"), + } + t.Cleanup(func() { _ = d.Close() }) + s := &Mipstack{tun: d, logger: logger.NOP()} + done := make(chan struct{}) + go func() { s.readLoop(); close(done) }() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("Windows read loop did not exit") + } +} + +func waitMipsStackClosed(t *testing.T, s *Mipstack) { + t.Helper() + result := make(chan error, 1) + go func() { + _, err := s.stack.Read([][]byte{make([]byte, s.mtu)}, make([]int, 1), 0) + result <- err + }() + select { + case err := <-result: + if !errors.Is(err, net.ErrClosed) && !errors.Is(err, io.EOF) && !errors.Is(err, os.ErrClosed) { + t.Fatalf("expected closed stack, got %v", err) + } + case <-time.After(time.Second): + t.Fatal("stack read was not unblocked by shutdown") + } +} + +// Hold both writes inside the device to expose scratch storage shared by callers. +type overlappingLinuxTun struct { + *linuxTun + entered chan struct{} + release chan struct{} +} + +func (d *overlappingLinuxTun) BatchWrite(packets [][]byte, offset int) (int, error) { + d.entered <- struct{}{} + <-d.release + return d.linuxTun.BatchWrite(packets, offset) +} + +func TestMipsConcurrentOutputAndReflection(t *testing.T) { + d := &overlappingLinuxTun{ + linuxTun: &linuxTun{newMemoryTun(), 10}, + entered: make(chan struct{}, 2), + release: make(chan struct{}), + } + t.Cleanup(func() { _ = d.Close() }) + s := &Mipstack{tun: d} + first, second := []byte("stack output"), []byte("reflected packet") + results := make(chan error, 2) + go func() { + var scratch [][]byte + results <- s.writePacketsWithBuffers([][]byte{first}, &scratch) + }() + go func() { results <- s.writePacket(second) }() + for i := 0; i < 2; i++ { + select { + case <-d.entered: + case <-time.After(time.Second): + close(d.release) + t.Fatal("independent writes were serialized") + } + } + close(d.release) + for i := 0; i < 2; i++ { + if err := <-results; err != nil { + t.Fatal(err) + } + } + seen := map[string]bool{string(readPacket(t, d.memoryTun)): true} + seen[string(readPacket(t, d.memoryTun))] = true + if !seen[string(first)] || !seen[string(second)] { + t.Fatalf("concurrent writes corrupted packet storage: %v", seen) + } +} + +type overlappingDarwinTun struct { + *darwinTun + entered chan struct{} + release chan struct{} +} + +func (d *overlappingDarwinTun) Write(packet []byte) (int, error) { + d.entered <- struct{}{} + <-d.release + return d.darwinTun.Write(packet) +} + +func TestMipsDarwinConcurrentOutputAndReflection(t *testing.T) { + d := &overlappingDarwinTun{ + darwinTun: &darwinTun{newMemoryTun()}, + entered: make(chan struct{}, 2), + release: make(chan struct{}), + } + t.Cleanup(func() { _ = d.Close() }) + s := &Mipstack{tun: d} + first := ipPacket(netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("198.18.0.2"), 253, []byte("output")) + second := ipPacket(netip.MustParseAddr("fd00::1"), netip.MustParseAddr("fd00::2"), 253, []byte("reflection")) + results := make(chan error, 2) + go func() { + var scratch [][]byte + results <- s.writePacketsWithBuffers([][]byte{first}, &scratch) + }() + go func() { results <- s.writePacket(second) }() + for i := 0; i < 2; i++ { + select { + case <-d.entered: + case <-time.After(time.Second): + close(d.release) + t.Fatal("independent Darwin writes were serialized") + } + } + close(d.release) + for i := 0; i < 2; i++ { + if err := <-results; err != nil { + t.Fatal(err) + } + } + seen := map[string]bool{string(readPacket(t, d.memoryTun)): true} + seen[string(readPacket(t, d.memoryTun))] = true + if !seen[string(first)] || !seen[string(second)] { + t.Fatalf("concurrent writes corrupted packet storage: %v", seen) + } +} diff --git a/stack_mipstack_linux_test.go b/stack_mipstack_linux_test.go new file mode 100644 index 00000000..52ba3f05 --- /dev/null +++ b/stack_mipstack_linux_test.go @@ -0,0 +1,30 @@ +//go:build linux + +package tun + +import ( + "context" + "net/netip" + "os" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestMipsLinuxGROMergesOutput(t *testing.T) { + file, err := os.CreateTemp(t.TempDir(), "tun-output") + require.NoError(t, err) + defer file.Close() + device := &NativeTun{tunFile: file, vnetHdr: true, tcpGROTable: newTCPGROTable(), udpGROTable: newUDPGROTable()} + s := &Mipstack{ctx: context.Background(), tun: device} + source, target := netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("198.18.0.2") + first := tcpPacket(source, target, 100, 200, 16, []byte("abcd")) + second := tcpPacket(source, target, 104, 200, 16, []byte("efgh")) + require.NoError(t, s.writePackets([][]byte{first, second})) + data, err := os.ReadFile(file.Name()) + require.NoError(t, err) + // A single virtio/IP/TCP header followed by both payloads proves that the + // batch survived the wrapper and had enough writable tailroom for GRO. + require.Len(t, data, virtioNetHdrLen+40+8) + require.Equal(t, "abcdefgh", string(data[virtioNetHdrLen+40:])) +} diff --git a/stack_mipstack_packet_test.go b/stack_mipstack_packet_test.go new file mode 100644 index 00000000..2c79c6ea --- /dev/null +++ b/stack_mipstack_packet_test.go @@ -0,0 +1,129 @@ +package tun + +import ( + "bytes" + "context" + "encoding/binary" + "net/netip" + "testing" + "time" + + "github.com/metacubex/sing/common/buf" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +func TestUDPFragmentReassemblyAndMTU(t *testing.T) { + d := newMemoryTun() + received := make(chan []byte, 2) + testStack(t, d, &testHandler{udp: func(_ context.Context, _ netip.AddrPort, b *buf.Buffer, m M.Metadata, init func(N.PacketConn) N.PacketWriter) { + received <- append([]byte(nil), b.Bytes()...) + b.Release() + reply := buf.NewSize(2000) + _, _ = reply.Write(bytes.Repeat([]byte{42}, 2000)) + _ = init(nil).WritePacket(reply, m.Destination) + }}, func(o *StackOptions) { o.TunOptions.MTU = 1280 }) + source, target := netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("8.8.8.8") + payload := bytes.Repeat([]byte{7}, 64) + packet := udpPacket(source, target, 53, payload) + for _, part := range []struct { + offset, end int + more bool + }{{0, 24, true}, {24, len(packet) - 20, false}} { + fragment := ipPacket(source, target, 17, packet[20+part.offset:20+part.end]) + binary.BigEndian.PutUint16(fragment[4:], 123) + flags := uint16(part.offset / 8) + if part.more { + flags |= 0x2000 + } + binary.BigEndian.PutUint16(fragment[6:], flags) + fragment[10], fragment[11] = 0, 0 + binary.BigEndian.PutUint16(fragment[10:], mipsTestChecksum(fragment[:20])) + d.in <- fragment + } + select { + case got := <-received: + if !bytes.Equal(got, payload) { + t.Fatal("wrong reassembled datagram") + } + case <-time.After(time.Second): + t.Fatal("fragments not reassembled") + } + var assembled []byte + for { + p := readPacket(t, d) + if len(p) > 1280 || mipsTestChecksum(p[:20]) != 0 { + t.Fatalf("invalid output fragment %x", p) + } + offset := int(binary.BigEndian.Uint16(p[6:])&0x1fff) * 8 + end := offset + len(p) - 20 + if end > len(assembled) { + assembled = append(assembled, make([]byte, end-len(assembled))...) + } + copy(assembled[offset:], p[20:]) + if binary.BigEndian.Uint16(p[6:])&0x2000 == 0 { + break + } + } + if len(assembled) != 2008 || !bytes.Equal(assembled[8:], bytes.Repeat([]byte{42}, 2000)) { + t.Fatal("fragmented response lost data") + } +} + +func TestMalformedPacketsDoNotStopStack(t *testing.T) { + d := newMemoryTun() + received := make(chan struct{}, 4) + testStack(t, d, &testHandler{udp: func(_ context.Context, _ netip.AddrPort, b *buf.Buffer, _ M.Metadata, _ func(N.PacketConn) N.PacketWriter) { + b.Release() + received <- struct{}{} + }}, nil) + p := udpPacket(netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("8.8.8.8"), 53, []byte("valid")) + d.in <- []byte{0x45} + bad := append([]byte(nil), p...) + bad[len(bad)-1] ^= 1 + d.in <- bad + d.in <- p + select { + case <-received: + case <-time.After(time.Second): + t.Fatal("valid packet not delivered after malformed traffic") + } + select { + case <-received: + t.Fatal("invalid checksum reached handler") + case <-time.After(50 * time.Millisecond): + } +} + +func TestCloseDuringTCPHandshake(t *testing.T) { + d := newMemoryTun() + s := testStack(t, d, &testHandler{}, nil) + d.in <- tcpPacket(netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("8.8.8.8"), 100, 0, 2, nil) + _ = readPacket(t, d) + done := make(chan struct{}) + go func() { _ = s.Close(); close(done) }() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("pending handshake prevented shutdown") + } +} + +func TestMipsStartFailure(t *testing.T) { + d := newMemoryTun() + defer d.Close() + base := StackOptions{Context: context.Background(), Tun: d, Handler: &testHandler{}, TunOptions: Options{MTU: 1500, Inet4Address: []netip.Prefix{netip.MustParsePrefix("198.18.0.1/30")}}} + s, err := NewMipstack(base) + if err != nil { + t.Fatal(err) + } + defer s.Close() + // Force the dependency's validation failure before any I/O loop starts. + s.(*Mipstack).mtu = 65536 + if err = s.Start(); err == nil { + t.Fatal("invalid stack configuration started") + } + if s.(*Mipstack).stack != nil { + t.Fatal("failed startup retained stack") + } +} diff --git a/stack_mipstack_regression_test.go b/stack_mipstack_regression_test.go new file mode 100644 index 00000000..782b39f6 --- /dev/null +++ b/stack_mipstack_regression_test.go @@ -0,0 +1,258 @@ +package tun + +import ( + "bytes" + "context" + "encoding/binary" + "net/netip" + "testing" + "time" + + mips "github.com/metacubex/mipstack" + "github.com/metacubex/sing/common/buf" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" + "github.com/stretchr/testify/require" +) + +func TestMipsFamiliesWithoutInterfaceAddresses(t *testing.T) { + for _, configuration := range []string{"none", "ipv4", "ipv6"} { + for _, family := range []string{"ipv4", "ipv6"} { + t.Run(configuration+"/"+family, func(t *testing.T) { + d := newMemoryTun() + s := testStack(t, d, &testHandler{udp: func(_ context.Context, _ netip.AddrPort, b *buf.Buffer, m M.Metadata, init func(N.PacketConn) N.PacketWriter) { + if err := init(nil).WritePacket(b, m.Destination); err != nil { + t.Error(err) + } + }}, func(o *StackOptions) { + if configuration != "ipv4" { + o.TunOptions.Inet4Address = nil + } + if configuration != "ipv6" { + o.TunOptions.Inet6Address = nil + } + }) + source, target := netip.MustParseAddr("198.18.0.2"), netip.MustParseAddr("8.8.8.8") + offset := 20 + if family == "ipv6" { + source, target = netip.MustParseAddr("fd00::2"), netip.MustParseAddr("2001:4860::8888") + offset = 40 + } + s.processPacket(udpPacket(source, target, 53, []byte("reply"))) + packet := readPacket(t, d) + require.Equal(t, "reply", string(packet[offset+8:])) + s.processPacket(tcpPacket(source, target, 100, 0, 2, nil)) + packet = readPacket(t, d) + require.Equal(t, byte(0x12), packet[offset+13]&0x12) + }) + } + } +} + +func TestMipsIPv4SmallMTU(t *testing.T) { + d := newMemoryTun() + testStack(t, d, &testHandler{}, func(o *StackOptions) { + o.TunOptions.MTU = 576 + o.TunOptions.Inet6Address = nil + }) +} + +func TestMipsUnknownProtocolRejection(t *testing.T) { + for _, ipv6 := range []bool{false, true} { + t.Run(map[bool]string{false: "ipv4", true: "ipv6"}[ipv6], func(t *testing.T) { + d := newMemoryTun() + s := testStack(t, d, &testHandler{}, nil) + source, target := netip.MustParseAddr("198.18.0.2"), netip.MustParseAddr("8.8.8.8") + kind, code, offset := byte(3), byte(2), 20 + if ipv6 { + source, target = netip.MustParseAddr("fd00::2"), netip.MustParseAddr("2001:4860::8888") + kind, code, offset = 4, 1, 40 + } + input := ipPacket(source, target, 253, bytes.Repeat([]byte{42}, 64)) + s.processPacket(input) + response := readPacket(t, d) + require.Equal(t, []byte{kind, code}, response[offset:offset+2]) + parsed, err := mips.ParseIPPacket(response) + require.NoError(t, err) + require.Equal(t, source, parsed.Destination) + require.Equal(t, target, parsed.Source) + if ipv6 { + value, err := mips.IPTransportChecksum(parsed.Source, parsed.Destination, 58, parsed.Payload) + require.NoError(t, err) + require.Zero(t, value) + require.Equal(t, uint32(6), binary.BigEndian.Uint32(response[offset+4:])) + } else { + require.Zero(t, mipsTestChecksum(parsed.Payload)) + } + require.True(t, bytes.HasPrefix(input, response[offset+8:])) + // An ICMP error must not trigger another error; IPv6 No Next Header + // is also explicitly silent rather than an unknown protocol. + if ipv6 { + s.processPacket(ipPacket(source, target, 59, nil)) + } + s.processPacket(response) + select { + case p := <-d.out: + t.Fatalf("recursive error: %x", p) + case <-time.After(30 * time.Millisecond): + } + }) + } +} + +func TestMipsLoopbackValidation(t *testing.T) { + for _, ipv6 := range []bool{false, true} { + t.Run(map[bool]string{false: "ipv4", true: "ipv6"}[ipv6], func(t *testing.T) { + d := newMemoryTun() + source, target := netip.MustParseAddr("198.18.0.2"), netip.MustParseAddr("198.18.0.9") + offset := 20 + if ipv6 { + source, target = netip.MustParseAddr("fd00::2"), netip.MustParseAddr("fd00::9") + offset = 40 + } + s := testStack(t, d, &testHandler{}, func(o *StackOptions) { + if ipv6 { + o.TunOptions.Inet6LoopbackAddress = []netip.Addr{target} + } else { + o.TunOptions.Inet4LoopbackAddress = []netip.Addr{target} + } + }) + s.processPacket(ipPacket(source, target, 6, []byte{1, 2})) + bad := tcpPacket(source, target, 100, 0, 2, nil) + bad[offset+12] = 0xf0 + s.processPacket(bad) + bad = tcpPacket(source, target, 100, 0, 2, []byte("bad checksum")) + bad[len(bad)-1] ^= 1 + s.processPacket(bad) + good := tcpPacket(source, target, 100, 0, 2, nil) + if ipv6 { + good = append(append(append([]byte(nil), good[:40]...), 6, 0, 0, 0, 0, 0, 0, 0), good[40:]...) + good[6] = 0 + binary.BigEndian.PutUint16(good[4:], uint16(len(good)-40)) + } + s.processPacket(good) + response := readPacket(t, d) + require.Len(t, response, len(good)) + parsed, err := mips.ParseIPPacket(response) + require.NoError(t, err) + protocol, payload, err := parsed.UpperLayer() + require.NoError(t, err) + value, err := mips.IPTransportChecksum(parsed.Source, parsed.Destination, protocol, payload) + require.NoError(t, err) + require.Zero(t, value) + select { + case p := <-d.out: + t.Fatalf("invalid packet reflected: %x", p) + case <-time.After(30 * time.Millisecond): + } + }) + } +} + +type ringFullTun struct { + *windowsTun + full bool +} + +func (d *ringFullTun) Write(p []byte) (int, error) { + if d.full { + d.full = false + return 0, nil + } + return d.memoryTun.Write(p) +} +func TestMipsWindowsRingFullKeepsStackAlive(t *testing.T) { + d := &ringFullTun{windowsTun: &windowsTun{newMemoryTun(), make(chan struct{}, 1)}, full: true} + s := testStack(t, d, &testHandler{}, nil) + packet := udpPacket(netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("224.0.0.1"), 53, nil) + s.processPacket(packet) + s.processPacket(packet) + require.Equal(t, packet, readPacket(t, d.memoryTun)) +} + +type batchLinuxTun struct { + *linuxTun + counts []int +} + +func (d *batchLinuxTun) BatchWrite(p [][]byte, offset int) (int, error) { + d.counts = append(d.counts, len(p)) + return d.linuxTun.BatchWrite(p, offset) +} + +type batchDarwinTun struct { + *darwinTun + counts []int +} + +func (d *batchDarwinTun) BatchWrite(p []*buf.Buffer) error { + d.counts = append(d.counts, len(p)) + return d.darwinTun.BatchWrite(p) +} +func TestMipsOutputBatching(t *testing.T) { + packets := [][]byte{ipPacket(netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("198.18.0.2"), 253, []byte{1})} + packets = append(packets, packets[0], packets[0]) + t.Run("linux", func(t *testing.T) { + d := &batchLinuxTun{linuxTun: &linuxTun{newMemoryTun(), 10}} + s := testStack(t, d, &testHandler{}, nil) + require.NoError(t, s.writePackets(packets)) + require.Equal(t, []int{3}, d.counts) + for range packets { + require.Equal(t, packets[0], readPacket(t, d.memoryTun)) + } + }) + t.Run("darwin", func(t *testing.T) { + d := &batchDarwinTun{darwinTun: &darwinTun{newMemoryTun()}} + s := testStack(t, d, &testHandler{}, nil) + require.NoError(t, s.writePackets(packets)) + require.Empty(t, d.counts, "Darwin output must not use shared batch descriptors") + for range packets { + require.Equal(t, packets[0], readPacket(t, d.memoryTun)) + } + }) +} + +func TestMipsLoopbackFragments(t *testing.T) { + for _, ipv6 := range []bool{false, true} { + t.Run(map[bool]string{false: "ipv4", true: "ipv6"}[ipv6], func(t *testing.T) { + source, target := netip.MustParseAddr("198.18.0.2"), netip.MustParseAddr("198.18.0.9") + if ipv6 { + source, target = netip.MustParseAddr("fd00::2"), netip.MustParseAddr("fd00::9") + } + d := newMemoryTun() + s := testStack(t, d, &testHandler{}, func(o *StackOptions) { + if ipv6 { + o.TunOptions.Inet6LoopbackAddress = []netip.Addr{target} + } else { + o.TunOptions.Inet4LoopbackAddress = []netip.Addr{target} + } + }) + parsed, err := mips.ParseIPPacket(tcpPacket(source, target, 100, 0, 2, bytes.Repeat([]byte{42}, 1600))) + require.NoError(t, err) + fragments, err := parsed.MarshalFragments(1280, 123) + require.NoError(t, err) + require.Greater(t, len(fragments), 1) + for _, fragment := range fragments { + s.processPacket(fragment) + response := readPacket(t, d) + reflected, err := mips.ParseIPPacket(response) + require.NoError(t, err) + require.Equal(t, target, reflected.Source) + require.Equal(t, source, reflected.Destination) + reflected.Source, reflected.Destination = source, target + original, err := reflected.MarshalRawBinary() + require.NoError(t, err) + require.Equal(t, fragment, original) + } + }) + } +} + +func TestMipsNonUnicastBeforeChecksumValidation(t *testing.T) { + d := newMemoryTun() + s := testStack(t, d, &testHandler{}, nil) + packet := udpPacket(netip.MustParseAddr("198.18.0.2"), netip.MustParseAddr("224.0.0.1"), 53, []byte("reflect")) + packet[10] ^= 1 + s.processPacket(packet) + require.Equal(t, packet, readPacket(t, d)) +} diff --git a/stack_mipstack_tcp.go b/stack_mipstack_tcp.go new file mode 100644 index 00000000..db0db697 --- /dev/null +++ b/stack_mipstack_tcp.go @@ -0,0 +1,24 @@ +package tun + +import ( + mips "github.com/metacubex/mipstack" + M "github.com/metacubex/sing/common/metadata" +) + +func (s *Mipstack) forwardTCP(request *mips.TCPForwarderRequest) { + flow := request.Flow() + conn, err := request.Accept(s.ctx) + if err != nil { + return + } + metadata := M.Metadata{ + Source: M.SocksaddrFromNetIP(flow.Source), + Destination: M.SocksaddrFromNetIP(flow.Destination), + } + go func() { + if err := s.handler.NewConnection(s.ctx, conn, metadata); err != nil { + _ = conn.SetLinger(0) + _ = conn.Close() + } + }() +} diff --git a/stack_mipstack_test.go b/stack_mipstack_test.go new file mode 100644 index 00000000..a3bc2adf --- /dev/null +++ b/stack_mipstack_test.go @@ -0,0 +1,376 @@ +package tun + +import ( + "bytes" + "context" + "encoding/binary" + "errors" + "io" + "net" + "net/netip" + "sync" + "testing" + "time" + + mips "github.com/metacubex/mipstack" + "github.com/metacubex/sing/common/buf" + "github.com/metacubex/sing/common/logger" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +type memoryTun struct { + in, out chan []byte + done chan struct{} + once sync.Once +} + +func newMemoryTun() *memoryTun { + return &memoryTun{in: make(chan []byte, 32), out: make(chan []byte, 32), done: make(chan struct{})} +} +func (d *memoryTun) Read(p []byte) (int, error) { + select { + case packet := <-d.in: + return copy(p, packet), nil + case <-d.done: + return 0, net.ErrClosed + } +} +func (d *memoryTun) Write(p []byte) (int, error) { + select { + case d.out <- append([]byte(nil), p...): + return len(p), nil + case <-d.done: + return 0, net.ErrClosed + } +} +func (d *memoryTun) Close() error { d.once.Do(func() { close(d.done) }); return nil } + +type testHandler struct { + tcp func(context.Context, net.Conn, M.Metadata) error + udp func(context.Context, netip.AddrPort, *buf.Buffer, M.Metadata, func(N.PacketConn) N.PacketWriter) + prepare func(DirectRouteContext) (DirectRouteDestination, error) +} + +func (h *testHandler) NewConnection(ctx context.Context, c net.Conn, m M.Metadata) error { + if h.tcp != nil { + return h.tcp(ctx, c, m) + } + return c.Close() +} +func (h *testHandler) NewPacket(ctx context.Context, k netip.AddrPort, b *buf.Buffer, m M.Metadata, init func(N.PacketConn) N.PacketWriter) { + if h.udp != nil { + h.udp(ctx, k, b, m, init) + } else { + b.Release() + } +} +func (h *testHandler) PrepareConnection(_ string, _, _ M.Socksaddr, writer DirectRouteContext, _ time.Duration) (DirectRouteDestination, error) { + if h.prepare != nil { + return h.prepare(writer) + } + return nil, nil +} +func (h *testHandler) NewError(context.Context, error) {} + +func testStack(t *testing.T, device Tun, handler *testHandler, modify func(*StackOptions)) *Mipstack { + t.Helper() + options := StackOptions{Logger: logger.NOP(), Tun: device, Handler: handler, Context: context.Background(), ICMPTimeout: 50 * time.Millisecond, + TunOptions: Options{MTU: 1500, Inet4Address: []netip.Prefix{netip.MustParsePrefix("198.18.0.1/30")}, Inet6Address: []netip.Prefix{netip.MustParsePrefix("fd00::1/126")}}} + if modify != nil { + modify(&options) + } + v, err := NewStack("mips", options) + if err != nil { + t.Fatal(err) + } + s := v.(*Mipstack) + if err = s.Start(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = s.Close(); _ = device.Close() }) + return s +} +func readPacket(t *testing.T, d *memoryTun) []byte { + t.Helper() + select { + case p := <-d.out: + return p + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for TUN response") + return nil + } +} +func mipsTestChecksum(p []byte) uint16 { + var sum uint32 + for len(p) >= 2 { + sum += uint32(binary.BigEndian.Uint16(p)) + p = p[2:] + } + if len(p) > 0 { + sum += uint32(p[0]) << 8 + } + for sum>>16 != 0 { + sum = (sum & 65535) + (sum >> 16) + } + return ^uint16(sum) +} +func ipPacket(source, target netip.Addr, protocol byte, payload []byte) []byte { + h := 20 + if source.Is6() { + h = 40 + } + p := make([]byte, h+len(payload)) + copy(p[h:], payload) + if h == 20 { + p[0] = 0x45 + p[8] = 64 + p[9] = protocol + binary.BigEndian.PutUint16(p[2:], uint16(len(p))) + copy(p[12:], source.AsSlice()) + copy(p[16:], target.AsSlice()) + binary.BigEndian.PutUint16(p[10:], mipsTestChecksum(p[:20])) + } else { + p[0] = 0x60 + p[6] = protocol + p[7] = 64 + binary.BigEndian.PutUint16(p[4:], uint16(len(payload))) + copy(p[8:], source.AsSlice()) + copy(p[24:], target.AsSlice()) + } + return p +} +func transportPacket(source, target netip.Addr, protocol byte, payload []byte) []byte { + payload = append([]byte(nil), payload...) + pseudo := append(append([]byte(nil), source.AsSlice()...), target.AsSlice()...) + if source.Is4() { + pseudo = append(pseudo, 0, protocol, byte(len(payload)>>8), byte(len(payload))) + } else { + pseudo = append(pseudo, 0, 0, byte(len(payload)>>8), byte(len(payload)), 0, 0, 0, protocol) + } + offset := 2 + if protocol == 17 { + offset = 6 + } + if protocol == 6 { + offset = 16 + } + if protocol == 1 { + pseudo = nil + } + binary.BigEndian.PutUint16(payload[offset:], mipsTestChecksum(append(pseudo, payload...))) + return ipPacket(source, target, protocol, payload) +} +func udpPacket(source, target netip.Addr, port uint16, data []byte) []byte { + p := make([]byte, 8+len(data)) + binary.BigEndian.PutUint16(p, 12345) + binary.BigEndian.PutUint16(p[2:], port) + binary.BigEndian.PutUint16(p[4:], uint16(len(p))) + copy(p[8:], data) + return transportPacket(source, target, 17, p) +} +func tcpPacket(source, target netip.Addr, seq, ack uint32, flags byte, payload []byte) []byte { + p := make([]byte, 20+len(payload)) + binary.BigEndian.PutUint16(p, 12345) + binary.BigEndian.PutUint16(p[2:], 443) + binary.BigEndian.PutUint32(p[4:], seq) + binary.BigEndian.PutUint32(p[8:], ack) + p[12] = 5 << 4 + p[13] = flags + binary.BigEndian.PutUint16(p[14:], 65535) + copy(p[20:], payload) + return transportPacket(source, target, 6, p) +} + +func TestUDPRepliesFromMultipleDestinations(t *testing.T) { + for _, pair := range [][2]string{{"198.18.0.1", "8.8.8.8"}, {"fd00::1", "2001:4860:4860::8888"}} { + t.Run(pair[0], func(t *testing.T) { + d := newMemoryTun() + source, target := netip.MustParseAddr(pair[0]), netip.MustParseAddr(pair[1]) + errorsCh := make(chan error, 4) + h := &testHandler{udp: func(_ context.Context, key netip.AddrPort, b *buf.Buffer, m M.Metadata, init func(N.PacketConn) N.PacketWriter) { + if key.Addr() != source || m.Source.Addr != source || m.Destination.Addr != target { + errorsCh <- errors.New("incorrect metadata") + } + writer := init(nil) // DNS uses this path, and replies asynchronously. + go func() { errorsCh <- writer.WritePacket(b, m.Destination) }() + }} + s := testStack(t, d, h, nil) + for _, port := range []uint16{53, 443} { + d.in <- udpPacket(source, target, port, []byte("hello")) + response := readPacket(t, d) + src, dst, proto, ok := mipsPacketAddresses(response) + if !ok || src != target || dst != source || proto != 17 { + t.Fatalf("wrong UDP response %x", response) + } + h := 20 + if source.Is6() { + h = 40 + } + if binary.BigEndian.Uint16(response[h:]) != port || !bytes.Equal(response[h+8:], []byte("hello")) { + t.Fatalf("wrong UDP payload %x", response) + } + if err := <-errorsCh; err != nil { + t.Fatal(err) + } + } + if err := s.Close(); err != nil { + t.Fatal(err) + } + if err := s.Close(); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestTCPHandshakeAndData(t *testing.T) { + for _, pair := range [][2]string{{"198.18.0.1", "8.8.8.8"}, {"fd00::1", "2001:4860:4860::8888"}} { + t.Run(pair[0], func(t *testing.T) { + d := newMemoryTun() + source, target := netip.MustParseAddr(pair[0]), netip.MustParseAddr(pair[1]) + result := make(chan error, 1) + h := &testHandler{tcp: func(_ context.Context, c net.Conn, m M.Metadata) error { + defer c.Close() + _ = c.SetDeadline(time.Now().Add(2 * time.Second)) + if m.Source.Addr != source || m.Destination.Addr != target { + result <- errors.New("wrong TCP metadata") + return nil + } + p := make([]byte, 5) + _, err := io.ReadFull(c, p) + if err == nil && !bytes.Equal(p, []byte("hello")) { + err = errors.New("wrong TCP data") + } + if err == nil { + _, err = c.Write([]byte("world")) + } + result <- err + return err + }} + testStack(t, d, h, nil) + d.in <- tcpPacket(source, target, 100, 0, 2, nil) + synAck := readPacket(t, d) + offset := 20 + if source.Is6() { + offset = 40 + } + if synAck[offset+13]&18 != 18 { + t.Fatalf("expected SYN ACK: %x", synAck) + } + ack := binary.BigEndian.Uint32(synAck[offset+4:]) + 1 + d.in <- tcpPacket(source, target, 101, ack, 16, nil) + d.in <- tcpPacket(source, target, 101, ack, 24, []byte("hello")) + for { + p := readPacket(t, d) + tcpOffset := int(p[offset+12]>>4) * 4 + if bytes.Contains(p[offset+tcpOffset:], []byte("world")) { + break + } + } + if err := <-result; err != nil { + t.Fatal(err) + } + }) + } +} + +func TestICMPEchoAndFilter(t *testing.T) { + for _, pair := range [][2]string{{"198.18.0.1", "8.8.8.8"}, {"fd00::1", "2001:4860:4860::8888"}} { + t.Run(pair[0], func(t *testing.T) { + d := newMemoryTun() + testStack(t, d, &testHandler{}, nil) + source, target := netip.MustParseAddr(pair[0]), netip.MustParseAddr(pair[1]) + protocol, kind, reply, offset := byte(1), byte(8), byte(0), 20 + if source.Is6() { + protocol, kind, reply, offset = 58, 128, 129, 40 + } + p := transportPacket(source, target, protocol, []byte{kind, 0, 0, 0, 1, 2, 3, 4, 5, 6}) + d.in <- p + response := readPacket(t, d) + src, dst, _, ok := mipsPacketAddresses(response) + if !ok || src != target || dst != source || response[offset] != reply || !bytes.Equal(response[offset+4:], p[offset+4:]) { + t.Fatalf("wrong echo %x", response) + } + }) + } + d := newMemoryTun() + testStack(t, d, &testHandler{}, func(o *StackOptions) { + o.TunOptions.Inet4LoopbackAddress = []netip.Addr{netip.MustParseAddr("10.0.0.1")} + }) + source := netip.MustParseAddr("198.18.0.1") + for _, target := range []string{"198.18.0.3", "224.0.0.1", "127.0.0.1"} { + p := udpPacket(source, netip.MustParseAddr(target), 53, []byte("filter")) + d.in <- p + if !bytes.Equal(readPacket(t, d), p) { + t.Fatal("filter did not reflect packet") + } + } + p := tcpPacket(source, netip.MustParseAddr("10.0.0.1"), 100, 0, 2, nil) + d.in <- p + r := readPacket(t, d) + src, dst, _, _ := mipsPacketAddresses(r) + if src.String() != "10.0.0.1" || dst != source || mipsTestChecksum(r[:20]) != 0 { + t.Fatal("invalid loopback rewrite") + } +} + +func TestCloseUnblocksStackOutput(t *testing.T) { + d := newMemoryTun() + s := testStack(t, d, &testHandler{}, nil) + closed := make(chan struct{}) + go func() { _ = s.Close(); close(closed) }() + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("Close waited for TUN Read") + } + outputDone := make(chan struct{}) + go func() { s.writeLoop(); close(outputDone) }() + select { + case <-outputDone: + case <-time.After(time.Second): + t.Fatal("closed stack did not stop output loop") + } + // The owner has not closed TUN yet, so its read loop can still reflect packets. + packet := udpPacket(netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("224.0.0.1"), 53, []byte("reflection")) + d.in <- packet + if response := readPacket(t, d); !bytes.Equal(response, packet) { + t.Fatal("stack close changed TUN reflection") + } +} + +func mipsPacketAddresses(packet []byte) (source, destination netip.Addr, protocol byte, ok bool) { + parsed, err := mips.ParseIPPacket(packet) + if err != nil { + return + } + upper, _, err := parsed.UpperLayer() + return parsed.Source, parsed.Destination, byte(upper), err == nil +} + +func TestMipsContextLifecycle(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + d := newMemoryTun() + s := testStack(t, d, &testHandler{}, func(o *StackOptions) { o.Context = ctx }) + if s.ctx != ctx { + t.Fatal("stack did not preserve the caller context") + } + if err := s.Close(); err != nil { + t.Fatal(err) + } + if err := s.ctx.Err(); err != nil { + t.Fatalf("stack close canceled the handler context: %v", err) + } + waitMipsStackClosed(t, s) + + cancel() + d2 := newMemoryTun() + testStack(t, d2, &testHandler{}, func(o *StackOptions) { o.Context = ctx }) + packet := udpPacket(netip.MustParseAddr("198.18.0.1"), netip.MustParseAddr("224.0.0.1"), 53, []byte("reflection")) + d2.in <- packet + if response := readPacket(t, d2); !bytes.Equal(response, packet) { + t.Fatal("caller cancellation changed packet reflection") + } +} diff --git a/stack_mipstack_udp.go b/stack_mipstack_udp.go new file mode 100644 index 00000000..82d4d6ec --- /dev/null +++ b/stack_mipstack_udp.go @@ -0,0 +1,47 @@ +package tun + +import ( + "os" + + mips "github.com/metacubex/mipstack" + "github.com/metacubex/sing/common/buf" + E "github.com/metacubex/sing/common/exceptions" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +func (s *Mipstack) forwardUDP(request *mips.UDPForwarderRequest) { + flow := request.Flow() + buffer := buf.NewSize(len(request.Payload())) + _, _ = buffer.Write(request.Payload()) + responder, err := request.DetachForReplies() + if err != nil { + buffer.Release() + return + } + s.handler.NewPacket( + s.ctx, + flow.Source, + buffer, + M.Metadata{ + Source: M.SocksaddrFromNetIP(flow.Source), + Destination: M.SocksaddrFromNetIP(flow.Destination), + }, + func(N.PacketConn) N.PacketWriter { + return &mipsUDPWriter{responder: responder} + }, + ) +} + +type mipsUDPWriter struct { + responder *mips.UDPForwarderResponder +} + +func (w *mipsUDPWriter) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { + defer buffer.Release() + if !destination.IsIP() { + return E.Cause(os.ErrInvalid, "invalid destination") + } + _, err := w.responder.ReplyFrom(buffer.Bytes(), destination.AddrPort()) + return err +}